
    \j                     "   d dl Z d dlZd dlZd dlZ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 d dlZd dlZd dlmZmZmZmZ d dlZddlmZmZmZ dej4                  d	ej4                  d
ej4                  fdZ	 	 d6dej4                  d	ej4                  dedz  ded
ej4                  f
dZ G d d      Z G d d      Z  G d de      Z! G d de jD                        Z# G d dejH                        Z%ddddddde&d
dfd Z'd7d!Z( G d" d#      Z) G d$ d%e)      Z* G d& d'e)      Z+ G d( d)e+      Z, G d* d+e+      Z- G d, d-e+      Z. G d. d/e jD                        Z/ G d0 d1e/      Z0dd2e!jb                  ddi fd3e2ez  d4e	e2   dz  fd5Z3y)8    N)Sequence)Enum)Path)
ModelProtoTensorProtohelpernumpy_helper   )
apply_plotload_model_with_shape_infersmooth_distributionpkqkreturnc                    t        j                  | j                  | j                        }| dd t        j                  | dd |dd z        z  |dd | dk(  |dk\  z  }d||<   | dkD  |dkD  z  }t         j
                  || <   |S )z
    See https://docs.scipy.org/doc/scipy/reference/generated/scipy.special.rel_entr.html#scipy.special.rel_entr.
    Python implementation.
    dtypeNr   )npemptyshaper   loginf)r   r   resc2c1s        U/root/.hermes/venv/lib/python3.12/site-packages/onnxruntime/quantization/calibrate.pyrel_entrr      s    
 ((288288
,CURVVBqEBqEM**CF
'bAg	BCG
q&R!V	BvvCHJ    baseaxisc                 R   ||dkD  sJ d       |J d       t        j                  |       j                  t         j                        } d| z  t        j                  | |d      z  } t        j                  |      j                  t         j                        }t        j
                  | |      \  } }d|z  t        j                  ||d      z  }t        | |      }t        j                  ||      }||t        j                  |      z  }|j                  | j                        S )z
    Simplifeied version of entropy.
    Source: https://docs.scipy.org/doc/scipy/reference/generated/scipy.stats.entropy.html.
    This avoids taking a dependency on scipy just for this function.
    r   z0base={base} must be a positive number or `None`.z
qk is None      ?T)r    keepdimsr    )	r   asarrayastypefloat32sumbroadcast_arraysr   r   r   )r   r   r   r    vecss         r   entropyr,   *   s     <4!8W%WW#>'<'>	B		rzz	*B	rBFF2D48	8B	B		rzz	*B  R(FB	rBFF2D48	8B
2r
C
sA	RVVD\88BHHr   c                   z    e Zd Z eg d      Z eg d      Zd Zed        Zed        Z	d Z
ededd fd	       Zy
)
TensorData)avgstdlowesthighesthist
hist_edgesbins)r/   r0   r1   r2   r4   c                    t        |j                               | _        |j                         D ]  \  }}|t        j
                  vr t        d|dt        j
                   d      |t        j                  v rmt        |d      st        dt        |       d|      |j                  t        j                  t        j                  fvrt        d|j                   d|      t        | ||        y )NzUnexpected value z not in .r   Unexpected type z for k=zUnexpected dtype )listkeys_attrsitemsr.   _allowed
ValueError_floatshasattrtyper   r   float16r'   setattr)selfkwargskvs       r   __init__zTensorData.__init__J   s    6;;=)LLN 	 DAq
+++ #4QE*BUBUAVVW!XYYJ&&&q'*$'7Qyu%MNN772::rzz"::$'8	%NOOD!Q	 r   c                     t        | d      rt        | d      st        dt        |        d      | j                  | j                  fS )Nr1   r2   z0Attributes 'lowest' and/or 'highest' missing in r7   )r@   AttributeErrordirr1   r2   rD   s    r   range_valuezTensorData.range_valueV   sF    tX&gdI.F #STWX\T]S^^_!`aaT\\**r   c                     t        | d      rt        | d      st        dt        |        d      | j                  | j                  fS )Nr/   r0   z)Attributes 'avg' and/or 'std' missing in r7   )r@   rJ   rK   r/   r0   rL   s    r   avg_stdzTensorData.avg_std\   sC    tU#74+? #LSQUYKWX!YZZ$((##r   c                     | j                   D ci c]  }|t        | |       }}| j                  j                  |d<   |S c c}w )NCLS)r;   getattr	__class____name__)rD   rF   datas      r   to_dictzTensorData.to_dictb   sB    -1[[974##99nn--U :s   A dr   c                    i }|j                         D ]  \  }}|dk(  r|}t        |t              rE|j                  d      dk(  r1t	        j
                  |d   t	        j                  |d               }nI|| j                  v r;t        |t        t        f      r%t	        j
                  |t        j                        }|||<     | di |S )z;Reconstruct a TensorData from a dict produced by to_dict().rQ   numpy.arrayrU   r   r    )r<   
isinstancedictgetr   arrayr   r?   intfloatr'   )clsrW   rE   rF   rG   values         r   	from_dictzTensorData.from_dicth   s     GGI 	DAqEzE%&599U+;}+LvbhhuW~6NOckk!je&Ebjj9F1I	 }V}r   N)rT   
__module____qualname__	frozensetr=   r?   rH   propertyrM   rO   rV   classmethodr\   rc   rZ   r   r   r.   r.   F   sl    Z[HIJG
  + +
 $ $
 $ <  r   r.   c                   r    e Zd Zdeeeez  f   fdZd Zd Z	d Z
d Zd Zd Zd	 Zd
 Zededd fd       Zy)TensorsDatarU   c           
      b   || _         i | _        |j                         D ]  \  }}t        |t              st        dt        |       d      t        |t              r|t        j                  k(  r/t        |      dk(  r!t        |d   |d         | j                  |<   t        |      dk(  r)t        |d   |d   |d   |d   	      | j                  |<   t        d
|ddt        |       d| d      t        |t              st        dt        |       d      || j                  |<    y )NzKeys must be strings not r7      r   r
   r1   r2         )r1   r2   r3   r5   zUnexpected tuple for rz	, it has z elements: zValues must be TensorData not )calibration_methodrU   r<   r[   str	TypeErrorrA   tupleCalibrationMethodMinMaxlenr.   )rD   rq   rU   rF   rG   s        r   rH   zTensorsData.__init__y   s$   "4	JJL 	DAqa%";DG9A FGG!U#%):)A)AAc!fPQk#-QqT1Q4#HDIIaLq6Q;#-QqT1Q4aPQdYZ[\Y]#^DIIaL"7!uIc!fX[YZX[[\ ]^^a,"@a	 KLLDIIaL	r   c              #   8   K   | j                   E d {    y 7 wNrU   rL   s    r   __iter__zTensorsData.__iter__   s     99s   c                     || j                   v S ry   rz   rD   keys     r   __contains__zTensorsData.__contains__   s    diir   c                      | j                   |   S ry   rz   r}   s     r   __getitem__zTensorsData.__getitem__   s    yy~r   c                 \    || j                   vrt        d|d      || j                   |<   y )Nz)Only an existing tensor can be modified, z is not.)rU   RuntimeError)rD   r~   rb   s      r   __setitem__zTensorsData.__setitem__   s1    dii!J3'QYZ[[		#r   c                 6    | j                   j                         S ry   )rU   r:   rL   s    r   r:   zTensorsData.keys   s    yy~~r   c                 6    | j                   j                         S ry   )rU   valuesrL   s    r   r   zTensorsData.values   s    yy!!r   c                 6    | j                   j                         S ry   )rU   r<   rL   s    r   r<   zTensorsData.items   s    yy  r   c                 b    | j                   j                  | j                  | j                  d}|S )N)rQ   rU   rq   )rS   rT   rU   rq   )rD   rU   s     r   rV   zTensorsData.to_dict   s/     >>**II"&"9"9

 r   rW   r   c                 *   |d   }t        |t              r5|j                  d      dk(  r!|d   j                  d      d   }t        |   }n|}|d   j                         D ci c]  \  }}|t        j                  |       }}} | ||      S c c}}w )z<Reconstruct a TensorsData from a dict produced by to_dict().rq   rQ   ru   rb   r7   rU   )r[   r\   r]   splitru   r<   r.   rc   )ra   rW   
