
    \j$                       d dl m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
Z
d dlZd dlmZmZmZ d dlmZmZmZ d dlmZ d dlmZmZmZmZ d d	lmZ d d
lmZmZmZ 	 d dl m!Z! dZ#dZ$dZ%dZ&dZ'dZ(dZ)dZ*dZ+dZ,i Z- e.e      D  ci c]  }  e/ e0e|       e1      s e0e|       |  c} Z2 G d de      Z3 G d de      Z4 G d de      Z5 G d de      Z6ej"                  jn                   e
jp                  d      ej"                  jr                   e
jp                  d      ej"                  jt                   e
jp                  d       ej"                  jv                   e
jp                  d!      ej"                  jx                  eej"                  jz                  eej"                  j|                  eiZ?ej"                  jr                   e
j                  d e
j                  "       e
j                  d#e
j                  "      fej"                  jn                   e
j                  d$e
j                  "       e
j                  d%e
j                  "      fej"                  jv                   e
j                  d e
j                  "       e
j                  d&e
j                  "      fej"                  jt                   e
j                  d'e
j                  "       e
j                  d(e
j                  "      fej"                  j|                   e
j                  d e"       e
j                  d)e"      fej"                  jz                   e
j                  d*e"       e
j                  d+e"      fiZEej"                  jn                   e
j                  d,e
j                  "       e
j                  d%e
j                  "      fej"                  jt                   e
j                  d-e
j                  "       e
j                  d(e
j                  "      fiZFej"                  jr                   e
j                  d e
j                  "       e
j                  d%e
j                  "      fej"                  jn                   e
j                  d.e
j                  "       e
j                  d/e
j                  "      fej"                  jv                   e
j                  d e
j                  "       e
j                  d(e
j                  "      fej"                  jt                   e
j                  d0e
j                  "       e
j                  d1e
j                  "      fej"                  j|                   e
j                  d e"       e
j                  d+e"      fej"                  jz                   e
j                  d2e"       e
j                  d3e"      fiZGd4d5d6ZHdbd7ZIdcd8ZJddded9ZKd: ZL	 	 	 	 	 	 	 	 	 	 	 	 dfd;ZM	 	 	 	 dg	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dhd<ZN	 dg	 did=ZO	 	 	 dj	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 dkd>ZPdld?ZQdld@ZRdmdAZSdndBZT G dC dD      ZU G dE dF      ZV G dG dH      ZWdI ZXdJ ZYdK ZZdL Z[dodMZ\dN Z]dpdOZ^dqdPZ_drdQZ`dsdRZadtdSZbdudTZcdtdUZddudVZedvdWZf	 	 	 dj	 	 	 	 	 	 	 	 	 	 	 dwdXZgdxdYZhdydZZidzd[Zjd{d\Zkd{d]Zld|d^Zmd|d_Znd|d`Zod|daZpy# e"$ r dZ!Y fw xY wc c} w )}    )annotationsN)Enum)Path)float8_e4m3fnint4uint4)
ModelProtoTensorProtoexternal_data_helper)onnx_pb)
make_graph
make_model	make_nodemake_tensor_value_info)ReferenceEvaluator)GraphOptimizationLevelInferenceSessionSessionOptions)to_array_extendedzonnx.quantizez0.1.0ai.onnxzcom.microsoftQuantizeLinear_QuantizeLinear_InputDequantizeLinear_DequantizeLinear_Output
_quantizedl        c                  *    e Zd ZdZdZd Zed        Zy)QuantizationModer      c                    | j                   S Nnameselfs    W/root/.hermes/venv/lib/python3.12/site-packages/onnxruntime/quantization/quant_utils.py__str__zQuantizationMode.__str__8       yy    c                D    	 t         |    S # t        $ r t               w xY wr    )r   KeyError
ValueError)modes    r%   from_stringzQuantizationMode.from_string;   s)    	#D)) 	,	    N)__name__
__module____qualname__
IntegerOps
QLinearOpsr&   staticmethodr-    r(   r%   r   r   4   s%    JJ  r(   r   c                  *    e Zd ZdZdZd Zed        Zy)QuantizedValueTyper   r   c                    | j                   S r    r!   r#   s    r%   r&   zQuantizedValueType.__str__G   r'   r(   c                D    	 t         |    S # t        $ r t               w xY wr    )r7   r*   r+   )vs    r%   r-   zQuantizedValueType.from_stringJ   s)    	%a(( 	,	r.   N)r/   r0   r1   InputInitializerr&   r4   r-   r5   r(   r%   r7   r7   C   s%    EK  r(   r7   c                  N    e Zd ZdZdZdZdZdZdZdZ	d Z
ed	        Zed
        Zy)	QuantTyper   r                  c                    | j                   S r    r!   r#   s    r%   r&   zQuantType.__str__[   r'   r(   c                D    	 t         |    S # t        $ r t               w xY wr    )r>   r*   r+   )ts    r%   r-   zQuantType.from_string^   s(    	Q< 	,	r.   c                
   | t         j                  k(  rt        j                  S | t         j                  k(  rt        j
                  S | t         j                  k(  rt        j                  S | t         j                  k(  rt        j                  S | t         j                  k(  rt        j                  S | t         j                  k(  rt        j                  S | t         j                  k(  rt        j                  S t!        d| d      )NzUnexpected value qtype=.)r>   QInt8r
   INT8QUInt8UINT8QUInt16UINT16QInt16INT16QFLOAT8E4M3FNFLOAT8E4M3FNQUInt4UINT4QInt4INT4r+   r#   s    r%   tensor_typezQuantType.tensor_typee   s    9??"###9###$$$9$$$%%%9###$$$9***+++9###$$$9??"###24(!<==r(   N)r/   r0   r1   rI   rK   rQ   rO   rM   rU   rS   r&   r4   r-   propertyrW   r5   r(   r%   r>   r>   R   sR    EFMFGEF   > >r(   r>   c                  *    e Zd ZdZdZd Zed        Zy)QuantFormatr   r   c                    | j                   S r    r!   r#   s    r%   r&   zQuantFormat.__str__|   r'   r(   c                D    	 t         |    S # t        $ r t               w xY wr    )rZ   r*   r+   )formats    r%   r-   zQuantFormat.from_string   s)    	v&& 	,	r.   N)r/   r0   r1   	QOperatorQDQr&   r4   r-   r5   r(   r%   rZ   rZ   x   s%    I
C  r(   rZ   int8uint8int16uint16dtype   i   i  i i     i   iii@   i i @  r@   zero_point_indexc                @   g }t        |      D ]  \  }}t        j                  t        |      t        j                        r%|j                  t        j                  |             n=t        |t        j                        r|j                  |       nt        d| d|       || k(  s|d   }|j                  t        j                  k(  s|j                  t        j                  k(  st        d|j                          t        |      dkD  rt        |      S |d   S )Nzarg z is not an array: rl   zzero_point cannot be r   r   )	enumeratenumpy
issubdtypetypenumberappendarray
isinstancendarray	TypeErrorre   float32float16lentuple)rn   argsnew_argsiar:   s         r%   _check_typer      s    H$ 
C1DGU\\2OOEKKN+5==)OOAd1#%7s;<<  Aww%--'177emm+C"7y ABB
C "(ma/5?@Xa[@r(   c                   | t         v sJ d|  d       | t        j                  j                  t        j                  j                  t        j                  j
                  t        j                  j                  fv r/|dk7  rt        d|d      |j                  t        j                  k(  rt        j                  }nG|j                  t        j                  k(  rt        j                  }nt        d|j                   d      t        t!        t#        dg dgt$        j&                  j)                  d| g dg      	      t#        d
