o
    ;ήc-                     @   s   d dl Z d dlZd dlZd dl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mZmZmZmZmZ ddlmZ G dd dZdS )	    N)onnx_pb   )	ONNXModel)TENSOR_NAME_QUANT_SUFFIXQuantizationModeQuantizedValueQuantizedValueType	QuantType__producer____version__add_infer_metadataattribute_to_kwargcompute_scale_zpfind_by_nameget_qmin_qmax_for_qTypeget_qrange_for_qTypemodel_has_infer_metadataquantize_datasave_and_reload_modeltensor_proto_to_array)CreateOpQuantizerc                   @   s.  e Zd Z	dDddZdd Zdd Zdd	 Zd
d Zdd Zdd Z	dd Z
dd Zdd Zdd Zdd Zdd Zdd Zdd Zd d! ZdEd"d#ZdEd$d%Zd&d' Zd(d) Zd*d+ ZdFd-d.Zd/d0 ZdGd2d3Z	1	1	4	1dHd5d6Z	7	1	1	4	1dId8d9ZdJd:d;Z	7	1dKd<d=Zd>d? Zd@dA Z dBdC Z!dS )LONNXQuantizerNc                 C   sH  t |st|}dd |jjD | _| jdd |jjD  | jdd |jjD  t|| _	|s8| j	
  || _|| _|| _|| _d| _|rK|ni | _d| jv oW| jd | _d| jv ob| jd | _d| jv om| jd | _|tjk}d	| jvr{|n| jd	 | _d
| jvrdn| jd
 | _|tjkrtjjntjj| _|tjkrtjjntjj| _	 || _|	| _|
| _ || _!g | _"d | _#d| _$i | _%| j%dd |jjD  | j%dd |jjD  | j	j	jj&D ]}| j%dd |jD  q| ' | _(| jt)vrt*d+| j| , | _-d| _.d| _/d| _0d| _1i | _2| j	3 | _4i | _5d S )Nc                 S      i | ]}|j |qS  name).0vir   r   N/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/quantization/onnx_quantizer.py
<dictcomp>7       z*ONNXQuantizer.__init__.<locals>.<dictcomp>c                 S   r   r   r   r   otr   r   r   r   8   r    c                 S   r   r   r   r   itr   r   r   r   9   r    FEnableSubgraphForceQuantizeNoInputCheckMatMulConstBOnlyWeightSymmetricActivationSymmetric/c                 S      i | ]}|j d qS r   r   r!   r   r   r   r   n   r    c                 S   r+   r,   r   r#   r   r   r   r   o   r    c                 S   s   i | ]}|d qS r,   r   )r   output_namer   r   r   r   q   s    z unsupported quantization mode {}fixed_quantization_range_uint8fixed_quantization_range_int8
fixed_zerofixed_zero_zp)6r   r   graph
value_infovalue_infosupdateoutputinputr   modelreplace_gemm_with_matmulper_channelreduce_rangemodestaticfuse_dynamic_quantextra_optionsenable_subgraph_quantizationforce_quantize_no_input_checkq_matmul_const_b_onlyr	   QInt8is_weight_symmetricis_activation_symmetric
onnx_protoTensorProtoINT8UINT8activation_qTypeweight_qTypetensors_rangenodes_to_quantizenodes_to_excludeop_types_to_quantize	new_nodesparentgraph_scopetensor_namesnodecheck_opset_versionopset_versionr   
ValueErrorformatcalculate_quantization_paramsquantization_paramsfixed_qrange_uint8_namefixed_qrange_int8_namefixed_zero_namefixed_zero_zp_namequantized_value_mapget_non_initializer_inputsgenerated_value_namesused_scale_zp_map)selfr8   r:   r;   r<   r=   rK   rJ   rL   rM   rN   rO   r?   is_weight_int8rT   r   r   r   __init__%   sh   





zONNXQuantizer.__init__c                 C   s|   t jj|d| jjjd}t| t|| j| j| j	| j
| j| j| j| j| j| j| j}| |_d| j||_|  |jjjS )z
        generate submodel for the subgraph, so that we re-utilize current quantization implementation.
        quantize the submodel
        update subgraph and set it back to node
        zonnx-quantizer)producer_nameopset_importsz{}{}/)onnxhelper