method_valnamemethodrF   rG   reconstructeds           r   rc   zTensorsData.from_dict   s     +,
j$'JNN5,AEX,Xg&,,S1"5D&t,FF@A&	@QR1J0033RR6=)) Ss   # BN)rT   rd   re   r\   rr   r.   rt   rH   r{   r   r   r   r:   r   r<   rV   rh   rc   rZ   r   r   rj   rj   x   sg    c:;M6M1N $ 
 "! 	*$ 	*= 	* 	*r   rj   c                       e Zd ZdZdZdZdZy)ru   r   r
   rl   ro   N)rT   rd   re   rv   Entropy
PercentileDistributionrZ   r   r   ru   ru      s    FGJLr   ru   c                   h    e Zd Zed        Zej                  defd       Zd Z	d Z
d Zdedefd	Zy
)CalibrationDataReaderc                 X    t        |d      xr t        |j                        xs t        S )Nget_next)r@   callabler   NotImplemented)ra   subclasss     r   __subclasshook__z&CalibrationDataReader.__subclasshook__   s%    *-M(8;L;L2M`R``r   r   c                     t         )z9generate the input data dict for ONNXinferenceSession runNotImplementedErrorrL   s    r   r   zCalibrationDataReader.get_next   s
     "!r   c                     | S ry   rZ   rL   s    r   r{   zCalibrationDataReader.__iter__   s    r   c                 6    | j                         }|t        |S ry   )r   StopIteration)rD   results     r   __next__zCalibrationDataReader.__next__   s    >r   c                     t         ry   r   rL   s    r   __len__zCalibrationDataReader.__len__       !!r   start_index	end_indexc                     t         ry   r   )rD   r   r   s      r   	set_rangezCalibrationDataReader.set_range   r   r   N)rT   rd   re   rh   r   abcabstractmethodr\   r   r{   r   r   r_   r   rZ   r   r   r   r      sY    a a 	"$ " """S "S "r   r   )	metaclassc                       e Zd ZdZd Zy)CalibrationCacheEncoderzShared JSON encoder for calibration caches.

    Handles numpy ndarrays and numpy scalar types (integer/floating) so
    calibration JSON output is consistent across ``save_tensors_data`` and
    ``quant_utils.write_calibration_table``.
    c                    t        |t        t        f      r|j                         S t        |t        j
                        r'|j                         t        |j                        ddS t        |t              r"|j                  j                  t        |      dS t        |t        j                        rt        |      S t        |t        j                        rt        |      S t         j"                  j%                  | |      S )NrY   )rU   r   rQ   )rQ   rb   )r[   r.   rj   rV   r   ndarraytolistrr   r   ru   rS   rT   integerr_   floatingr`   jsonJSONEncoderdefault)rD   objs     r   r   zCalibrationCacheEncoder.default   s    cJ45;;= c2::&JJL3syy>-XXc,-==11CHEEc2::&s8Oc2;;':''c22r   N)rT   rd   re   __doc__r   rZ   r   r   r   r      s    3r   r   F)smooth_quanttensors_datapath
str | Pathr   c                F   t        |      }|j                  j                  dd       t        j                  |j                  dd      \  }}	 t        j                  |d      5 }| j                         }||d<   t        j                  ||t               |j                          d	d	d	       t        j                  ||       y	# 1 sw Y    xY w# t        $ rE t        j                  t               5  t        j"                  |       d	d	d	        # 1 sw Y    xY ww xY w)
zSerialize calibration tensor ranges to a JSON file at *path*.

    :param smooth_quant: whether the producing run used SmoothQuant.  Stored in
        the cache so a later load can detect a mismatch and recompute.
    T)parentsexist_okz.calibcache_z.tmp)rK   prefixsuffixwr   )ra   N)r   parentmkdirtempfilemkstemposfdopenrV   r   dumpr   flushreplaceBaseException
contextlibsuppressFileNotFoundErrorunlink)r   r   r   fdtmp_namefpayloads          r   save_tensors_datar      s     :DKKdT2##NSYZLB
YYr3 	1"**,G&2GN#IIgq&=>GGI		
 	

8T"	 	    !23 	 IIh	 	 s=   C %AC'C CC "D 4D
	D D	D c                 0   t        |       } | j                         st        d|        | j                         st	        d|        | j                  d      5 }t        j                  |      }ddd       t        j                        S # 1 sw Y   xY w)zOLoad calibration tensor ranges from a JSON file written by save_tensors_data().zCalibration cache not found: z&Calibration cache path is not a file: rp   N)
r   existsr   is_filer>   openr   loadrj   rc   )r   r   rW   s      r   load_tensors_datar     s    :D;;="?v FGG<<>A$HII	3 1IIaL  ## s   BBc                   |    e Zd Z	 	 	 	 	 ddeez  dee   dz  fdZdgfdZd Zde	fd	Z
d
 Zd ZdefdZdefdZy)CalibraterBaseN
model_pathop_types_to_calibratec                 "   t        |t              rt        t        |            | _        n,t        |t              rt        |      | _        nt        d      || _        || _        || _        || _	        || _
        d| _        d| _        dg| _        y)a  
        :param model_path: ONNX model to calibrate. It should be a model file path
        :param op_types_to_calibrate: operator types to calibrate. By default, calibrate all the float32/float16 tensors.
        :param augmented_model_path: save augmented model to this path.
        :param symmetric: make range of tensor symmetric (central point is 0).
        :param use_external_data_format: use external data format to store model which size is >= 2Gb.
        :param per_channel: whether to compute ranges per each channel.
        z model_path should be model path.NCPUExecutionProvider)r[   rr   r   r   modelr>   r   augmented_model_path	symmetricuse_external_data_formatper_channelaugment_modelinfer_sessionexecution_providers)rD   r   r   r   r   r   r   s          r   rH   zCalibraterBase.__init__  s    " j#&4T*5EFDJ
