o
    ;ήcE                     @   s   d dl Z d dlmZ d dlZd dlZd dlmZ d dlmZ ddlm	Z	 ddl
mZmZmZmZmZmZmZmZmZmZmZmZ ddlmZ G d	d
 d
eZG dd dZG dd de	ZdS )    N)Enum)TensorProto)onnx_pb   )ONNXQuantizer)DEQUANT_OP_NAMEQUANT_OP_NAMEQuantizedValueQuantizedValueType__producer____version__add_dequant_output_suffixadd_dequant_suffixadd_quant_input_suffixadd_quant_output_suffixadd_quant_suffixfind_by_name)CreateQDQQuantizerc                   @   s   e Zd ZdZdZdZdS )QDQQuantTensorTyper   r      N)__name__
__module____qualname__
ACTIVATIONWEIGHTBIAS r   r   M/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/quantization/qdq_quantizer.pyr       s    r   c                   @   s   e Zd ZejddfddZdS )QDQTensorQuantInfoNc                 C   s    || _ || _|| _|d u| _d S N)tensor_typequant_para_provideraxis	is_shared)selfr    r!   r"   r   r   r   __init__'   s   zQDQTensorQuantInfo.__init__)r   r   r   r   r   r%   r   r   r   r   r   &   s    r   c                   @   s   e Zd Z	d'ddZdd ZdejfddZd'dd	Zd'd
dZ	dd Z
d(ddZdd Zdd Zdd Zdd Z	d'ddZd'ddZdd Zdd  Zd!d" Zd#d$ Zd%d& ZdS ))QDQQuantizerNc                 C   s   t | |||||||||	|
|| i | _g | _g | _d|vr g n|d | _d|vr+dn|d | _d|vr6dn|d | _| jrAi | _d|vrJi | _	d S |d | _	d S )N#OpTypesToExcludeOutputQuantizatioinAddQDQPairToWeightFDedicatedQDQPair QDQOpTypePerChannelSupportToAxis)
r   r%   tensors_to_quantizebias_to_quantizenodes_to_remove'op_types_to_exclude_output_quantizationadd_qdq_pair_to_weightdedicated_qdq_pairtensor_to_its_receiving_nodes'qdq_op_type_per_channel_support_to_axis)r$   modelper_channelreduce_rangemodestaticweight_qTypeactivation_qTypetensors_rangenodes_to_quantizenodes_to_excludeop_types_to_quantizeextra_optionsr   r   r   r%   /   sB   	
zQDQQuantizer.__init__c                 C   s~   t || j }|dur|jtjjkrdS dS || j v r5| j| }|j	
dr3|j	jjtjkr3dS dS td| dS )z2
        Check if tensor can be quantized
        NTr    z\failed to infer the type of tensor: {}. Skip to quantize it. Please check if it is expected.F)r   r3   initializer	data_type
onnx_protor   FLOATvalue_infoskeystypeHasFieldr    	elem_typeloggingwarningformat)r$   tensor_nameweightvir   r   r   _is_tensor_quantizables   s    
z#QDQQuantizer._is_tensor_quantizablec                 C   sJ   |  |r!|rt||d| j|< dS || jvr#t|d| j|< dS dS dS )a  
        Quantize tensors. If quant_param_tensor is not None, tensor with name tensor_name will be quantized with same
        quantization parameters as tensor quant_param_tensor

        Args:
            tensor_name: name of the tensor to quantize
            quant_sharing_param: name of the tensor that provides quantization parameter
            tensor_type: QDQQuantTensorType default ACTIVATION
        )r    r!   )r    N)rN   r   r+   )r$   rK   quant_sharing_paramr    r   r   r   __quantize_tensor   s   


zQDQQuantizer.__quantize_tensorc                 C      |  ||tjS )z
        Quantize Activation Tensor
        Args:
            tensor_name: name of the tensor to quantize
            quant_sharing_param: name of the tensor that provides quantization parameter

        )_QDQQuantizer__quantize_tensorr   r   r$   rK   rO   r   r   r   quantize_activation_tensor      z'QDQQuantizer.quantize_activation_tensorc                 C   rQ   )z
        Quantize Weight Tensor
        Args:
            tensor_name: name of the tensor to quantize
            quant_sharing_param: name of the tensor that provides quantization parameter

        )rR   r   r   rS   r   r   r   quantize_weight_tensor   rU   z#QDQQuantizer.quantize_weight_tensorc                 C   sR   t || j }|r|jtjjkrttj	|d| j
|< d S d S td| d S )N)r    r"   zMonly support per-channel quantization on weight. Tensor: {} is not quantized.)r   r3   r?   r@   rA   r   rB   r   r   r   r+   rH   rI   rJ   )r$   rK   r"   rL   r   r   r   "quantize_weight_tensor_per_channel   s   z/QDQQuantizer.quantize_weight_tensor_per_channel      ?c                 C   sV   t || j }|d ur!|jtjjkr| j||||f d S d S t	
d| d S )NzExpected {} to be a weight)r   r3   r?   r@   rA   r   rB   r,   appendrH   rI   rJ   )r$   	bias_name
input_nameweight_namebetarL   r   r   r   quantize_bias_tensor   s   z!QDQQuantizer.quantize_bias_tensorc                 C   s   | j | d S r   )r-   rY   )r$   noder   r   r   remove_node   s   zQDQQuantizer.remove_nodec                 C   s   | j | j d S r   )r3   remove_nodesr-   )r$   r   r   r   ra      s   zQDQQuantizer.remove_nodesc                 C   s   | j  D ]+}| |r0t| |}|  | jr0|jD ]}|| jvr'g | j|< | j| | qq| 	  | 
  |   |   | jsI| j   t| j j _t| j j _| j j S r   )r3   nodesshould_quantize_noder   quantizer0   inputr1   rY   _quantize_normal_tensors_quantize_sharing_param_tensors_quantize_bias_tensorsra   r/   clean_initializersr   producer_namer   producer_version)r$   r_   op_quantizerrK   r   r   r   quantize_model   s&   