make_modelr8   opset_importr   r   r:   r;   r<   r=   rK   rJ   rL   rM   rN   rO   r?   rQ   rX   rR   quantize_modelr2   )rc   subgraph	graph_keywarped_modelsub_quanitzerr   r   r   quantize_subgraph   s0   
zONNXQuantizer.quantize_subgraphc           	      C   s  dd |j D }t|dkr|S |jdkr|jn	d|jt| j}i }|j D ]I}|jtjj	kr@|j| 
|jd||ji}n+|jtjjkrgg }|jD ]}|| 
|d||jt|g qL|j|i}nt|}|| q'tjj|j|j|jfd|ji|S )	z|
        Check subgraph, if any, quantize it and replace it.
        return new_nodes added for quantizing subgraph
        c                 S   s,   g | ]}|j tjjks|j tjjkr|qS r   )typerh   AttributeProtoGRAPHGRAPHS)r   attrr   r   r   
<listcomp>   s    z>ONNXQuantizer.quantize_node_with_sub_graph.<locals>.<listcomp>r    z{}_node_count_{}z{}:{}z{}:{}:{}r   )	attributelenr   rX   op_typerP   rr   rh   rs   rt   rq   gru   graphsextendr   r5   ri   	make_noder7   r6   )	rc   rT   graph_attrs	node_namekwargsrv   kvvaluerm   r   r   r   quantize_node_with_sub_graph   s0   $
 
$z*ONNXQuantizer.quantize_node_with_sub_graphc                 C   s   dd | j j jD }dt|krtd|d j}|dkr'td| dS |dk rMtd| | j j j|d  | j j j	t
jd	d
g d
}d| _|S )Nc                 S   s    g | ]}|j r|j d kr|qS )zai.onnx)domain)r   opsetr   r   r   rw      s    z5ONNXQuantizer.check_opset_version.<locals>.<listcomp>r   z$Failed to find proper ai.onnx domainr   
   zThe original model opset version is {}, which does not support node fusions. Please update the model to opset >= 11 for better performance.zThe original model opset version is {}, which does not support quantization. Please update the model to opset >= 11. Updating the model automatically to opset 11. Please verify the quantized model.rx      T)r8   rk   rz   rW   versionloggingwarningrX   remover~   rh   ri   make_opsetidr>   )rc   ai_onnx_domainrV   r   r   r   rU      s0   
z!ONNXQuantizer.check_opset_versionc                 C   s   t dd | j D S )zQ
        Detect if model already has QuantizeLinear or DequantizeLinear.
        c                 s   s$    | ]}|j d kp|j dkV  qdS )QuantizeLinearDequantizeLinearN)r{   )r   rT   r   r   r   	<genexpr>   s    
z.ONNXQuantizer.has_QDQ_nodes.<locals>.<genexpr>)anyr8   nodesrc   r   r   r   has_QDQ_nodes   s   zONNXQuantizer.has_QDQ_nodesc                 C   s2   t || j d urdS | jd ur| j|S dS )NTF)r   r8   initializerrQ   find_initializer_in_path)rc   initializer_namer   r   r   r      s
   
z&ONNXQuantizer.find_initializer_in_pathc                 C   s2   | j | |D ]}|jD ]}| j| qqd S N)rP   r~   r6   ra   add)rc   r   rT   r-   r   r   r   add_new_nodes   s   
zONNXQuantizer.add_new_nodesc                 C   s   |   r	td | j D ]2}| jr| |}t| j}t	| |}|
  t|t| jD ]}| j| jD ]}| j| q6q.q|   | j d | j j| j | jd u rq| j \}}t|dkrqtdt| t| jj_t| jj_| jjS )NzPlease check if the model is already quantized.Note you don't need to quantize a QAT model. OnnxRuntime support to run QAT model directly.rT   r   z0Invalid model with unknown initializers/tensors.)r   r   r   r8   r   r@   r   rz   rP   r   quantizeranger6   ra   r   _dequantize_outputsr2   
ClearFieldrT   r~   rQ   clean_initializersRuntimeErrorstrr
   rf   r   producer_version)rc   rT   number_of_existing_new_nodesop_quantizerir-   _initializers_not_foundr   r   r   rl      s2   