D)4Z@DJ?@@%:"$8!"(@%&!!$:#; r   r   c                 2    || _         | j                          y)zz
        reset the execution providers to execute the collect_data. It triggers to re-creating inference session.
        N)r   create_inference_session)rD   r   s     r   set_execution_providersz&CalibraterBase.set_execution_providers4  s     $7 %%'r   c                     t        j                         }t         j                  j                  |_        t        j
                  | j                  || j                        | _        y)z9
        create an OnnxRuntime InferenceSession.
        )sess_options	providersN)	onnxruntimeSessionOptionsGraphOptimizationLevelORT_DISABLE_ALLgraph_optimization_levelInferenceSessionr   r   r   )rD   r   s     r   r   z'CalibraterBase.create_inference_session;  sN     #1130;0R0R0b0b-(99%%%..
r   r   c                    |j                   j                  D ci c]  }|j                  | }}|j                  |j                   j                  D ci c]  }|j                  | c}       |j                  |j                   j
                  D ci c]  }|j                  | c}       |j                   j                  D ch c]  }|j                   }}t               }t        j                  t        j                  h}	|j                   j                  D ]  }
| j                  r|
j                  | j                  v s(t        j                  |
j
                  |
j                        D ]a  }||v s||   }|j                   j#                  d      s)|j                   j$                  j&                  |	v sL||vsQ|j)                  |       c  ||fS c c}w c c}w c c}w c c}w )z
        select input/output tensors of candidate nodes to calibrate.
        returns:
            tensors (set): set of tensor name.
            value_infos (dict): tensor name to value info.
        tensor_type)graph
value_infor   updateoutputinputinitializersetr   FLOATFLOAT16noder   op_type	itertoolschainrA   HasFieldr   	elem_typeadd)rD   r   vivalue_infosotitinitr   tensors_to_calibratetensor_type_to_calibrater  tensor_names               r   select_tensors_to_calibratez*CalibraterBase.select_tensors_to_calibrateG  s    .3[[-C-CDrrww{DD%++2D2DEBBGGRKEF%++2C2CDBBGGRKDE-2[[-D-DETtyyEE"u$/$5$5{7J7J#K KK$$ 
	BD--A[A[1[#,??4::t{{#K BK"k1(5GG,,];!#!4!4!>!>BZ!Z!,K!?044[AB
	B $[00) EEDEs   GGGG#c                     | j                   S )zP
        return: augmented onnx model. Call after calling augment_graph
        )r   rL   s    r   get_augment_modelz CalibraterBase.get_augment_modeld  s     zzr   c                     t         )z
        abstract method: augment the input model to prepare for collecting data. It will:
            1. augment the model to be able to collect desired statistics data
            2. save augmented model to augmented_model_paths
        r   rL   s    r   augment_graphzCalibraterBase.augment_graphj  s
     "!r   data_readerc                     t         )z
        abstract method: collect the tensors that will be used for range computation. It can be called multiple times.
        r   )rD   r  s     r   collect_datazCalibraterBase.collect_datar  
     "!r   r   c                     t         )ze
        abstract method: compute data based on the calibration method stored in TensorsData
        r   rL   s    r   compute_datazCalibraterBase.compute_datax  r  r   )Naugmented_model.onnxFFF)rT   rd   re   rr   r   r   rH   r   r   r   r  r  r  r   r  rj   r  rZ   r   r   r   r     sz     7;3!& <$J <  (}t3 <D <R:R (

1 1:""(= ""k "r   r   c                   v     e Zd Z	 	 	 	 	 	 	 	 ddeez  dee   dz  f fdZd Zd Zde	fdZ
d	 Zd
efdZ xZS )MinMaxCalibraterNr   r   c
                    t         |   ||||||	       g | _        d| _        t	        | j
                  j                  j                        | _        | j
                  j                  j                  D 
ch c]  }
|
j                   c}
| _
        || _        |r|dk  s|dkD  rt        d      || _        || _        yc c}
w )aw  
        :param model_path: ONNX model to calibrate. It is a model path
        :param op_types_to_calibrate: operator types to calibrate. By default, calibrate all the float32/float16 tensors.
        :param augmented_model_path: save augmented model to this path.
        :param symmetric: make range of tensor symmetric (central point is 0).
        :param use_external_data_format: use external data format to store model which size is >= 2Gb
        :param moving_average: compute the moving average of the minimum and maximum values instead of the global minimum and maximum.
        :param averaging_constant: constant smoothing factor to use when computing the moving average.
        :param max_intermediate_outputs: maximum number of intermediate outputs before an intermediate range is computed.
        :param per_channel: whether to compute ranges per each channel.
        )r   r   r   r   r   Nr   r
   z;Invalid averaging constant, which should not be < 0 or > 1.)superrH   intermediate_outputscalibrate_tensors_rangerw   r   r   r   num_model_outputsr   model_original_outputsmoving_averager>   averaging_constantmax_intermediate_outputs)rD   r   r   r   r   r   r&  r'  r(  r   r   rS   s              r   rH   zMinMaxCalibrater.__init__  s    . 	"7!5%=# 	 	
 %'!'+$!$TZZ%5%5%<%<!=AEAQAQAXAX&Yvv{{&Y#,1A59Ka9OZ[["4(@% 'Zs   5B=c                      j                   j                        \  }}t        t        j                               t        j                  t        j                  dgt        j                              } j                  j                  j                  j                  |       d  fd fd}|D ]  } ||d        ||d        t        j                   j                   j                   j                          y	)
z
        Adds ReduceMin and ReduceMax nodes to all quantization_candidates op type nodes in
        model and ensures their outputs are stored as part of the graph output
        :return: augmented ONNX model
        r   r   c                     |j                   D ]:  }t        j                  j                  | |j                        s.|j
                  c S  t        d|  d      )Nz&Model does not contain a version for 'z'.)opset_importonnxdefshasdomainversionr   )r  r   r+  s      r   get_op_versionz6MinMaxCalibrater.augment_graph.<locals>.get_op_version  sS     % 2 2 099==,*=*=>'///0 !GyPRSTTr   c                 F    t         fdt        j                  j                  j                        D        t        j                  j                  j                              }|D ]7  }j                  j                  j                  j                  ||       |dz  }9 y )Nc              3   F   K   | ]  \  }}|j                   v s|  y wry   )r   ).0ixr  s      r   	<genexpr>zGMinMaxCalibrater.augment_graph.<locals>.insert_nodes.<locals>.<genexpr>  s#     Ztq!;RSRYRYCYZs   !!r
   )next	enumerater   r   r  rw   insert)r  	new_nodesindexr  rD   s   `   r   insert_nodesz4MinMaxCalibrater.augment_graph.<locals>.insert_nodes  s    Zy)9)9)>)>?Z\_`d`j`j`p`p`u`u\vE " 

  %%,,UD9
