Ë
    ÝÍ:j@#  ã                   ó’  — d dl Z d dlZd dlZd dlZd dlmZ d dlZd dl	m
Z
mZmZmZ d dlmZ ddlmZ ddlmZ  ej(                  e«      Z G d„ d«      Zd	„ Zed
k(  rë e«       Zej4                  rej7                  ej8                  «       ej:                  Zej>                  Z ejB                  jE                  e «      r!ejG                  de › d�«        e$de › d�«      ‚ ejJ                  e«      Z& ee&ejN                  ejP                  ejR                  ¬«      Z*e*jW                  «        e*jL                  jY                  e d«       yy)é    N)Ú
GraphProtoÚ
ModelProtoÚ	NodeProtoÚTensorProto)Úquantize_matmul_bnb4é   )Ú	ONNXModel)Úattribute_to_kwargc                   óÂ   — e Zd ZdZdZdZddededefd„Ze	d	e
e   d
eeef   fd„«       Zdej                   d
ej$                  fd„Zdede
e   d
efd„Zde
e   fd„Zd„ Zy)ÚMatMulBnb4QuantizerzMPerform 4b quantization of constant MatMul weights using FP4 or NF4 data typer   r   NÚmodelÚ
quant_typeÚ
block_sizec                 ó´   — |xs g }|t         j                  t         j                  fv sJ ‚t        |«      | _        || _        || _        t        |«      | _        y ©N)	r   ÚFP4ÚNF4r	   r   r   r   ÚsetÚnodes_to_exclude)Úselfr   r   r   r   s        úƒ/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/onnxruntime/quantization/matmul_bnb4_quantizer.pyÚ__init__zMatMulBnb4Quantizer.__init__%   sV   € Ø+Ò1¨rÐØÔ1×5Ñ5Ô7J×7NÑ7NÐOÑOÐOÐOÜ˜uÓ%ˆŒ
Ø$ˆŒØ$ˆŒÜ #Ð$4Ó 5ˆÕó    Ú
graph_pathÚreturnc                 óš   — t        t        |«      dz
  dd«      D ]/  }||   }|j                  D ]  }|j                  | k(  sŒ||fc c S  Œ1 y)Nr   éÿÿÿÿ)NN)ÚrangeÚlenÚinitializerÚname)r!   r   ÚgidÚgraphÚtensors        r   Ú__get_initializerz%MatMulBnb4Quantizer.__get_initializer-   s\   € äœ˜Z›¨1Ñ,¨b°"Ó5ò 	)ˆCØ˜s‘OˆEØ×+Ñ+ò )�Ø—;‘; $Ó&Ø! 5˜=Ô(ñ)ð	)ð
 r   Úfpweightc           	      ó˜  — t        |j                  «      dk7  rt        d«      ‚|j                  «       j	                  «       }|j                  \  }}||z  }| j
                  }||z   dz
  |z  }|dz   dz  }t        j                  |d¬«      }	t        j                  ||j                  ¬«      }