zONNXQuantizer.quantize_modelc                 C   s   t || j }|d uS r   )r   r8   r   )rc   
input_namer   r   r   r   is_input_a_initializer$  s   z$ONNXQuantizer.is_input_a_initializerc                 C   s   | j S r   )r:   r   r   r   r   is_per_channel(  s   zONNXQuantizer.is_per_channelc                 C   sF   t || j }|d ur|jtjjkS | jr| jd u rdS | j	|S )NF)
r   r8   r   	data_typerF   rG   FLOATr@   rQ   is_valid_quantize_weight)rc   weight_nameweightr   r   r   r   +  s   z&ONNXQuantizer.is_valid_quantize_weightc                 C   s~   |  |r
| |S || j v r)| j| }|jdr'|jjjtj	j
kr'dS dS | jr5| jr5| j|S td| dS )Ntensor_typeTzzFailed to infer data type of tensor: {}. Please add data type info for this tensor if your model has customized operators.F)r   r   r4   keysrr   HasFieldr   	elem_typerF   rG   r   r@   rQ   is_float_tensorr   r   rX   )rc   tensor_namer   r   r   r   r   3  s   


	zONNXQuantizer.is_float_tensorc                 C   sV   | j d urt| j dkr|j| j vrdS |j| jvrdS | jd ur)|j| jv r)dS dS )Nr   FT)rM   rz   r   r{   rO   rN   )rc   rT   r   r   r   should_quantize_nodeE  s   
z"ONNXQuantizer.should_quantize_nodec                 C   s$   |t jjkr| ||S | ||S )aZ  
        Create nodes for dynamic quantization of input and add them to nodes_list.
            parameter input_name: Name of the input.
            parameter nodes_list: new nodes are appended to this list.
            parameter qType: type to quantize to.
            return: scale_name, zero_point_name, scale_shape, zero_point_shape.
        )rF   rG   rH   +_get_dynamic_input_quantization_params_int8,_get_dynamic_input_quantization_params_uint8)rc   r   
nodes_listqTyper   r   r   &_get_dynamic_input_quantization_paramsU  s   z4ONNXQuantizer._get_dynamic_input_quantization_paramsc                 C   s  t jj}|d }|d }tjjd|g|d g|dd}|| |d }tjjd|g|d g|dd}|| |d	 }	tjd
|jd g|	d g|	}
||
 |d	 }tjd
|jd g|d g|}|| |d }tjd|
jd |jd g|d g|}|| tj| j	t jj
g t|d g}| j| |d }tjd|jd | j	g|g|}|| tj| j|g dg}| j| || jg g fS )a/  
        Create nodes for dynamic quantization of input to int8 and add them to nodes_list
            parameter input_name: Name of the input.
            parameter nodes_list: new nodes are appended to this list.
            return: scale_name, zero_point_name, scale_shape, zero_point_shape.
        _scale
_ReduceMin	ReduceMin:0r   keepdims
_ReduceMax	ReduceMax_AbsAbs_Abs_MaxMaxg       @	scale_DivDiv)rF   rG   rH   rh   ri   r   appendr6   make_tensorr\   r   r   r8   add_initializerr^   )rc   r   r   r   input_scale_namereduce_min_namereduce_min_nodereduce_max_namereduce_max_nodereduce_min_abs_namereduce_min_abs_nodereduce_max_abs_namereduce_max_abs_nodeabs_max_nameabs_max_nodeinitializer_divscale_div_namescale_div_nodeinitializer_zpr   r   r   r   b  s|   







z9ONNXQuantizer._get_dynamic_input_quantization_params_int8c                 C   s  t jj}|d }|d }|d }tjjd|g|d g|dd}|| |d }tjjd	|g|d g|dd}	||	 tj| jt jj	g t
|g}
| j|
 tj| jt jj	g d
g}| j| |d }tjd|	jd |jd g|d g|}|| |d }tjd|jd | jg|g|}|| |d }tjd| j|jd g|d g|}|| |d }tjd|jd |g|d g|}|| |d }tjd|j|d g|}|| |d }tjjd|j|g||d}|| ||g g fS )a0  
        Create nodes for dynamic quantization of input to uint8 and add them to nodes_list
            parameter input_name: Name of the input.
            parameter nodes_list: new nodes are appended to this list.
            return: scale_name, zero_point_name, scale_shape, zero_point_shape.
        r   _zero_pointr   r   r   r   r   r   r   g        