r   c                    d}| dz   |z   }|dz   }t         j                  j                  || g|g||      }t         j                  j                  d|g|g|      }j                  j                  j
                  D ci c]  }|j                  | }}|j                  j                  j                  j                  D 	ci c]  }	|	j                  |	 c}	       |j                  j                  j                  j                  D 
ci c]  }
|
j                  |
 c}
       | |v r$||    j                  j                  j                  }nt        d| d      j                  r+t        ||    j                  j                  j                   j"                        }d	gt%        d
|      } |j                        dk  r0|j&                  j)                  t        j*                  d|             nt-        t/        j0                               }t3        j4                  t7        j8                  |t6        j:                        |      }|j                  j)                  |       j                  j                  j<                  j)                  |        | ||g       j                  j                  j                  j)                  t        j>                  ||d g             y c c}w c c}	w c c}
w )Nr
   __Reshape)r#   r   Reshape)inputsoutputsr   z'Unable to guess tensor type for tensor zE, running shape inference before quantization may resolve this issue.r   rl      axesr   ) r,  r   	make_noder   r   r   r   r   r   r   rA   r   r  r>   r   rw   r   dimrange	attributeappendmake_attributerr   uuiduuid4r	   
from_arrayr   r^   int64r   make_tensor_value_info)r  reduce_op_namer#   reduce_outputintermediate_outputreduce_nodereshape_noder
  r  or5  	onnx_typetensor_rankreduced_axesreduce_axes_namereduce_axesr1  r=  reshape_shape_namerD   s                   r   add_reduce_min_maxz:MinMaxCalibrater.augment_graph.<locals>.add_reduce_min_max  s    H (#->M"/*"<++//0C/Dx^k 0 K  ;;00+-?@&(	 1 L 261A1A1L1LM2277B;MKM4::3C3C3J3JKa	KL4::3C3C3I3IJa	JKk)'499EEOO	 =k_ MZ Z  !+k":"?"?"K"K"Q"Q"U"UV !:E![$9:!.$**=B))001F1Fv|1\]'*4::<'8$"."9"9"((<WYW_W_:`br"sK%%,,-=>JJ$$0077D{L&ABJJ##**6+H+HXadhci+jk3 NKJs   ?K%K*
K/	ReduceMin	ReduceMaxsave_as_external_dataN)r  r   rr   rL  rM  r	   rN  r   r^   rO  r   r   rJ  r,  saver   r   )	rD   tensorsr?  reshape_shaper]  tensorr1  r=  r\  s	   `     @@@r   r  zMinMaxCalibrater.augment_graph  s     55djjA
 .$//"RXX0NPbc

$$++M:	U	,	l\  	4Fv{3v{3	4 			JJ%%"&"?"?	
r   c                     g | _         y ry   r"  rL   s    r   clear_collected_dataz%MinMaxCalibrater.clear_collected_data  
    $&!r   r  c           	         	 |j                         }|sn| j                  j                  t        | j                  j                         | j                  j                  d |      d      D cg c]!  \  }}|j                  | j                  vr|nd # c}}       | j                  2t        | j                        | j                  k(  r| j                          t        | j                        dk(  r| j                  t        d      | j                         }t        |t               st#        dt%        |       d      | j                          y c c}}w )NFstrictr   No data is collected.z+compute_data must return a TensorsData not r7   )r   r"  rJ  zipr   get_outputsrunr   r%  r(  rw   rh  r#  r>   r  r[   rj   rs   rA   )rD   r  rB  sess_orb   ts         r   r  zMinMaxCalibrater.collect_data  s8    ))+F%%,, *-**668$:L:L:P:PQUW]:^gl*% $[[0K0KKEQUU --9112d6S6SS))+! $ t(()Q.43O3O3W455![)I$q'RSTUU!!#'s   -&E
c                 >   |s|S |j                         D ]  \  }}t        |t              r|j                  d   }|j                  d   }n|\  }}t        ||   t              r%||   j                  d   }||   j                  d   }n||   \  }}| j                  r+|| j
                  ||z
  z  z   }	|| j
                  ||z
  z  z   }
nt        ||      }	t        ||      }
t        |t              st        ||   t              rt        |	|
      ||<   |	|
f||<    |S )Nr   r
   rm   )r<   r[   r.   rM   r&  r'  minmax)rD   	old_range	new_ranger~   rb   old_minold_maxnew_minnew_max	min_value	max_values              r   merge_rangezMinMaxCalibrater.merge_range  s1   #//+ 	8JC%,++A.++A.#( )C.*5#C.44Q7#C.44Q7#,S> ""#d&=&=7AR&SS	#d&=&=7AR&SS	1	1	 %,
9S>:0V!+9i!P	#"+Y!7	#3	86 r   r   c           
         t        | j                        dk(  r| j                  S t        t        | j                  d               D cg c])  }| j                  j                         |   j                  + }}| j                  D cg c]  }t        t        ||d             }}i }|D ];  }|j                         D ]&  \  }}|j                  |g       j                  |       ( = || j                  d }	t        dt        |	      d      D cg c]  }|	|   j                  d      d    }
}|D ci c]  }|| j                  vs|||    }}g }t        dt        |	      d      D ]  }| j                  r>t!        j"                  ||	|      d      }t!        j"                  ||	|dz         d      }n=t!        j$                  ||	|      d      }t!        j&                  ||	|dz         d      }| j(                  rTt!        j&                  t!        j*                  |      t!        j*                  |      gd      }|j                  | |f       |j                  ||f        t-        t.        j0                  t        t        |
|d                  }| j                  r-| j3                  | j                  |      | _        | j                  S || _        | j                  S c c}w c c}w c c}w c c}w )	z
        Compute the min-max range of tensor
        :return: dictionary mapping: {added node names: (ReduceMin, ReduceMax) pairs }
        r   Frk  Nrl   r?  r$   r
   )rw   r"  r#  rH  r   ro  r   r\   rn  r<   
setdefaultrJ  r$  
rpartitionr%  r&  r   nanmeannanminnanmaxr   absrj   ru   rv   r~  )rD   r5  output_namesrS  output_dicts_listmerged_output_dictrW   rF   rG   added_output_namescalibrate_tensor_namesmerged_added_output_dictpairsmin_value_arraymax_value_arraymax_absolute_valuenew_calibrate_tensors_ranges                    r   r  zMinMaxCalibrater.compute_data9  s    t(()Q.///JOPSTXTmTmnoTpPqJrsQ**668;@@ss (,'@'@
# \#6uEF
 

  " 	?A	 ?1"--a4;;A>?	? *$*@*@*BC>CAsK]G^`a>b"