g ddg      gdt+        d|d       t+        d|d       gt+        d| d       g            }t-        |      }t/        |j1                  d ||d      d         S t         |    }	t3        | dd      \  }
}|t5        |
|      n|
}|t7        ||      n|}t        j8                  |j;                  t        j                        |z  j=                         |z         }t        j>                  ||||       t/        |j;                  |	            S )NUnexpected data type > requested. Only INT8, UINT8, INT16, and UINT16 are supported.r   z2zero_point is expected to be null for float 8 not rH   zUnexpected dtype Constant
zero_point)valuer   )Xscaler   Yqur   r   )r   r   Freduce_range	symmetric)out) ONNX_TYPE_TO_NP_TYPE
onnx_protor
   rR   FLOAT8E4M3FNUZ
FLOAT8E5M2FLOAT8E5M2FNUZNotImplementedErrorre   rq   rz   FLOATr{   FLOAT16r+   r   r   r   onnxhelpermake_tensorr   r   r   runget_qmin_qmax_for_qTypemaxminasarrayastyperoundclip)qTypearrr   r   lowhigh	onnx_type
onnx_modelrefre   qminqmaxcliplowcliphigharr_fp32s                  r%   quantize_nparrayr      s-   (( 
w&de( ++--))--	  ?%(Z[eZhhi&jkk99%#))IYY%--'#++I01=>>"Bdkk>U>UVbdikmpqor>s .0LseT	 *3	4@*7ItD (UD9:

  !,3774sU)CDQGHH %U+,URWX
d$'O#dC.&*&63tT?D==#**U]]";e"C!J!J!Lz!YZ

8WhH=8??5122r(   c           	        |dkD  s|dk  rt        d| d|       t        j                  | t        j                  d| j                              } t        j
                  |t        j                  d|j                              }|.t        || t        j                  || j                        z         }|rBt        j
                  t        j                  |       t        j                  |            }| } |}||k  sJ d|  d|        t        j                  || z
  t        j                        }t        j                  |t        j                        t        j                  |t        j                        z
  }t        j                  ||z        }	|	dk\  sJ d       |	t        j                  |j                        j                  k  rFt        j                  d|j                        }	t        j                  d|j                        }
|
|	gS |r^t        j                  t        j                  ||z   t        j                  d	t        j                        z        |j                        }
n:t        j                  t        j                  || |	z  z
        |j                        }
|	j                  |j                        }	|
|	gS )
a  Calculate the scale s and zero point z for the quantization relation
    r = s(q-z), where r are the original values and q are the corresponding
    quantized values.

    r and z are calculated such that every value within [rmin,rmax] has an
    approximate representation within [qmin,qmax]. In addition, qmin <= z <=
    qmax is enforced. If the symmetric flag is set to True, the interval
    [rmin,rmax] is symmetrized to [-absmax, +absmax], where
    absmax = max(abs(rmin), abs(rmax)).

    :parameter rmin: minimum value of r
    :parameter rmax: maximum value of r
    :parameter qmin: minimum value representable by the target quantization data type
    :parameter qmax: maximum value representable by the target quantization data type
    :parameter symmetric: True if the floating-point range should be made symmetric. Defaults to False.
    :parameter min_real_range: Minimum floating-point range (i.e., rmax - rmin) to enforce. Defaults to None.
    :return: zero and scale [z, s]

    r   Bqmin and qmax must meet requirement: qmin <= 0 <= qmax while qmin:, qmmax:rd   zqmin=z > qmax=zscale issue      ?       @)r+   rq   minimumrv   re   maximumr   r   absfloat64finfotinyr   r   )rminrmaxr   r   r   min_real_rangeabsmaxdrdqr   r   s              r%   compute_scale_zpr      s    ( ax4!8]^b]ccklpkqrss
 ==u{{1DJJ?@D==u{{1DJJ?@D !4nDJJ OOPuyy		$@ww4<55htf55<	TD[	6B	T	/%++d%--2X	XBKKR EA:$}$:u{{4::&+++Ctzz2[[$**5
   TD[EKK5==,QQRZ^ZdZdJ U[[u1D%ETZZXJTZZ(r(   c                   t        |      }t        |      }||z   dz   dz  }t        t        j                  |             } t        t        j                  |            }t	        | d      } t        |d      }||dkD  rt        || t        |      z         }|| k  r| dk\  r|n|}t        t        |       t        |            }	||k(  r||z
  nt        d||z
  dz        }
|	dkD  r|	ndt        d|
      z  }|||||z
  z  k  r|||z
  z  }t        j                  |t        j                        t        j                  |t        j                        fS | dk\  rQt        j                  |t        j                        }t        j                  |||z
  z  t        j                        }net        j                  |t        j                        }|  ||z
  z  }|||z
  z  }t        j                  t        ||      t        j                        }|?t        |      |||z
  z  k  r+t        j                  |||z
  z  t        j                        }||fS )a6  Snap a uint8 activation zero-point to qmin (when rmin >= 0) or mid (when rmin < 0).

    Used by the ActivationRestrictedAsymmetric quantization option. Recomputes scale so the
    dequantized range still covers [rmin, rmax] without clipping.

    :parameter rmin: calibrated minimum activation value (numpy scalar)
    :parameter rmax: calibrated maximum activation value (numpy scalar)
    :parameter qmin: minimum quantized value (int, default 0)
    :parameter qmax: maximum quantized value (int, default 255)
    :parameter min_real_range: minimum floating-point range to enforce (same semantics as compute_scale_zp).
        When not None and > 0, rmax is adjusted to max(rmax, rmin + min_real_range) before scale computation.
    :return: (zero_point, scale) with zero_point dtype uint8 and scale dtype float32
    r   r?           r   r   rd   )