_scale_SubSub
_scale_Divr   _zero_point_Sub_zero_point_Div_zero_point_FloorFloor_zero_point_CastCast)to)rF   rG   rI   rh   ri   r   r   r   r[   r   r   r8   r   r]   r6   )rc   r   r   r   r   input_zp_namer   r   r   r   initializer_qrangeinitializer_qvaluescale_sub_namescale_sub_noder   r   zp_sub_namezp_sub_nodezp_div_namezp_div_nodezp_floor_namezp_floor_nodezp_cast_namezp_cast_noder   r   r   r     s   







z:ONNXQuantizer._get_dynamic_input_quantization_params_uint8c                 C   s   |du s|du r>| j du s|| j vrtd| dS | j | }|du s+t|dkr3td|||d g}|d g}n|g}|g}g }|d }| j}	g }
|d	 }tj	||	||}| j
| tj	|tjj|
|}| j
| d
|||
|fS )a\  
        Create initializers and inputs in the graph for zero point and scale of output.
        Zero point and scale values are obtained from self.quantization_params if specified.
            parameter param_name: Name of the quantization parameter.
            return: result, scale_name, zero_point_name, scale_shape, zero_point_shape.
        Nz5Quantization parameters for tensor:"{}" not specified)Frx   rx   rx   rx      z_Quantization parameters should contain zero point and scale. Specified values for output {}: {}r   r   r   r   T)rZ   r   inforX   rz   rW   rJ   rh   ri   r   r8   r   rF   rG   r   )rc   
param_name	use_scaleuse_zeropointparamszero_point_valuesscale_valueszero_point_shapezero_point_namezero_point_typescale_shape
scale_nameinit_zp
init_scaler   r   r   _get_quantization_params  s0   

z&ONNXQuantizer._get_quantization_paramsc                 C   s  |j | }|t }|d }|dur|durd||}	}
}n
| |\}	}
}}}g }|	r:tjd||
|g|g|}n<| jr?dS | jr^|tj	j
kr^|d }
|d }tjd|g||
|g|}n| |||\}
}}}tjd||
|g|g|}t|||
||| j|< ||g S )a  
        Given an input for a node (which is not a initializer), this function

        - add nodes to compute zero point and scale for this input if they don't exist.
        - add new QuantizeLinear node to quantize the input.

        :param node: node being quantized in NodeProto format.
        :param input_index: index of input in node.input.
        :param qType: type to quantize to.
        :param given_scale_name: if those inputs need to be quanitzed using this scale tensor.
        :param given_zp_name: if those inputs to be quantized using this zeropoint tensor.
        :return: List of newly created nodes in NodeProto format.
        _QuantizeLinearNTr   r   r   DynamicQuantizeLinear)r7   r   r  rh   ri   r   r=   r>   rF   rG   rI   r   r   r_   )rc   rT   input_indexr   given_scale_namegiven_zp_namer   r-   ql_node_name
data_foundr  zp_namer   r   qlinear_noder  zp_shaper   r   r   _get_quantize_input_nodes9  sN   

z'ONNXQuantizer._get_quantize_input_nodesc                 C   sD   t |trt|dksJ d|| jvsJ | d|| j|< d S )Nr   z(value must be scale(float) and zeropointz has been setted before)
isinstancetuplerz   rb   )rc   r   r   r   r   r   set_quant_scale_zpw  s   z ONNXQuantizer.set_quant_scale_zpc                 C   s.   || j v r
| j | S | jd ur| j|S dS )NNN)rb   rQ   find_quantized_valuerc   r   r   r   r   find_quant_scale_zp|  
   


z!ONNXQuantizer.find_quant_scale_zpc                 C   s.   || j v r
| j | S | jd ur| j|S d S r   )r_   rQ   r  r  r   r   r   r    r  z"ONNXQuantizer.find_quantized_value      ?c                 C   s  || j v r| j | jS | j | j}t|| j }t|}t|| j }t|}	|t }
|| j v r9| j | j}n|| jv rI| 	|\}}}}}nt
d|t|| j }t|}|| | }t|	|  tj}tj|tjd|j}tj||
}| j |g |
d }tj|tjdd}tj||}| j |g |
d }tj|jtjdd}tj||}| j |g || j vsJ t||
||tj|jdkrdnd}|| j |< |
S )	z]
        Quantized the bias. Zero Point == 0 and Scale == Input_Scale * Weight_Scale
        z@Expected {} to be in quantized value map for static quantizationdtyper   r   r   r   N)r_   q_namer  r   r8   r   r   r   rZ   r  rW   rX   npasarrayroundastypeint32reshapedimsrh   numpy_helper