9:q!,,S1!4"
 "

 /A$
)*ATMhMhDhA!!$$$
  $
 q#0115 	AA"""$**-EFXYZF[-\cd"e"$**-EFXYZ]^Y^F_-`gh"i"$)),DEWXYEZ,[bc"d"$)),DEWXY\]X]E^,_fg"h~~%'YY0GP_I`/ahi%j"113EFGo?@	A '2$$d3/EuUZ+[&\'
# ''+/+;+;D<X<XZu+vD( +++ ,GD(+++U t
"
$
s   .K#K(K-3K2K2)Nr  FFF{Gz?NF)rT   rd   re   rr   r   r   rH   r  rh  r   r  r~  rj   r  __classcell__rS   s   @r   r  r    sp     7;3!&!%'A$J'A  (}t3'ARO
b'$(= $6B3,k 3,r   r  c                   r     e Zd Z	 	 	 	 	 	 	 	 	 ddeez  dee   dz  f fdZd Zd Zde	fdZ
d	efd
Z xZS )HistogramCalibraterNr   r   c                    t         |   |||||       g | _        d| _        t	        | j
                  j                  j                        | _        | j
                  j                  j                  D ch c]  }|j                   c}| _
        d| _        || _        || _        || _        |	| _        d| _        |
| _        yc c}w )a=  
        :param model_path: ONNX model to calibrate. It is a model path.
        :param op_types_to_calibrate: operator types to calibrate. By default, calibrate all the float32/float16 tensors.
        :param augmented_model_path: save augmented model to this path.
        :param use_external_data_format: use external data format to store model which size is >= 2Gb
        :param method: A string. One of ['entropy', 'percentile'].
        :param symmetric: make range of tensor symmetric (central point is 0).
        :param num_bins: number of bins to create a new histogram for collecting tensor values.
        :param num_quantized_bins: number of quantized bins. Default 128.
        :param percentile: A float number between [0, 100]. Default 99.99.
        :param scenario: see :class:`DistributionCalibrater`
        )r   r   r   r   N)r!  rH   r"  r#  rw   r   r   r   r$  r   r%  	collectorr   num_binsnum_quantized_bins
percentiler  scenario)rD   r   r   r   r   r   r   r  r  r  r  r   rS   s               r   rH   zHistogramCalibrater.__init__p  s    2 	"7!5%= 	 	
 %'!'+$!$TZZ%5%5%<%<!=AEAQAQAXAX&Yvv{{&Y# "4$$(!  'Zs   4Cc                 Z   | j                  | j                        \  | _        }| j                  D ]C  }|| j                  vs| j                  j                  j
                  j                  ||          E t        j                  | j                  | j                  | j                         y)z
        make all quantization_candidates op type nodes as part of the graph output.
        :return: augmented ONNX model
        r`  N)r  r   r  r%  r   r   rJ  r,  rb  r   r   )rD   r  re  s      r   r  z!HistogramCalibrater.augment_graph  s    
 261Q1QRVR\R\1].!;// 	DFT888

  ''..{6/BC	D 			JJ%%"&"?"?	
r   c                     g | _         y ry   rg  rL   s    r   rh  z(HistogramCalibrater.clear_collected_data  ri  r   r  c           
         | j                   j                         D ch c]  }|j                   }}| j                   j                         D cg c]  }|j                   }}	 |j	                         }|sn| j                   j                  d|      }g }t        |      D ]B  \  }}	||   |v r%|j                  t        j                  |	             2|j                  |	       D | j                  j                  |       t        | j                        dk(  rt        d      | j                  D 
cg c]  }
t        t        ||
d             }}
i }|D ];  }|j                         D ]&  \  }}|j                  |g       j                  |       ( = |D ci c]  }|| j                   v s|||    }}| j"                  sRt%        | j&                  | j(                  | j*                  | j,                  | j.                  | j0                        | _        | j"                  j3                  |       | j5                          yc c}w c c}w c c}
w c c}w )zy
        Entropy Calibrator collects operators' tensors as well as generates tensor histogram for each operator.
        Nr   rm  Frk  )r   r   r  r  r  r  )r   
get_inputsr   ro  r   rp  r9  rJ  copyr"  rw   r>   r\   rn  r<   r  r  r  HistogramCollectorr   r   r  r  r  r  collectrh  )rD   r  node_arginput_names_setr  rB  rC  fixed_outputsoutput_indexr   rS  r  merged_dictrW   rF   rG   r5  clean_merged_dicts                     r   r  z HistogramCalibrater.collect_data  s.    :>9K9K9V9V9XYX8==YY6:6H6H6T6T6VW(WW ))+F((,,T6:G M(1'(: 1$f-@!((6):;!((0	1 %%,,]; " t(()Q.455 (,'@'@
# \#6uEF
 

 " 	8A	 81&&q"-44Q78	8 9Df1qDLeLeGeQA.ff~~/{{..#'#:#:??DN 	01!!#] ZW,
 gs   I I2I
I,Ir   c                 n   | j                   st        d      t        | t              rt        j
                  }nZt        | t              rt        j                  }n9t        | t              rt        j                  }nt        dt        |        d      t        || j                   j                               S )z
        Compute the min-max range of tensor
        :return: dictionary mapping: {tensor name: (min value, max value)}
        z9No collector created and can't generate calibration data.zUnknown calibrater z". This method must be overwritten.)r  r>   r[   EntropyCalibraterru   r   PercentileCalibraterr   DistributionCalibraterr   rs   rA   rj   compute_collection_result)rD   cals     r   r  z HistogramCalibrater.compute_data  s    
 ~~XYYd-.#++C23#..C45#00C1$t*=_`aa3 H H JKKr   )	Nr  Fr  F      -X@same)rT   rd   re   rr   r   r   rH   r  rh  r   r  rj   r  r  r  s   @r   r  r  o  sk     7;3!&*!$J*!  (}t3*!X
 '2$(= 2$hLk Lr   r  c                   J     e Zd Z	 	 	 	 	 	 	 ddeez  dee   dz  f fdZ xZS )r  Nr   r   c	           
      4    t         	|   ||||||||       y)a  
        :param model_path: ONNX model to calibrate. It is a model path
        :param op_types_to_calibrate: operator types to calibrate. By default, calibrate all the float32/float16 tensors.
        :param augmented_model_path: save augmented model to this path.
        :param use_external_data_format: use external data format to store model which size is >= 2Gb
        :param method: A string. One of ['entropy', 'percentile', 'distribution'].
        :param symmetric: make range of tensor symmetric (central point is 0).
        :param num_bins: number of bins to create a new histogram for collecting tensor values.
        :param num_quantized_bins: number of quantized bins. Default 128.
        )r   r   r  r  Nr!  rH   )