intfloatrq   squeezer   r   r   rv   ra   rz   )r   r   r   r   r   qmin_valqmax_valmiddegenerate_zpabs_maxdenom	scale_valr   r   	scale_neg	scale_poss                   r%   snap_zero_point_to_uint8r   ,  s    4yH4yHh"q
(Ct$%Dt$%D tS>DtS>D !nq&84n 556t| %)CKSc$iT+)6()BH$APX[cPchiOiHj '!WAuE	%)nS[H[6\*\&(X*=>I{{=<ekk)[`[h[h>iiis{[[=
DHx$78N [[EKK8
ES8^,	HsN+	C	95U]]K !eEl^xRZGZ5[&[Nh.AB%--Xur(   c                   d}| t         vr| t        j                  k(  rddlm} |}t        d      D cg c]  }t        |       }}t        j                  |D cg c]0  }t        j                  |      rt        j                  |      r/|2 c}t        j                        }nt        d|  d      |t         | <   n| t        j                  k(  rddlm} |}|t        d|  d	      t        j                  t         |          }t        j                  d|      }	t        j                  ||z  |j                        }
|	|
gS c c}w c c}w )
ar  Calculate the scale s for a float8 type (E4M3FN).
    The function assumes the coefficient distribution and the float 8
    distribution are similar to two gaussian laws.

    :return: zero and scale [z, s]

    More details in notebook `quantization_fp8.ipynb
    <https://github.com/microsoft/onnxruntime/blob/main/docs/python/notebooks/quantization_fp8.ipynb>`_.
    Nr   )r      rd   zQuantization to element_type=z not implemented.zUnexpected element_type rH   )FLOAT8_DISTRIBUTIONSr
   rR   	ml_dtypesr   ranger   rq   rv   isnanisinfrz   r+   ry   stdre   )element_typer   zp_dtyper   r   
all_valuesfvaluesstd_f8zeror   s              r%   compute_scale_zp_float8r   g  s    H//;333/$H,1#J7q%(7J7[[&Tqekk!nU[[QR^T\a\i\iF <\NJ[\]]-3\*	11	1+ 2<.BCCYY+L9:F;;q)DKKfCII6E%=# 8Ts   EE5EEc           
        | j                   dk7  r&t        d| j                    d| j                   d      | j                  |   }||z   dz
  |z  }t        t	        j
                  t        | j                        D cg c]  \  }}||k7  s| c}}            }	t	        j                  | |d      }
|
j                  ||	      }
t        |d|      \  }}t        |   d   j                  }||z  |z
  }|dkD  r=t	        j                  ||	f|
j                  	      }t	        j                  |
|gd
      }n|
}|j                  |||	      }|j                  d
      }|j                  d
      }t	        j                   |t	        j"                  |            }t	        j$                  |t	        j"                  |            }|rAt	        j$                  t	        j&                  |      t	        j&                  |            }| }|}t	        j(                  |      }t	        j(                  |      }||z
  j+                  t        j(                        }||z
  }||z  j+                  | j                        }t	        j,                  | j                        j.                  }||k  }t	        j0                  |t	        j2                  |      |      }|r?t        t	        j4                  ||z   dz              }t	        j6                  ||	f||	      }nt	        j4                  ||j+                  t        j(                        |j+                  t        j(                        z  z
        }t	        j8                  |||      }|j+                  |      }d||<   t	        j                  |d|      }t	        j                  |d|      }||fS c c}}w )a  Compute per-block scale and zero-point for a weight tensor.

    The weight is sliced along *axis* into blocks of *block_size* elements.
    Per the ONNX opset-21 spec, QuantizeLinear/DequantizeLinear require the
    scale and zero_point tensors to have the **same rank** as the input tensor.
    Only rank-2 weight tensors are supported; rank > 2 is explicitly rejected.

    Returns arrays with the same rank and dimensions as *weight*, except
    ``shape[axis] == ceil(weight.shape[axis] / block_size)``. This matches
    the ONNX opset-21 QuantizeLinear/DequantizeLinear blocked-quantization spec.

    :param weight: Float32/float16 weight array (must be rank-2).
    :param quant_type: ONNX tensor data type for quantization.
    :param axis: Axis along which to apply block-wise quantization.
    :param block_size: Number of elements per block along *axis*.
    :param symmetric: Whether to use symmetric quantization per block.
    :return: Tuple of (zero_point, scale), each with shape matching *weight*
        except ``shape[axis] == n_blocks``, where ``n_blocks == ceil(weight.shape[axis] / block_size)``.
    :raises NotImplementedError: If weight rank is not 2 (opset-21 constraint).
    r?   zXPer-block (opset-21) quantization is only supported for rank-2 weight tensors. Got rank-z tensor with shape zY. For rank > 2 tensors, reshape to 2-D before quantizing or use per-channel quantization.r   r   Fr   rd   )axisr   )ndimr   shaper   rq   prodrp   moveaxisreshaper   ONNX_INT_TYPE_RANGEre   zerosconcatenater   r   r   