from_arrayr~   float32zerosshaper   r   Initializersize)rc   	bias_namer   r   betaweight_scale_nameweight_initializerweight_scalebias_initializer	bias_dataquantized_bias_namer   r   inputscale_initializerinput_scale
bias_scalequantized_databias_np_datapacked_bias_initializerquantized_bias_scale_namebias_scale_datapacked_bias_scale_initializerquantized_bias_zp_namebias_zp_datapacked_bias_zp_initializerquantized_valuer   r   r   quantize_bias_static  sN   



z"ONNXQuantizer.quantize_bias_staticc                 C   s   || j v p|| jv p|| jv S )zq
        only check for value info and newly generated tensor names, initializers are checked separately
        )r4   rS   ra   )rc   r   r   r   r   contains_tensor  s
   
zONNXQuantizer.contains_tensorFc              	   C   s   | j ||dddd|dS )NFr  rT   indicesinitializer_use_weight_qTyper;   op_level_per_channelaxisfrom_subgraph_ONNXQuantizer__quantize_inputs)rc   rT   rD  rH  r   r   r   quantize_activation  s   z!ONNXQuantizer.quantize_activationr  c              	   C   s   | j ||d||||dS )NTrC  rI  )rc   rT   rD  r;   rF  rG  rH  r   r   r   quantize_weight  s   	zONNXQuantizer.quantize_weightTc              
   C   s8  g }g }	g }
g }|D ]
}|j | }|| jv r/| j| }||j |	|j |
|j q
t|| j }|durs| j	rS|rS| 
|j|rI| jn| j||\}}}n| ||r[| jn| j|\}}}|
| |	| || q
| |r| j|d | j| j }|du r| ||| j}|du r dS |r| | n|| |d }|jdkr|
|j ||j d  |	|j d  q
|
|jd  ||jd  |	|jd  q
| jdur| jj||g||||d	d
\}}}}|
|d  ||d  |	|d  q
td|| j|
|	||fS )a  
        Given a node, this function quantizes the inputs as follows:
            - If input is an initializer, quantize the initializer data, replace old initializer
              with new initializer
            - Else, add QuantizeLinear nodes to perform quantization
            parameter node: node being quantized in NodeProto format.
            parameter indices: input indices to quantize.
            return: (List of quantized input names,
                     List of zero point names used for input quantization,
                     List of scale names used for input quantization,
                     List of new QuantizeLinear nodes created)
        Nr  )NNNNr  r   r   r   r   T)rE  r;   rF  rG  rH  z2Invalid tensor name to quantize: {} @graph scope{})r7   r_   r   r  r  r  r   r8   r   r:   quantize_weight_per_channelr   rK   rJ   quantize_initializerrB  find_node_by_namerP   r2   r  r   r~   r{   r6   rQ   rJ  rW   rX   rR   )rc   rT   rD  rE  r;   rF  rG  rH  scale_nameszero_point_namesquantized_input_namesr   r  
node_inputr@  r   q_weight_namer  r  r  quantize_input_nodesparent_quantized_input_namesparent_zero_point_namesparent_scale_namesr   r   r   r   __quantize_inputs  s   











zONNXQuantizer.__quantize_inputsc                 C   s$  |j | jv r| j|j  }|j|j|jfS |j t }|j d }|j d }t|}	t|	 	 || j
| jo4|\}
}
}}}tj|tjjg |g}tj||g |g}| j ||g |s|tj|tjj| d|j}tj||}| j |g t|j |||tjd}|| j|j < |||fS )a  
        :param weight: TensorProto initializer
        :param qType: type to quantize to
        :param keep_float_weight: Whether to quantize the weight. In some cases, we only want to qunatize scale and zero point.
                                  If keep_float_weight is False, quantize the weight, or don't quantize the weight.
        :return: quantized weight name, zero point name, scale name
        r   r   r  N) r   r_   r  r  r  r   r   r   flattentolistrD   r;   rh   ri   r   rF   rG   r   r8   r   r~   r  r  mappingTENSOR_TYPE_TO_NP_TYPEr#  r$  r%  r&  r   r   r*  )rc   r   r   r;   keep_float_weightr@  rT  r  r  weight_datar   