rD   r   r   r   r   r   r   r  r  rS   s
            r   rH   zEntropyCalibrater.__init__  s/    * 	! $1 	 		
r   )Nr  Fr,   Fr  r  rT   rd   re   rr   r   r   rH   r  r  s   @r   r  r    sC     7;3!&
$J
  (}t3
 
r   r  c                   J     e Zd Z	 	 	 	 	 	 	 ddeez  dee   dz  f fdZ xZS )r  Nr   r   c	           
      4    t         	|   ||||||||       y)a  
        :param model_path: ONNX model to calibrate. It is a model path
        :param op_types_to_calibrate: operator types to calibrate. By default, calibrate all the float32/float16 tensors.
        :param augmented_model_path: save augmented model to this path.
        :param use_external_data_format: use external data format to store model which size is >= 2Gb
        :param method: A string. One of ['entropy', 'percentile', 'distribution'].
        :param symmetric: make range of tensor symmetric (central point is 0).
        :param num_quantized_bins: number of quantized bins. Default 128.
        :param percentile: A float number between [0, 100]. Default 99.99.
        )r   r   r  r  Nr  )
rD   r   r   r   r   r   r   r  r  rS   s
            r   rH   zPercentileCalibrater.__init__  s/    * 	! $! 	 		
r   )Nr  Fr  Fr  r  r  r  s   @r   r  r    sC     7;3!&
$J
  (}t3
 
r   r  c                   H     e Zd Z	 	 	 	 	 	 ddeez  dee   dz  f fdZ xZS )r  Nr   r   c           	      2    t         |   |||||||       y)a  
        :param model_path: ONNX model to calibrate. It is a model path
        :param op_types_to_calibrate: operator types to calibrate. By default, calibrate all the float32/float16 tensors.
        :param augmented_model_path: save augmented model to this path.
        :param use_external_data_format: use external data format to store model which size is >= 2Gb
        :param method: A string. One of ['entropy', 'percentile', 'distribution'].
        :param symmetric: make range of tensor symmetric (central point is 0).
        :param num_bins: number of bins to create a new histogram for collecting tensor values.
        :param scenario: for float 8 only, if `scenario="same"`,
            the algorithm weights and float 8 follow the same distribution,
            if `scenario="p3"`, it assumes the weights follow
            a gaussian law and float 8 ~ X^3 where X is a gaussian law
        )r   r  r  Nr  )	rD   r   r   r   r   r   r  r  rS   s	           r   rH   zDistributionCalibrater.__init__;  s,    . 	! $ 	 	
r   )Nr  Fdistributionr  r  r  r  s   @r   r  r  :  s@     7;3!&
$J
  (}t3
 
r   r  c                   X    e Zd ZdZej
                  d        Zej
                  d        Zy)CalibrationDataCollectorzL
    Base class for collecting data for calibration-based quantization.
    c                     t         )z
        Generate informative data based on given data.
            name_to_arr : dict
                tensor name to NDArray data
        r   rD   name_to_arrs     r   r  z CalibrationDataCollector.collectb  s
     "!r   c                     t         )z?
        Get the optimal result among collection data.
        r   rL   s    r   r  z2CalibrationDataCollector.compute_collection_resultk  s
    
 "!r   N)rT   rd   re   r   r   r   r  r  rZ   r   r   r  r  ]  s;     	" " 	" "r   r  c                   d    e Zd ZdZd Zd Zd Zd Zd Zd Z	d Z
d	 Zd
 Zedd       Zd Zd Zy)r  a`  
    Collecting histogram for each tensor. Percentile and Entropy method are supported.

    ref: https://github.com//apache/incubator-mxnet/blob/master/python/mxnet/contrib/quantization.py
    ref: https://docs.nvidia.com/deeplearning/tensorrt/pytorch-quantization-toolkit/docs/_modules/
                 pytorch_quantization/calib/histogram.html
    c                 f    i | _         || _        || _        || _        || _        || _        || _        y ry   )histogram_dictr   r   r  r  r  r  )rD   r   r   r  r  r  r  s          r   rH   zHistogramCollector.__init__|  s5     " "4$ r   c                     | j                   S ry   )r  rL   s    r   get_histogram_dictz%HistogramCollector.get_histogram_dict  s    """r   c                     t        d       | j                  dv r| j                  |      S | j                  dk(  r.| j                  r| j	                  |      S | j                  |      S t        d      )Nz/Collecting tensor data and making histogram ...>   r,   r  r  DOnly 'entropy', 'percentile' or 'distribution' methods are supported)printr   collect_valuer   collect_absolute_valuer>   r  s     r   r  zHistogramCollector.collect  sl    ?@ ;;55%%k22[[L(~~22;??))+66cddr   c                    |j                         D ]K  \  }}t        |t              r|D ]2  }t        |t        j                        rJ dt        |       d|        |D ch c]  }|j                   }}t        |      dk(  sJ d| d|       t        j                  |      }n6t        |t        j                        st        dt        |       d|      |}|j                         }|j                  dkD  r+t        j                  |      }t        j                  |      }	nBt        j                  d|j                        }t        j                  d|j                        }	t        j                  |      }|| j                   vrxt        j"                  || j$                        \  }
}|j'                  |j                        }|j                  t        j(                  k7  sJ d       |
|||	f| j                   |<   | j                   |   }|d	   }|d
   }t+        |d      sJ dt        |              t+        |d      sJ dt        |              |d   }|d   }t        j                  |      }||d   kD  rB|d   |d   z
  }t        j,                  |d   |z   ||z   |      }t        j.                  ||f      }t        j"                  ||      \  }
}|j'                  |j                        }|
dt        |      xxx |z  ccc |j                  t        j(                  k7  sJ d       |
|t1        ||      t3        ||	      f| j                   |<   N yc c}w )z5
        Collect histogram on absolute value
        r8   z for tensor=r
   z6The calibration expects only one element type but got r   r   )r5   zMonly float32 or float16 is supported, every constant must be explicitly typedrl   ro   r   z'old_min should be a numpy array but is r   N)r<   r[   r9   r   r   rA   r   rw   r%   r>   flattensizer  r  r^   absoluter  	histogramr  r&   float64r@   arangehstackrt  ru  )rD   r  re  data_arrarradtypesdata_arr_npr|  r}  r3   r4   old_histogramrx  ry  old_histold_hist_edges	temp_amaxwidthnew_bin_edgess                       r   r  z)HistogramCollector.collect_absolute_value  sZ    !, 1 1 3 4	sFH(D)# mC%c2::6l:J4PS9+Uabhak8ll6m+34a!''446{a' LVHT`ag`jk' !jj2"**5 #3DN3C<PVz!Z[[&%--/K!#IIk2	IIk2	HHQk.?.?@	HHQk.?.?@	++k2KT000#%<<$--#P j'..{/@/@A
"((BJJ6 c6 04ZI.V##F+ $ 3 3F ;'*'*w0k4[\`ah\i[j2kk0w0k4[\`ah\i[j2kk0(+!.q!1IIk2	~b11*1-q0AAE$&IInR.@5.H)V[J[]b$cM%'YY/N%ON#%<<.#Q j'..{/@/@A
_s8}%1%"((BJJ6 c6 04ZWiAXZ]^egpZq.r##F+i4	s 5s   #M!c           	         |j                         D ]a  \  }}t        j                  |      }|j                         }|j                  dkD  r+t        j
                  |      }t        j                  |      }nBt        j                  d|j                        }t        j                  d|j                        }t        j                  t        t        |      t        |            |j                        }|| j                  v r3| j                  |   }| j                  |||||      | j                  |<   &t        j                  || j                  | |f      \  }}	||	|||f| j                  |<   d y)z1
        Collect histogram on real value
        r   r   rH  N)r<   r   r%   r  r  r  r  r^   r   ru  r  r  merge_histogramr  r  )
rD   r  re  r  r|  r}  	thresholdr  r3   r4   s
             r   r  z HistogramCollector.collect_value  s<    !, 1 1 3 	FHzz(+H'')H}}q IIh/	IIh/	HHQhnn=	HHQhnn=	S^S^!DHNN[I,,, $ 3 3F ;.2.B.B!8Y	9/##F+ $&<<$--QZPZ\eOf#g j/##F+)	r   c                    |\  }}}}	}
||
k  rEt        j                  |t        |      |
 |
f      \  }}||z   |t        ||      t	        |	|      |
fS |
dk(  r-t        j                  |t        |      | |f      \  }}||z  }nft        |      }d|
z  |z  }t        ||
z
  |z  dz         }|d|z  z   }||z  |
z   }t        j                  ||| |f      \  }}||||z
  xxx |z  ccc ||t        ||      t	        |	|      |fS )Nr  r   rl   r
   )r   r  rw   rt  ru  r_   )rD   r  r  rz  r{  new_thresholdr  r  rx  ry  old_thresholdnew_histr?  r3   r4   old_num_bins
old_stridehalf_increased_binsnew_num_binss                      r   r  z"HistogramCollector.merge_histogram  sT   FSC>7G]M),,xX~WdFefKHa8#GW%GW%  !#%<<#h-Q^P^`mOn#o j "8}.=
&)==+HZ*WZ[*[&\#+a2E.EE 3j @= P#%<<,P]~_lNm#n j(<:M+MNRZZNGW%GW% r   c                 b   | j                   rt        | j                         dk(  rt        d      t        d| j                  d       | j                  dk(  r| j                         S | j                  dk(  r| j                         S | j                  dk(  r| j                         S t        d      )	Nr   z=Histogram has not been collected. Please run collect() first.z0Finding optimal threshold for each tensor using z algorithm ...r,   r  r  r  )r  rw   r>   r  r   compute_entropycompute_percentilecompute_distributionrL   s    r   r  z,HistogramCollector.compute_collection_result  s    ""c$*=*=&>!&C\]]@~^_;;)#''))[[L(**,,[[N*,,..cddr   c                    | j                   dk  s| j                   dkD  rt        d      | j                  }| j                   }i }t        dt	        |              t        d| j
                          t        dd|z
   d| d	       |j                         D ]  \  }}|d   }|d
   }|j                         }t        j                  ||z        }	| j                  rft        j                  |	|dz        }