zeros_liker   r   r   r   r   r   where	ones_liker   fullr   ) weight
quant_typer   
block_sizer   kn_blocksr   dothermovedr   r   r   pad_lenpadmoved_paddedblocksr   r   r   r   r   r   r   	raw_scaler   
degeneratescaleszp_valzero_pointsraw_zps                                    r%   compute_scale_zp_blockedr    s/   6 {{a!}$7~ Fff
 	
 	TAJ"z1H 

)FLL*AO$!QQ$YAOPQENN64+EMM!U#E(%S\]JD$":.q177H #a'G{kk7E*%++>((%A> !!(J>F ::1:D::1:D ==u//56D==u//56Duyy		$@w}}T"H}}T"H
+		emm	,B	H	Bb  .I;;v||$))DT!J[[U__Y%?KFU[[(X"5!<=>jj(E!2F(KXEMM(BV]]SXS`S`Ea(aabFHh7mmH-"#J
 ^^FAt,F..a6Ku Ps   <N>
N>c                   t        | t        j                        st        dt	        |        d      ||}nt        |       r| j                         nd}||}nt        |       r| j                         nd}t        j                  || j                        }t        j                  || j                        }t        j                  d| j                        }	|t        j                  k(  r?|rt        d      t        j                  |       }
t        ||
      \  }}	t        ||	d      S |t        j                   t        j"                  t        j$                  t        j&                  t        j(                  t        j*                  fv r_t-        |||	      \  }}t        |       rt/        ||||||      \  }}	n!t        j                  d|j                        }t        ||	d      S t1        d
| d      )a  
    Returns the zero_point and scale for the given data.

    :param data: The data for which to compute quantization parameters.
    :param quant_type: The quantization data type.
    :param symmetric: whether symmetric quantization is used or not.
    :parameter reduce_range: True if the quantization range should be reduced. Defaults to False.
    :parameter min_real_range: Minimum floating-point range (i.e., rmax - rmin) to enforce. Defaults to None.
    :parameter rmin_override: The value of rmin to use if not None. Otherwise, uses min(data).
    :parameter rmax_override: The value of rmax to use if not None. Otherwise, uses max(data).
    :return: zero point and scale
    z%Weight must be given as an array not rH   r   rd   r   z1Unsupported option reduce_range=True for float 8.r   rm   r   z Unexpected value for quant_type=)rw   rq   rx   ry   rs   r|   r   r   rv   re   r
   rR   RuntimeErrorr   r   r   rJ   rL   rP   rN   rV   rT   r   r   r+   )datar   r   r   r   rmin_overridermax_overrider   r   r   r   r   r   r   s                 r%   compute_data_quant_paramsr	    s   * dEMM*?T
|1MNN  YtxxzC  YtxxzC;;t4::.D;;t4::.DKK4::.E[---RSSiio3JD
E:uqAA  -ZQZ[
dt9 0tT4Tb cJQdjj9J:uqAA
7
|1E
FFr(   c                   t        | ||||||      \  }}|t        j                  k(  rt        || ||      }	t	        |	j                  t        j                        j                         dz  dk(        ret        j                  |       }
t        d|
j                          d|
j                          d|	j                          d|	j                          d	      |||	fS |t        j                  t        j                  t        j                  t        j                   t        j"                  t        j$                  fv rt        || ||      }	|||	fS t'        d| d      )al  
    :param data: data to quantize
    :param qType: data type to quantize to.
    :param symmetric: whether symmetric quantization is used or not.
    :parameter reduce_range: True if the quantization range should be reduced. Defaults to False.
    :parameter min_real_range: Minimum floating-point range (i.e., rmax - rmin) to enforce. Defaults to None.
    :parameter rmin_override: The value of rmin to use if not None. Otherwise, uses min(data).
    :parameter rmax_override: The value of rmax to use if not None. Otherwise, uses max(data).
    :return: minimum, maximum, zero point, scale, and quantized weights

    To pack weights, we compute a linear transformation

    - when data `type == uint8` mode, from `[rmin, rmax]` -> :math:`[0, 2^{b-1}]` and
    - when data `type == int8`, from `[-m , m]` -> :math:`[-(2^{b-1}-1), 2^{b-1}-1]` where
        `m = max(abs(rmin), abs(rmax))`

    and add necessary intermediate nodes to transform quantized weight to full weight using the equation

    :math:`r = S(q-z)`, where

    - *r*: real original value
    - *q*: quantized value
    - *S*: scale
    - *z*: zero point
    rg   z+One of the quantized value is NaN data in [z, z], quantized_data in [z].zUnexpected value for qType=rH   )r	  r
   rR   r   anyviewrq   ra   ravelr   r  r   r   rJ   rL   rP   rN   rV   rT   r+   )r  r   r   r   r   r  r  r   r   quantized_datanp_datas              r%   quantize_datar  ,  sZ   8 2J ((()%ujI##EKK06683>3FGmmD)G=gkkm_Bw{{}o ^&&4&8&8&:%;2n>P>P>R=SSUW  5.00  *%ujI5.00
25';
<<r(   c                P
   t        |       }d}||dkD  r|j                  |   }	t        t        j                  t        |j                        D 
cg c]  \  }
}|
|k7  s| c}}
            }t        j                  ||d      j                  |	|      }t        j                  ||d      }t        j                  ||d      }|j                  d   }t        j                  j                  |      }t        j                  ||      }t        |      D ]Z  }||z  }t        ||z   |	      }t        |      D ]6  }t        |||||f   j                         |||f   |||f         ||||f<   8 \ t        j                  |j                  |	gt        |j                        D 
cg c]  \  }
}|
|k7  s| c}}
z         d|      }n|t        ||j                         ||      }n|j                  |   }t!        |j                        }d||<   g }t        |      D ]m  }
|j#                  |
|      }||
   }||
   }t        ||j                         ||      }|j%                  t        j&                  |      j                  |             o t        j(                  ||      }|r|n| j*                   t,         }|t        j.                  j0                  k(  r"t        j.                         }||_        |j4                  j7                  | j4                         ||_        |j9                         j;                         j=                         |_        t@        tA        |      } | j                  |j                  k7  s!| j=                         |j=                         k7  r]tC        d|j                   d|j=                         dd  d| j=                         dd  d	| j                   d
tE        |      dd  d      |S |t        j.                  jF                  t        j.                  jH                  fv ry|jJ                  tL        tN        fvrtC        d| d      tQ        tS        |j=                                     }!t        j                  jU                  ||| j4                  |!d      }|S t        j                  j                  |      }t        j&                  ||      j                  | j4                        }t        jV                  jY                  ||      }|S c c}}
w c c}}
w )a  
    Returns a quantized version of the given ONNX initializer.

    :param weight: The ONNX initializer to quantize.
    :param quant_type: The final quantized data type.
    :param zero_point: The zero-point value to use for quantization.
    :param scale: The scale value to use for quantization.
    :param axis: The quantization axis if quantizing per-channel or per-block. Defaults to None.
    :param quant_weight_name: The name of the quantized initializer.
                              If not specified, the quantized name is generated.
    :param block_size: Block size for opset-21 block-wise quantization. 0 means disabled.
    :return: The quantized ONNX initializer.
    Nr   rd   r   zThe initializer of shape z! could not be created, expecting 
   z, got z and shape=z