zero_pointscaleq_weight_datascale_initializerzero_initializerq_weight_initializerr   r   r   rN  b  sF   	




z"ONNXQuantizer.quantize_initializerc                  C   s  || j v r| j | }|j|j|jfS t|| j }|d u r#td|t|}|j	| }	g }
g }g }g }g }t
|	D ];}|||}t|  || jpQ|tjjk| joU|\}}}}}|
| || || || || q:t|j	}d||< t|d |}t
dt|D ]}t|| |}t||f|}q|t }|d }|d }t||||tjd }|| j |< |j| g}t j!"|tjj#||}t j!"||||}| j $||g |stj|t j%j&| d|j}t j'(||}| j $|g |||fS )Nz{} is not an initializerr   r   r   r   r  ))r_   r  r  r  r   r8   r   rW   r   r)  r   taker   rZ  r[  rD   rF   rG   rH   r;   r   listr  r  r#  rz   concatenater   r   r   r*  r$  rh   ri   r   r   r~   r\  r]  r%  r&  ) rc   r   rK   channel_axisr;   r^  r@  r   weightschannel_count	rmin_list	rmax_listzero_point_list
scale_listquantized_per_channel_data_listr   per_channel_datarminrmaxr`  ra  quantized_per_channel_datareshape_dimsquantized_weightschannel_weightsrT  r  r  zero_scale_shaperc  rd  re  r   r   r   rM    s~   
	











z)ONNXQuantizer.quantize_weight_per_channelc                 C   s   || j v r@|| jvr@| j | }|d }| j|| j| j }|du r7|j|j|jg}t	j
d||g|}|S ||jd ks@J dS )a  
        Given a value (input/output) which is quantized, add a DequantizeLinear node to dequantize
        it back to float32
            parameter value_name: value to dequantize
            parameter new_nodes_list: List of new nodes created before processing current node
            return: None if there is already a DequantizeLinear node that dequantizes it
                    A DequantizeLinear node otherwise
        _DequantizeLinearNr   r   )r_   ra   r8   rO  rP   r2   r  r  r  rh   ri   r   r6   )rc   
value_namer@  dqlinear_namedqlinear_nodedqlinear_inputsdequantize_noder   r   r   _dequantize_value  s   	

zONNXQuantizer._dequantize_valuec                 C   s6   | j  jD ]}| |j}|dur| j| qdS )z
        Dequantize output if it is quantized
            parameter new_nodes_list: List of new nodes created before processing current node
            return: List of new nodes created
        N)r8   r2   r6   r  r   rP   r   )rc   r6   r~  r   r   r   r     s   z!ONNXQuantizer._dequantize_outputsc                 C   s   | j d u rd S | j D ]D}|jdvrq| jrq| |sqt| j |jd  dkr-q|jd | j 	 vsA|j
d | j 	 vrBq| j |j
d  | j |jd < qi }| j 	 D ]}| j | \}}t| j| jd\}}t||||| j||< qX|S )N)ClipRelur   r   )	symmetric)rL   r8   r   r{   rE   r   rz   input_name_to_nodesr7   r   r6   r   rJ   r   )rc   rT   rZ   r   rr  rs  qminqmaxr   r   r   rY     s(   


(z+ONNXQuantizer.calculate_quantization_paramsr   r  )r  )F)FFr  F)TFFr  F)FF)TF)"__name__
__module____qualname__re   rq   r   rU   r   r   r   rl   r   r   r   r   r   r   r   r   r  r  r  r  r  rA  rB  rK  rL  rJ  rN  rM  r  r   rY   r   r   r   r   r   $   sX    
g"%S
]
'>
B




l:
Sr   )r   numpyr  rh   onnx.numpy_helperr   rF   
onnx_modelr   quant_utilsr   r   r   r   r	   r
   r   r   r   r   r   r   r   r   r   r   r   registryr   r   r   r   r   r   <module>   s   L