t        j                  ||
   |j                         t        j                  ||
   |j                        f||<   nd|z
  dz  }t        j                  |	d|z
        }
t        j                  |	|      }t        j                  ||   |j                        t        j                  ||
   |j                        f||<   |d   }|d   }||   d   |k  r|||   d
   f||<   ||   d
   |kD  r||   d   |f||<   g ||   |d d ||<   t        j                  j!                  dd      dv st#        ||        |S )Nr   d   z<Invalid percentile. Must be in range 0 <= percentile <= 100.Number of tensors : Number of histogram bins : zPercentile : (g      Y@,)r
   r   g      i@r"   rl   ro   QUANTIZATION_DEBUG0r
   1)r  r>   r  r  rw   r  r<   r(   r   cumsumr   searchsortedr^   r   r   environr]   r   )rD   r  r  thresholds_dictre  r  r3   r4   totalcdf	idx_rightpercent_to_cut_one_sideidx_leftr|  r}  s                  r   r  z%HistogramCollector.compute_percentile  sZ   ??Q$//C"7[\\,,__
$S%8$9:;+DMM?;<uz12!J<qAB!/!5!5!7 	-FIQ<D"1JHHJE))D5L)C~~OOCe1CD	 XXj3:;K;KLLHHZ	2*:J:JK+'
 ,1:+=*F'OOC7N1NO	??30GHHHZ19I9IJHHZ	2*:J:JK+' "!I!!Iv&q)I5+4of6Ma6P*Q'v&q)I5+:6+B1+Ey*Q'&K(?&K$r(&KOF#zz~~2C8HD4,;	-> r   c                    | j                   }| j                  }i }t        dt        |              t        d| j                   d       t        d| j                          |j                         D ]^  \  }}| j                  ||      }|||<   g ||d d ||<   t        j                  j                  dd      dv sMt        |d	   |d
          ` |S )Nr  r  z: (The number may increase depends on the data it collects)zNumber of quantized bins : rl   r  r  r  r   r
   )r  r  r  rw   r  r<   get_entropy_thresholdr   r  r]   r   )rD   r  r  r  re  r  optimal_thresholds          r   r  z"HistogramCollector.compute_entropyM  s    ,,!44$S%8$9:;+DMM?:tuv+D,C,C+DEF!/!5!5!7 	7FI $ : :9FX Y&7OF#&J(9&JIbqM&JOF# zz~~2C8HD9Q<16	7 r   c                    |dk  rt        d| d      |d d |dd  z   dz  }|dk(  r| |z  j                         | j                         z  }| |dz  z  j                         | j                         z  |dz  z
  dz  }t        j                  ||j                        t        j                  ||j                        fS t        |      |k(  rt        |      dz  dk(  r| ||z  z  j                         | j                         z  }| ||z  |z
  dz  z  j                         | j                         z  dz  }t        j                  ||j                        t        j                  ||j                        fS t        j                  |      |z  }d|t        j                  |      <   d|t        j                  |      <   t        j                  |      |z  |z  }| |z  j                         | j                         z  }| |dz  z  j                         | j                         z  |dz  z
  dz  }t        j                  ||j                        t        j                  ||j                        fS )	Nr   zpower=z <= 0 is invalid.r   r
   g      ?rl   r   )	r>   r(   r   r^   r   r_   r  isnanisinf)r3   r4   powerr   r/   r0   facts          r   _avg_stdzHistogramCollector._avg_stdb  s   A:veW,=>??Sb/JqrN2c9A:&=%%'$((*4C619$))+dhhj836AcIC88Cz'7'78"((3jN^N^:___u:3u:>Q#6&%-',,.;CFEMC/A55::<txxzIcQC88Cz'7'78"((3jN^N^:___vvf~& RXXd^ RXXd^5(4/f}!!#dhhj0vqy %%'$((*4sAv=#Exx:#3#34bhhs*JZJZ6[[[r   c           
         | j                   dk  rt        d      | j                  }i }t        dt	        |              t        d| j                           t        d| j
                  d       |j                         D ]E  \  }}|d   }|d   }|j                  t        j                  k7  sJ | j
                  d	k(  r| j                  ||d
      \  }}n2| j
                  dk(  r| j                  ||d
      \  }}nt        d      |j                  t        j                  k7  sJ |j                  t        j                  k7  sJ |j                  t        j                  k7  sJ t        |||||j                         |j                               ||<   t        j                  j!                  dd      dv s:t#        ||       H |S )Ni   z3Invalid num_bins. Must be in range 512 <= num_bins.r  r  zScenario : r  r   r
   r  )r  p3gUUUUUU?z,Invalid scenario. Must be in {'same', 'p3'}.)r/   r0   r3   r4   r1   r2   r  r  r  )r  r>   r  r  rw   r  r<   r   r   r  r
  r.   rt  ru  r   r  r]   r   )	rD   r  r  re  r  r3   r4   avg_coefstd_coefs	            r   r  z'HistogramCollector.compute_distributionx  s   ==3RSS,,$S%8$9:;+DMM?;<DMM,A./!/!5!5!7 	-FIQ<D"1J##rzz111}}&%)]]41]%M"($&%)]]49]%U"( !OPP>>RZZ///>>RZZ///##rzz111&0%!~~'"('OF# zz~~2C8HD4,3	-6 r   c           	         |d   }|d   }|j                   }|dz  }|dz  }|d   j                  }t        j                  ||z
  dz         }	t	        |	j                         D 
cg c]0  }
t        j
                  d|      t        j
                  d|      f2 }}
t	        ||dz   d      D ]  }
||
z
  }t        ||
z   dz   |      }||   ||   f||
|z
  <   t        j                  |||       }|j                         }t        |d|       }t        ||d       }|dxx   |z  cc<   |dxx   |z  cc<   |dk7  j                  t        j                        }t        j                  |t        j                        }|j                   |z  }t	        |      D ]  }||z  }||z   }t        |||       ||<    |dxx   t        |||z  d       z  cc<   t        j                  |j                   t        j                        }t	        |      D ]+  }||z  }||z   }t        |||       }|dk7  s!||   |z  ||| - t        |      }t        |      }||&t        j
                  t        j                  |      }n!t        j
                  t        ||      |      }||	|
|z
  <    t        j                  |	      }||   }|d   }|d   }|d   |k  r||d   f}|d   |kD  r|d   |f}t!        |d   d      sJ t!        |d   d      sJ |S c c}
w )	aF  Given a dataset, find the optimal threshold for quantizing it.
        The reference distribution is `q`, and the candidate distribution is `p`.
        `q` is a truncated version of the original distribution.
        Ref: http://on-demand.gputechconf.com/gtc/2017/presentation/s7310-8-bit-inference-with-tensorrt.pdf
        r   r
   rl   r   Nr   ro   r   )r  r   r   zerosrH  r^   rt  r  deepcopyr(   r&   rO  r   r   r,   argminr@   )rD   r  r  r3   r4   r  zero_bin_indexnum_half_quantized_binr   kl_divergencer5  
thresholdsr   r   sliced_distributionpleft_outliers_countright_outliers_countnonzerosquantized_binsnum_merged_binsr<  startendqnormdivmin_kl_divergence_idxr  r|  r}  s                                  r   r  z(HistogramCollector.get_entropy_threshold  sw    |q\
99!Q!3q!8!""2H!H1!LMTYZgZlZlTmnqrxx/!51IJn
n  -~/A1E .	<A(1,KNQ.2H=I6@6MzZcOd5eJq112"&--[0K"L $((*A"%d<K&8"9#&tIJ'7#8 aD''DbE))E Qrxx0H  XX&8IN166:LLO 12 L/o-(+,?c,J(Ku%L 2#&9:L:^:`&a"bb rxx0A12 @/o-8E#./19#1%#84#?AeCL@ $A&A#A&AyAIhhrvvU3hhwq!}E:8;M!445].	<` !#		- 8&'<=aL	aL	Q)+!*,=a,@ AQ)+!21!5y A(+W555(+W555  U os   "5L	N)r
   )rT   rd   re   r   rH   r  r  r  r  r  r  r  r  staticmethodr
  r  r  rZ   r   r   r  r  s  s]    !#e8st@@e,\* \ \*&PX!r   r  r  r   r   c                    d }|t         j                  k(  rp|j                  dd      }|j                  dd      }	|j                  dd      }