raw=   rH   zQuantized weights for z. must be 8-bit before packing as 4-bit values.T)raw)-tensor_proto_to_arrayr   r   rq   r   rp   r   r   r   r   tensor_dtype_to_np_dtype
empty_liker   r   r   r  listtakeru   r   r   r"   TENSOR_NAME_QUANT_SUFFIXr
   rR   	data_typedimsextendflattencopytobytesraw_datar   r  strrV   rT   re   r   r   bytespack_bytes_to_4bitr   numpy_helper
from_array)"r   r   r   r   r   quant_weight_namer   weight_dataq_weight_datar   r   r   r   r   scale_movedzp_movedr   quant_np_dtypeq_movedblkstartendcolchannel_countchannel_dimsquantized_channel_data_listchannel_datachannel_scalechannel_zero_pointquantized_channel_dataq_weight_nameq_weight_initializercheckpacked_datas"                                     r%   quantize_onnx_initializerr=  i  s   , (/K*.MJNd#EJJi8I8I.JXdaaSWiXYZ{D!4<<QF nnUD!4>>*dA6$$Q'==jI""5?? 	C*$Eej(!,CU| *:eCin 5 ; ; ={3PS8?TV^_bdg_gVh+c	3'	 OOQC;;L;L1M"[AQRVZQZ1"[[\^_ae
 
([5F5F5H%Q[\#))$/K--.T&(#}% 	lA&++At4L!!HM!+A%5L..0-AS&" (..u}}=S/T/\/\]i/jk	l ))*EtL):%6;;-PhOi@jMT%%222#//1)3&!!((5$1!(5(=(=(?(D(D(F(N(N(P%( &&:;E{{k///5==?mF[F[F]3]"/0A0A/BBc$,,.s34F5==?3B;O:PP[\b\h\h[iS!56t<=Q@ (   
((--t/?/?/E/EF	FtUm3!7Ftuvv .}/D/D/FGH  ${{66}jRXR]R]_jpt6u  	 ==jIm>JRRSYS^S^_#00;;M=YQ  Y" #\s   T T&T"4T"c                j   | t         j                  j                  k(  rt        d      d}|rt        j                  |       }n)|r| t        v r
t        |    }nt        j                  |       }|st        d|  d      |\  }}|dkD  s|dk  r't        d| d| d|j                   d	| d
| d|        |S )z
    Return qmin and qmax, the minimum and maximum value representable by the given qType
    :parameter qType: onnx.onnx_pb.TensorProto.UINT8 or onnx.onnx_pb.TensorProto.UINT8
    :return: qmin, qmax
    z;This function is not implemented for float 8 as not needed.Nr   r   r   r   r   z, dtype=z, reduce_range=z, symmetric=z, qType=)
r   r
   rR   r   ONNX_INT_TYPE_REDUCED_RANGEgetONNX_INT_TYPE_SYMMETRIC_RANGEr   r+   re   )r   r   r   qranger   r   s         r%   r   r     s     
&&333!"_``F,007	u ==.u5$((/07uvwwJD$ax4!86$x

|?<. Y"8E74
 	
 Mr(   c                .    t        | ||      \  }}||z
  S )z
    Helper function to get the quantization range for a type.
        parameter qType: quantization type.
        return: quantization range.
    r  )r   )r   r   r   r   r   s        r%   get_qrange_for_qTyperD    s      )	RJD$$;r(   c                :    | dk  r| |z   n| }|dk\  xr ||k  }||fS )z
    Helper function that tries to return a normalized axis in the range [0, rank - 1].
    :parameter axis: The axis to normalize.
    :parameter rank: The tensor rank (number of dimensions).
    :return (is_valid, axis_norm)
    r   r5   )r   rank	axis_normis_valids       r%   normalize_axisrI    s3      $axtTIA~2)d"2HYr(   c                    t        |       }|dk(  r
t               S |dz   dz  }t        |      }d}d}||dz
  k  r-| |dz      dz  dz  | |   dz  z  ||<   |dz  }|dz  }||dz
  k  r-||k  r| |   dz  ||<   |S )aB  
    Copies a source array of 8-bit values into a destination bytearray of packed 4-bit values.
    Assumes that the source values are already in the appropriate int4 range.
    :parameter src_8bit: The 8-bit element values to pack.
    :return A bytearray with every two 8-bit src elements packed into a single byte.
    r   r   r?   rh   rA   )r|   	bytearray)src_8bit	num_elemsdst_sizedstsrc_idst_is         r%   r$  r$    s     HIA~{A!#H
H
CEE )a-
	*S0Q68E?S;PQE


 )a-

 ye_s*E
Jr(   c                      e Zd ZdZg g dfdZy)QuantizedInitializerzJ
    Represents a linearly quantized weight input from ONNX operators
    Nc
                    || _         || _        || _        || _        || _        || _        || _        || _        |	| _        y r    )	r"   initializerrminsrmaxsr   r   r  r  r   )
r$   r"   rU  rV  rW  r   r   r  r  r   s
             r%   __init__zQuantizedInitializer.__init__(  sF     	&

&	,	r(   r/   r0   r1   __doc__rX  r5   r(   r%   rS  rS  #  s     r(   rS  c                       e Zd ZdZ	 	 	 	 ddZy)QuantizedValuezI
    Represents a linearly quantized value (input\output\intializer)
    Nc
                    || _         || _        || _        || _        || _        || _        || _        || _        |	| _        y r    )	original_nameq_name
scale_namezp_name
value_typer   	node_type
node_qtype
scale_type)
r$   r"   new_quantized_namer`  zero_point_namequantized_value_typer   rc  rd  re  s
             r%   rX  zQuantizedValue.__init__G  sD     "($&.	"$$r(   )NNNNrY  r5   r(   r%   r\  r\  B  s     %r(   r\  c                      e Zd ZdZd Zy)BiasToQuantizez+
    Represents a bias to be quantized
    c                .    || _         || _        || _        y r    )	bias_name
input_nameweight_name)r$   rl  rm  rn  s       r%   rX  zBiasToQuantize.__init__c  s    "$&r(   NrY  r5   r(   r%   rj  rj  ^  s    'r(   rj  c                   | j                   dk(  rt        d| j                   d      | j                   dk(  r| j                  }n#| j                   dk(  r| j                  }n| j                   dk(  r| j
                  }n| j                   dk(  r| j                  }n| j                   dk(  r| j                  }n| j                   d	k(  r| j                  }n| j                   d