t        |	||
|| j                  ||«       |	|
fS )z4b quantize fp32/fp16 weighté   z9Current bnb4 block quantization only supports 2D tensors!r   Úuint8)Údtype)r   ÚshapeÚ
ValueErrorÚ	transposeÚcopyr   ÚnpÚzerosr*   r   r   )r   r&   Ú
fpweight_tÚrowsÚcolsÚnumelr   Ú
num_blocksÚquantized_numelÚpackedÚabsmaxs              r   Úbnb4_block_quantz$MatMulBnb4Quantizer.bnb4_block_quant6   sÄ   € ô ˆx�~‰~Ó !Ò#ÜÐXÓYÐYð ×'Ñ'Ó)×.Ñ.Ó0ˆ
à—^‘^‰
ˆˆdØ�t‘ˆØ—_‘_ˆ
Ø˜jÑ(¨1Ñ,°Ñ;ˆ
Ø  1™9¨Ñ*ˆä—‘˜/°Ô9ˆÜ—‘˜*¨H¯N©NÔ;ˆä˜V Z°¸ÀTÇ_Á_ÐVZÐ\`Ôaà˜ÐÐr   ÚnodeÚgraph_stackc                 óJ  — |j                   dk7  r|S t        j                  d|j                  › d�«       |j                  | j                  v r%t        j                  d|j                  › d�«       |S |j
                  d   }t        j                  ||«      \  }}|€t        j                  d«       |S t        j                  j                  |«      }t        |j                  «      dk7  rt        j                  d	«       |S | j                  |«      \  }}t        j                  j                  |«      }	|j                  d
z   |	_        |j
                  D ].  }
|
j                  |k(  sŒ|j
                  j                  |
«        n t        j                  j                  |«      }|j                  dz   |_        |j                   j#                  |	|g«       i }|j                  \  }}||d<   ||d<   | j$                  |d<   | j&                  |d<   t        j(                  j*                  	 d|j
                  d   |	j                  |j                  g|j,                  d   g|j                  r|j                  d
z   ndddœ|¤Ž}t        j                  d|j                  › d�«       |S )zdIf the node is MatMul with fp32 const weight, quantize the weight with int4, and return the new nodeÚMatMulzstart to quantize z ...zexclude to quantize z$ as specified by nodes_to_exclude...r   z2MatMul doesn't have const weight. Skip to quantizer(   z)MatMul weight is not 2D. Skip to quantizeÚ_Bnb4Ú_absmaxÚKÚNr   r   r   Ú úcom.microsoft)ÚinputsÚoutputsr!   Údomainzcomplete quantization of )Ú
MatMulBnb4)Úop_typeÚloggerÚdebugr!   r   Úinputr   Ú%_MatMulBnb4Quantizer__get_initializerÚonnxÚnumpy_helperÚto_arrayr   r+   r9   Ú
from_arrayÚremover    Úextendr   r   ÚhelperÚ	make_nodeÚoutput)r   r:   r;   ÚinputBÚBÚBs_graphÚB_arrayr7   r8   ÚB_quantrK   Úabsmax_tensorÚkwargsr2   r3   Úmatmul_bnb4_nodes                   r   Ú_bnb4_matmul_node_weightz,MatMulBnb4Quantizer._bnb4_matmul_node_weightL   sN  € ð �<‰<˜8Ò#ØˆKä�‰Ð)¨$¯)©)¨°DÐ9Ô:Ø�9‰9˜×-Ñ-Ñ-Ü�L‰LÐ/°·	±	¨{Ð:^Ð_Ô`ØˆKà—‘˜A‘ˆÜ)×;Ñ;¸FÀKÓP‰ˆˆ8Øˆ9Ü�L‰LÐMÔNØˆKä×#Ñ#×,Ñ,¨QÓ/ˆÜˆw�}‰}Ó Ò"Ü�L‰LÐDÔEØˆKà×.Ñ.¨wÓ7‰ˆ�Ü×#Ñ#×.Ñ.¨vÓ6ˆØ—v‘v Ñ'ˆŒØ—^‘^ò 	ˆEØ�z‰z˜VÓ#Ø—‘×%Ñ% eÔ,Ùð	ô
 ×)Ñ)×4Ñ4°VÓ<ˆØŸV™V iÑ/ˆÔà×Ñ×#Ñ# W¨mÐ$<Ô=àˆØ—]‘]‰
ˆˆdØˆˆs‰Øˆˆs‰Ø#Ÿ™ˆˆ|ÑØ#Ÿ™ˆˆ|ÑäŸ;™;×0Ñ0Øð
à—J‘J˜q‘M 7§<¡<°×1CÑ1CÐDØ—[‘[ ‘^Ð$Ø(,¯	ª	�—‘˜WÒ$°rØ"ñ
ð ñ
Ðô 	�‰Ð0°·±°¸4Ð@ÔAàÐr   c                 ó~  — g }|d   }|j                   D �]ä  }|j                  D �cg c]R  }|j                  t        j                  j
                  k(  s'|j                  t        j                  j                  k(  r|‘ŒT }}|�rVi }|j                  D ]ù  }|j                  t        j                  j
                  k(  r9|j                  |j                  «       |j                  | j                  |«      i}n†|j                  t        j                  j                  k(  rTg }	|j                  D ]4  }
|j                  |
«       |	j                  | j                  |«      g«       Œ6 |j                  |	i}nt        |«      }|j                  |«       Œû t        j                  j                   |j"                  |j$                  |j&                  fd|j                  i|¤Ž}|j                  | j)                  ||«      «       �Œç |j+                  d«       |j                   j                  |«       |j-                  «        |S c c}w )Nr   r!   r:   )r:   Ú	attributeÚtyperM   ÚAttributeProtoÚGRAPHÚGRAPHSÚappendÚgr!   Ú_process_subgraphÚgraphsrR   r
   ÚupdaterS   rT   rH   rK   rU   r^   Ú
ClearFieldÚpop)r   r;   Ú	new_nodesr#   r:   ÚattrÚgraph_attrsr\   ÚkvÚvalueÚsubgraphs              r   rg   z%MatMulBnb4Quantizer._process_subgraphƒ   sç  € Øˆ	Ø˜B‘ˆà—J‘Jó 	OˆDð !ŸN™NöàØ—9‘9¤× 3Ñ 3× 9Ñ 9Ò9¸T¿Y¹YÌ$×J]ÑJ]×JdÑJdÒ=dò ðˆKð ò
 Ø�Ø ŸN™Nò &�DØ—y‘y¤D×$7Ñ$7×$=Ñ$=Ò=à#×*Ñ*¨4¯6©6Ô2Ø"Ÿi™i¨×)?Ñ)?ÀÓ)LÐM™ØŸ™¤d×&9Ñ&9×&@Ñ&@Ò@Ø "˜Ø(,¯©ò P˜Hà'×.Ñ.¨xÔ8Ø!ŸL™L¨$×*@Ñ*@ÀÓ*MÐ)NÕOðPð #Ÿi™i¨Ð/™ä/°Ó5˜Ø—M‘M "Õ%ð&ô —{‘{×,Ñ,Ø—L‘L $§*¡*¨d¯k©kñØ@DÇ	Á	ðØMSñ�ð ×Ñ˜T×:Ñ:¸4ÀÓMÖNð7	Oð: 	×Ñ˜Ô Ø�
‰
×Ñ˜)Ô$Ø�‰ÔØˆùò?s   ¦AH:c                 ó^  — | j                   j                  «       g}| j                   j                  «       }d}|D ]  }|j                  dk(  sŒd}Œ |s0|j	                  t
        j                  j                  dd«      g«       | j                  |«       | j                   j                  «        y )NFrC   Tr   )
r   r#   Úopset_importrF   rR   rM   rS   Úmake_opsetidrg   Úclean_initializers)r   r;   rs   Úhas_ms_domainÚopsets        r   ÚprocesszMatMulBnb4Quantizer.process©   s“   € à—z‘z×'Ñ'Ó)Ð*ˆØ—z‘z×.Ñ.Ó0ˆàˆØ!ò 	%ˆEØ�|‰|˜Ó.Ø $‘ð	%ñ Ø×Ñ¤§¡×!9Ñ!9¸/È1Ó!MÐ NÔOà×Ñ˜{Ô+Ø�
‰
×%Ñ%Õ'r   r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r   Úintr   ÚstaticmethodÚlistr   Útupler   rL   ÚnptÚ	ArrayLiker/   Úndarrayr9   r   r^   rg   rx   © r   r   r   r      s²   „ ÙWð €Cð €Cñ6˜jð 6°cð 6Àsó 6ð ð¨D°Ñ,<ð ÀÀ{ÐT^ÐG^ÑA_ò ó ðð ¨¯©ð  ¸2¿:¹:ó  ð,5 ¨Yð 5 ÀTÈ*ÑEUð 5 ÐZcó 5 ðn$¨T°*Ñ-=ó $óL(r   r   c                  ó´  — t        j                  d¬«      } | j                  ddd¬«       | j                  ddd¬«       | j                  d	d
dt        j                  t        j
                  gd¬«       | j                  dd
dd¬«       | j                  ddd
d¬«       | j                  d
¬«       | j                  ddt        d
g d¬«       | j                  «       S )Na  Blockwise FP4/NF4 quantization for MatMul 2D weight matrices.

A weight matrix is partitioned into blocks, where each block is a contiguous
subset inside the flattened transposed weight matrix. Each block is quantized
into a set of 4b integers with an absolute value scaling factor.
)Údescriptionz--input_modelTzPath to the input model file)ÚrequiredÚhelpz--output_modelzPath to the output model filez--quant_typeFr   z&Quantization data type. 0: FP4, 1: NF4)r‡   ÚdefaultÚchoicesrˆ   z--block_sizeé@   zVBlock size for blockwise quantization. Note: bnb.nn.Linear4bit only uses block_size=64)r‡   r‰   rˆ   z-vz	--verboseÚ
store_true)r‡   Úaction)Úverbosez--nodes_to_excludeú+zBSpecify the nodes to be excluded from quantization with node names)Únargsra   r‡   r‰   rˆ   )	ÚargparseÚArgumentParserÚadd_argumentr   r   r   Úset_defaultsÚstrÚ
parse_args)Úparsers    r   r–   r–   ¹   sû   € Ü×$Ñ$ðô€Fð ×Ñ˜°$Ð=[ÐÔ\Ø
×ÑÐ(°4Ð>]ÐÔ^Ø
×ÑØØØÜ$×(Ñ(Ô*=×*AÑ*AÐBØ5ð ô ð ×ÑØØØØeð	 ô ð ×Ñ˜˜k°EÀ,ÐÔOØ
×Ñ ÐÔ&Ø
×ÑØØÜØØØQð ô ð ×ÑÓÐr   Ú__main__zfile z already exists)r   T)-r‘   ÚloggingÚosÚnumpyr/   Únumpy.typingÚtypingr�   rM   Úonnx.onnx_pbr   r   r   r   Úonnxruntime.capi._pybind_stater   Ú
onnx_modelr	   Úquant_utilsr
   Ú	getLoggerry   rI   r   r–   ÚargsrŽ   ÚsetLevelÚDEBUGÚinput_modelÚinput_model_pathÚoutput_modelÚoutput_model_pathÚpathÚexistsÚerrorÚ	ExceptionÚloadr   r   r   r   Úquantrx   Úsave_model_to_filer„   r   r   ú<module>r±      s  ðó Û Û 	ã Ý Û ß GÓ Gå ?å !Ý +à	ˆ×	Ñ	˜8Ó	$€÷^(ñ ^(òB$ðN ˆzÒÙ‹<€DØ‡|‚|Ø�‰˜Ÿ™Ô&à×'Ñ'ÐØ×)Ñ)Ðà	‡w�w‡~�~Ð'Ô(Ø�‰�uÐ.Ð/¨Ð?Ô@Ù˜%Ð 1Ð2°/ÐBÓCÐCàˆD�I‰IÐ&Ó'€EÙ  t§¡¸¿¹ÐZ^×ZoÑZoÔp€EØ	‡M�M„OØ	‡K�K×"Ñ"Ð#4°dÕ;ð r   