|j                  dd       }|j                  dd      }t        | |||||	|
||	      }n |t         j                  k(  rI|j                  d	d
      }|j                  dd
      }|j                  dd      }t        | ||||||      }n|t         j                  k(  rI|j                  d	d      }|j                  dd      }|j                  dd      }t        | ||||||      }nH|t         j                  k(  r5|j                  d	d      }|j                  dd      }t        | |||||      }|r+|j                          |r||_        |j                          |S t        d|       )Nr   Fr&  r'  r  r(  r   )r   r   r&  r'  r(  r   r  r  r  )r   r   r  r  r  r  r  T)r   r   r  r  r  r  )r   r  r  zUnsupported calibration method )ru   rv   r]   r  r   r  r   r  r   r  r  r   r   r>   )r   r   r   calibrate_methodr   r   extra_options
calibratorr   r&  r'  r(  r   r  r  r  r  s                    r   create_calibratorr)    s    J,333!%%k59	&**+;UC*../CTJ#0#4#45OQU#V #''u=%! %=)1%=#


 
.66	6 $$Z5*../CSI!%%k59	&! %=1

 
.99	9 $$Z6"&&|V<
!%%k48	)! %=!

 
.;;	; $$Z6 $$Z8+! %=

   "-6J*++-
67G6HI
JJr   )Nr   )r   r   r   rj   )4r   r   r  r  r   r   r   rL  collections.abcr   enumr   pathlibr   numpyr   r,  r   r   r   r	   r   quant_utilsr   r   r   r   r   r`   r_   r,   r.   rj   ru   ABCMetar   r   r   boolr   r   r   r  r  r  r  r  r  r  rv   rr   r)  rZ   r   r   <module>r1     s        	   $     > >  U U  

 " 	





 $, 	
 ZZ8/ /d=* =*@ "ckk "43d.. 3, `e M  X\ im ,	$k" k"\m,~ m,`DL. DLN
+ 
D
. 
D 
0  
F" ",E!1 E!T 37/&--"NK:NK#C=4/NKr   