k(  r| j                  }nz| j                   dk(  r| j                  }n^| j                   dk(  r| j                  }nB| j                   dk(  r| j                  }n&t        d| j                   d| j                    d      | j                  |iS )z
    Convert attribute to kwarg format for use with onnx.helper.make_node.
        :parameter attribute: attribute in AttributeProto format.
        :return: attribute in {key: value} format.
    r   z
attribute z does not have type specified.r   r?   r@   rA   rB   rC   ri      	   r  z has unsupported type rH   )rs   r+   r"   r   r   srF   gfloatsintsstringstensorsgraphs)	attributer   s     r%   attribute_to_kwargrz  i  s;    ~~:inn%55STUU ~~	1		1		1		1		1	  	1		1	!!	1	!!	2	  :inn%55KINNK[[\]^^NNE""r(   c                t    |D cg c]  }|j                   | k(  s| }}t        |      dkD  r|d   S dS c c}w )z
    Helper function to find item by name in a list.
        parameter item_name: name of the item.
        parameter item_list: list of items.
        return: item if found. None otherwise.
    r   N)r"   r|   )	item_name	item_listitemitemss       r%   find_by_namer    sA     (Bd499	+ATBEB5zA~58/4/ Cs   55c                R    d}t        t        |            D ]  }||   | k(  s|} |S )zC
    Helper function to return index of an item in a node list
    rl   )r   r|   )	elem_name	elem_listelem_idxr   s       r%   get_elem_indexr    s9     H3y>" Q<9$H Or(   c                H    t         j                  j                  d| |g|      S )z
    Helper function to create a Mul node.
        parameter inputs: list of input names.
        parameter output: output name.
        parameter name: name of the node.
        return: Mul node in NodeProto format.
    Mul)r   r   r   )inputsoutputr"   s      r%   get_mul_noder    s!     ;;  $??r(   c                l    | j                   j                  | j                  |z   | j                  z         S )zp
    Helper function to generate a identifiable filepath by concatenating the given identifier as a suffix.
    )parentjoinpathstemsuffix)filename
identifiers     r%   generate_identified_filenamer    s+     ??##HMMJ$>$PQQr(   c                `   dd l }dd lm} dd l} |j                  |j
                         t        d       t        |        t        d       t        |       |j                  | |d       |j                  d       |j                  d       |j                  d	       |j                          y )
Nr   )	thresholdz
Histogram:zHistogram Edges:T)fillzTensor valueCountszTensor value V.S. Counts)sysmatplotlib.pyplotpyplotrq   set_printoptionsmaxsizeprintstairsxlabelylabeltitleshow)hist
hist_edgesr  pltrq   s        r%   
apply_plotr    s    #ES[[1	,	$K	
	*JJtZdJ+JJ~JJxII()HHJr(   c           	        ddl }ddl}ddl}ddlmc mc m} ddlmc mc m} ddl	m
} t        j                  d|         |j                  | |      }t        t        j                   j#                  |d      d      5 }	|	j%                  |       ddd       |j'                  d      }
|j)                  d      }g }t+        | j-                               D ]  }| |   }|j/                         }t1        |j3                  d	|
      j5                               t1        |j3                  d
|
      j5                               g}t7        t9        |            }|j;                  |      }|j;                  |      }|j=                  |       |j?                  ||       |jA                  ||       |jC                  |      }|jE                  |        |jG                  |tI        |             |D ]  }|jK                  |        |jM                         }|jO                  |       |jQ                  ||       |jS                  |      }|jU                  |       |jW                         }t        t        j                   j#                  |d      d      5 }	|	j%                  |       ddd       t        jX                  j3                  dd      dv r|j                  j[                  |d      }|j]                         }t_        |      D ]Y  }|ja                  |      }t        j                  |jc                                t        j                  |je                                [ t        t        j                   j#                  |d      d      5 }	t+        | j-                               D ]  }| |   }|j/                         }t1        |j3                  d	|
      j5                               t1        |j3                  d
|
      j5                               g}|dz   t7        t9        |            z   }|	j%                  |       |	j%                  d        	 ddd       y# 1 sw Y   xY w# 1 sw Y   xY w# 1 sw Y   yxY w)z>
    Helper function to write calibration table to files.
    r   N)CalibrationCacheEncoderzcalibration cache: )clszcalibration.jsonwi   highestlowestzcalibration.flatbufferswbQUANTIZATION_DEBUG0)r   1zcalibration.cache 
)3jsonflatbuffersrq   5onnxruntime.quantization.CalTableFlatBuffers.KeyValuequantizationCalTableFlatBuffersKeyValue5onnxruntime.quantization.CalTableFlatBuffers.TrtTableTrtTable"onnxruntime.quantization.calibrater  logginginfodumpsopenospathjoinwriterv   Buildersortedkeysto_dictr   r@  r~  r"  r   CreateStringKeyValueStartKeyValueAddKeyKeyValueAddValueKeyValueEndru   TrtTableStartDictVectorr|   PrependUOffsetTRelative	EndVectorTrtTableStartTrtTableAddDictTrtTableEndFinishOutputenvironGetRootAsTrtTable
DictLengthr   DictKeyValue)calibration_cachedirr  r  npr  r  r  	json_datafiler   builderkey_value_listkeyr   d_valuesrt  r   flat_key
flat_value	key_value	main_dict	cal_tablebufdict_lenr   s                             r%   write_calibration_tabler    s   
 LLLL KLL&'8&9:;

,2I
JI	bggll3 23S	9 T