zQDQQuantizer.quantize_modelc                 C   sd   || j  v r0t| j | dkr0| j|s0| j|s0| j|| || jv r.| j|= dS dS )Nr   TF)	quantization_paramsrD   lenr3   input_name_to_nodesis_graph_outputis_graph_inputreplace_output_of_all_nodesr+   )r$   upstream_output_nameoutput_namer   r   r   try_replacing_upstream_output   s   


z*QDQQuantizer.try_replacing_upstream_outputc
                 C   sP   t jjt|||g|g||	d}
t jjt|||g|g||	d}| j|
|g d S )Nr"   )onnxhelper	make_noder   r   r3   	add_nodes)r$   q_inputq_outputquant_node_namedq_input	dq_outputdequant_node_name
scale_namezp_namer"   qlinear_nodedequant_noder   r   r   _create_qdq_nodes   s   zQDQQuantizer._create_qdq_nodesc                 C   s   |j }|d ur | jdk rtd| j|tjj|| jd\}}}n| j||t	j
u r+| jn| j| jd\}}}t|}| j|| | jrZt|}	| ||	t||	|t||||	 d S tjjt|||g|gt||d}
| j|
 d S )N   zLPer-Channel support with QDQ format requires onnx opset version 13 or above.)keep_float_weightrw   )nameopset_version
ValueErrorquantize_weight_per_channelrA   r   INT8r/   quantize_initializerr   r   r8   r9   r   r3   replace_input_of_all_nodesr   r   r   r   rx   ry   rz   r   add_node)r$   weight_protor    r"   r\   q_weight_namer   r   weight_dequant_outputweight_quant_outputr   r   r   r   _add_qdq_pair_for_initializer   sF   
z*QDQQuantizer._add_qdq_pair_for_initializerc                 C   sT  | j re|| jv ret| j| dkret| j| }t|D ]F}d|d  }t|| }t|| }| ||t|||t||| | j| | }	| j	
|	|| |dkrbt||||tj}
|
| j|< qd S |}t|}| j	|rt|}|}| j	|| n| j	|| | |t|t|t||t||| t||||tj}
|
| j|< d S )Nr   _r   )r0   r1   ro   ranger   r   r   r   r   r3   replace_node_inputr	   r
   Inputquantized_value_maprq   r   rs   r   )r$   rK   r   r   num_dedicated_qdq_pairipostfix tensor_name_quant_output_postfix"tensor_name_dequant_output_postfixr_   quantized_valuer|   r   r   r   r   _add_qdq_pair_for_activation)  sv   
z)QDQQuantizer._add_qdq_pair_for_activationc           
      C   s   | j   D ]K\}}|| j v rq|jsRt|| j }|r*| 	||j
|j n$| |\}}| |||\}}}}	}	|sGtd| d| ||| | j |= qd S )Nz4Quantization parameters are not specified for param zb. In static mode quantization params for inputs and outputs of nodes to be quantized are required.)r+   copyitemsr   rD   r#   r   r3   r?   r   r    r"   find_quant_scale_zp_get_quantization_paramsr   r   )
r$   rK   tensor_infor?   
used_scaleused_zp
data_foundr   r   r   r   r   r   rf   h  s&   
z%QDQQuantizer._quantize_normal_tensorsc                 C   s   | j r>| j   D ].\}}|j}|| jv r8| j |= | j| }t|| j }|d ur/td| 	||j
|j q
| j sd S d S )NzBQuantization parameter shared mode is not supported for weight yet)r+   r   r   r!   r   r   r3   r?   r   r   r   r   )r$   rK   r   tensor_provider_namer   r?   r   r   r   rg     s   

z,QDQQuantizer._quantize_sharing_param_tensorsc           	      C   s   | j D ]V\}}}}|| j v rq| |||| | jt|| j  | j| }|j|j	|j
g}t|}|jd urItjjd||g||jd}n
tjd||g|}| j| qd S )NDequantizeLinearrw   )r,   r   rD   quantize_bias_staticr3   remove_initializerr   r?   q_namer   r   r   r"   rx   ry   rz   r   )	r$   rZ   r[   r\   r]   quant_valueinputs	node_namer   r   r   r   rh     s0   

z#QDQQuantizer._quantize_bias_tensorsc                 C   s   || j v p	|| jv S r   )r+   r,   )r$   rK   r   r   r   is_tensor_quantized  s   z QDQQuantizer.is_tensor_quantizedr   )rX   )r   r   r   r%   rN   r   r   rR   rT   rV   rW   r^   r`   ra   rm   rv   r   r   r   rf   rg   rh   r   r   r   r   r   r&   .   s*    
D






)?r&   )rH   enumr   rx   onnx.numpy_helperr   r   rA   onnx_quantizerr   quant_utilsr   r   r	   r
   r   r   r   r   r   r   r   r   registryr   r   r   r&   r   r   r   r   <module>   s   8