9 88A;D!!$'GN',,./ )"3'>>#(,,y$/4467(,,x.3356
 CK '',))%0
w'2!!':6((1	i(#)& $$Wc..AB# 3	''	23!!#I7#Wi0$$W-INN9
..
C	bggll3 9:D	A T

3 
zz~~*C0H<%%77Q?	'')x 	,A!q)ILL)LL*+	, 
bggll3 34c	: 
d+0023 		C&s+F~~'Hhll9d388:;hll8T2779:F #ICK 00EJJuJJt		
 
g L 
 
s%    QQ$CQ1Q!$Q.1Q:c                   | dk(  j                  t        j                        }| dk7  j                  t        j                        }|j                         }| j                  |z
  }|sy|t        |      z  t        |      z  }|dk  sJ d| d| d|        | j                  t        j                        }|||z  | |z  z   z  }|dk  j                         dk(  sJ |S )a~  Given a discrete distribution (may have not been normalized to 1),
    smooth it by replacing zeros with eps multiplied by a scaling factor
    and taking the corresponding amount off the non-zero values.
    Ref: http://web.engr.illinois.edu/~hanj/cs412/bk3/KL-divergence.pdf
         https://github.com//apache/incubator-mxnet/blob/master/python/mxnet/contrib/quantization.py
    r   Nr   zn_zeros=z, n_nonzeros=z, eps1=)r   rq   rz   sumsizer   )pepsis_zerosis_nonzerosn_zeros
n_nonzeroseps1r  s           r%   smooth_distributionr    s     Qu}}-H6//%--0KllnG'!Jw%
"33D#:Q'-
|74&QQ:88EMM"DC(Nte{222DAI??!!!Kr(   c                    t        j                  | j                         d      }t        d |j                  j
                  D              S )NF)load_external_datac              3  F   K   | ]  }t        j                  |        y wr    )r   uses_external_data).0
intializers     r%   	<genexpr>z*model_has_external_data.<locals>.<genexpr>8  s     mz#66zBms   !)r   loadas_posixr  graphrU  )
model_pathmodels     r%   model_has_external_datar  6  s9    IIj))+FEmUZU`U`UlUlmmmr(   c                    t               }|j                         |_        t        j                  |_        i }dg|d<   t        | j                         |fddgi|}y)z
        Generate model that applies graph optimization (constant folding, etc.)
        parameter model_path: path to the original onnx model
        parameter opt_model_path: path to the optimized onnx model
    :return: optimized onnx model
    ConstantSharingdisabled_optimizers	providersCPUExecutionProviderN)r   r  optimized_model_filepathr   ORT_ENABLE_BASICgraph_optimization_levelr   )r   opt_model_pathsess_optionkwargs_s        r%   optimize_modelr  ;  sb     !"K+9+B+B+DK(+A+R+RK(F%6$7F !,,.jH^G_jcijAr(   c                    ddi}| j                   r8| j                   D ])  }|j                  |j                  |j                  i       + t        j
                  j                  | |       y)z>Tag the model that it went through quantization pre-processingonnx.quant.pre_processonnxruntime.quantNmetadata_propsupdater  r   r   r   set_model_props)r  r  props      r%   add_pre_process_metadatar  K  sZ    .0CDN(( 	:D!!488TZZ"89	:KK~6r(   c                    | j                   r2| j                   D ]#  }|j                  dk(  s|j                  dk(  s# y y)zCCheck the model whether it went through quantization pre-processingr  r  TFr  r  r   )r  r  s     r%   model_has_pre_process_metadatar  T  sA    (( 	Dxx33

FY8Y	 r(   c                    ddi}| j                   r8| j                   D ])  }|j                  |j                  |j                  i       + t        j
                  j                  | |       y )N
onnx.inferr  r  )r  r  r  s      r%   add_infer_metadatar  ]  sZ    "$78N%% 	4A!!155!''"23	4KK~6r(   c                    | j                   r2| j                   D ]#  }|j                  dk(  s|j                  dk(  s# y y)Nr  r  TFr  )r  r  s     r%   model_has_infer_metadatar   e  s@    %% 	Auu$4G)G	 r(   c                    | j                   D cg c]   }|j                  r|j                  dk(  s|" }}t        |      dk7  rt        d      |d   j                  }|S c c}w )Nr   r   z$Failed to find proper ai.onnx domainr   )opset_importdomainr|   r+   version)r  opsetai_onnx_domainopset_versions       r%   get_opset_versionr(  m  se    ).););m5<<SXS_S_clSlemNm
>a?@@"1%--M ns
    A A c                0   t        |       }|}t        |d|      }|t        |d|      nd }t        j                  j                  t        j                  j
                  f}	||	v xs ||	v }
|
s|rt        j                  t        j                  h}	 |j                         D ]S  }|D ]H  }|j                  d      }||v rd}
 n/|j                  d      }|0|j                  d      }||v sFd}
 n |
sS n |dk  r!|dkD  rt        j                  d| d	       d}n|d
k  r9|t        j                  j                   k(  rt        j                  d| d       d
}nb|dk  r|
rt        j                  d| d       d}n?|dk(  rt        j                  d| d       n |dk  rt        j                  d| d       d}||k7  r+t        j"                  j%                  | |      } t'        |       } | S # t        t        f$ r t        j                  d       Y w xY w)NrW   r   TconvertzYSkipping 16-bit opset bump heuristic for TensorQuantOverrides: structure not as expected.   r   z$The original model opset version is z, which does not support block-wise quantization natively. Please update the model to opset >= 21. Automatically updating the model to opset 21. Please verify the quantized model.   z, which does not support quantization to float 8. Please update the model to opset >= 19. Automatically update the model to opset 19. Please verify the quantized model.z, which does not support 16-bit integer quantization natively. Please update the model to opset >= 21. Automatically update the model to opset 21. Please verify the quantized model.r  ze, which does not support node fusions. Please update the model to opset >= 11 for better performance.z, which does not support quantization. Please update the model to opset >= 11. Automatically update the model to opset 11. Please verify the quantized model.   )r(  getattrr   r
   rN   rP   r>   rO   rM   r   r@  AttributeErrorry   r  debugwarningrR   version_converterconvert_version&save_and_reload_model_with_shape_infer)r  weight_typeactivation_typetensor_quant_overridesr   r'  target_opset_versionweight_quant_typeactivation_quant_type_int16_typesneeds_opset21_for_16bit_int16_quant_typesoverrides_listoverrideqtr*  
convert_qts                    r%   update_opset_versionrB  v  sz    &e,M(]KHDSD_@ei  $$++T-=-=-C-CDL/<?hCX\hCh #'='..	0A0AB	w"8"?"?"A  . 
"H!l3B//26/&ll95G*%,[[%>
%);;6:3!
" +& rj1n2=/ B1 1	
  "		 1T5E5E5R5R R2=/ B1 1	

  "		 72=/ B1 1	
  "	"	2=/ BM M	

 
	2=/ B1 1	

  "},&&66u>RS 7u=Lg 	* 	w MMuv	ws%   AG- G- *G- 2G- -$HHc                    t        | d      }t        j                  j                  t	        |       t	        |             t        j
                  |j                               }t        |       |j                          |S )Nz	-inferred)	r  r   shape_inferenceinfer_shapes_pathr"  r  r  r  unlink)r   inferred_model_pathr  s      r%   load_model_with_shape_inferrH    s`    6z;O**3z?C@S<TUII)2245Eu Lr(   c                   t        j                  d      5 }t        j                  |       }t	        |      j                  d      }t        j                  ||j                         d       t        |      cd d d        S # 1 sw Y   y xY w)Nz
ort.quant.)prefixz
model.onnxT)save_as_external_data)
tempfileTemporaryDirectoryr  deepcopyr   r  r   
save_modelr  rH  )r  quant_tmp_dir
model_copyr   s       r%   r4  r4    sl    		$	$L	9 7]]]5)
-(11,?

J$7$7$9QUV*:6	7 7 7s   A BB
c                   | j                   t        j                  j                  t        j                  j                  fv rt
        j                  j                  |       S t        d| j                   dt        | j                             )Nz&Only float type is supported. Weights z is )r  r   r
   r   r   r   r%  to_arrayr+   r"   type_to_name)rU  s    r%   r  r    su    !7!7!=!=z?U?U?]?] ^^  ))+66

01A1A0B$|T_TiTiGjFkl r(   c                    | dz   S )N_QuantizeLinearr5   tensor_names    r%   add_quant_suffixrY    s    ***r(   c                    | t         z   S r    )QUANT_INPUT_SUFFIXrW  s    r%   add_quant_input_suffixr\    s    +++r(   c                    | dz   S )N_QuantizeLinear_Outputr5   rW  s    r%   add_quant_output_suffixr_    s    111r(   c                    | dz   S )N_DequantizeLinearr5   rW  s    r%   add_dequant_suffixrb    s    ,,,r(   c                    | dz   S )N_DequantizeLinear_Inputr5   rW  s    r%   add_dequant_input_suffixre    s    222r(   c                    | t         z   S r    )DEQUANT_OUTPUT_SUFFIXrW  s    r%   add_dequant_output_suffixrh    s    ...r(   )NN)FN)r   rf   N)r   r   r   r   r   float | None)r   numpy.ndarrayr   r   r   r   r   r   r   boolreturn#tuple[numpy.ndarray, numpy.ndarray])FNNN)r  rj  r   onnx.TensorProto.DataTyper   rk  r   rk  r   ri  r  ri  r  ri  rl  rm  )rl  z2tuple[numpy.ndarray, numpy.ndarray, numpy.ndarray])NNr   )r   onnx.TensorProtor   rn  r   rj  r   rj  r   z
int | Noner'  z
str | Noner   r   rl  ro  )FF)r   r   rF  r   rl  ztuple[bool, int])rL  r#  rl  rK  )r  r   r  r"  rl  r   )rH   )g-C6?)r   r   )r   r   r  r   )r  r	   )r  r	   rl  rk  )r  r	   rl  r   )r  r	   r5  r>   r6  zQuantType | Noner7  zdict | Noner   r   rl  r	   )r   r   rl  r	   )r  r	   rl  r	   )rU  r
   rl  rj  )rX  r"  rl  r"  )rl  r"  )q
__future__r   r  r  r  rL  enumr   pathlibr   rq   r   r   r   r   r   r	   r
   r   r   r   onnx.helperr   r   r   r   onnx.referencer   onnxruntimer   r   r   onnx.reference.op_runr   ImportError__producer____version__onnx_domain	ms_domainQUANT_OP_NAMEr[  DEQUANT_OP_NAMErg  r  MODEL_SIZE_THRESHOLDr   r  rw   r.  r   rT  r   r7   r>   rZ   rJ   re   rL   rP   rN   rR   rV   rT   r   rv   ra   r`   rc   rb   r   rA  r?  r   r   r   r   r   r  r	  r  r=  r   rD  rI  r$  rS  r\  rj  rz  r  r  r  r  r  r  r  r  r  r  r  r  r   r(  rB  rH  r4  r  rY  r\  r_  rb  re  rh  )r   s   0r%   <module>r     s>   #   	      0 0 > > & Q Q - P P7 	 , $2 ' !  474Dqq
SZ[fhiSjloHpQ'*qt  #> #>L$   V!4  +%++g"6  +%++g"6!!;5;;x#8''  %    ;5;;q#DkekkRU]b]h]hFi"j+%++d%**"E{u{{SV^c^h^hGi!j!!KEKK$FTYafamamHn#o  ;5;;vU[[#I;5;;W\didodoKp"q  ;5;;q#>BV[@\"]+%++b"={u{{1TX?Y!Z  +%++d%**"E{u{{SV^c^h^hGi!j  ;5;;vU[[#I;5;;W\didodoKp"q!    ;5;;q#DkekkRU]b]h]hFi"j+%++c"DkekkRT\a\f\fFg!h!!KEKK$FTYafamamHn#o  ;5;;vU[[#I;5;;W\didodoKp"q  ;5;;q#>AUZ@["\+%++b"={u{{1TX?Y!Z  )+ A 13h<~8v!H``` ` 	`
 ` )`N #'"&"&;G
;G);G ;G 	;G
 !;G  ;G  ;G );G~ hl:=7:=D $(c c )c  c  	c 
 c  "c  c  c L@	< >% %8' '"#J0@R$Rj2n
k 77 )-*.WWW &W (	W
 W Wt7+,2-3/G'  $ rs   "[ [[[[