Ë
    îÍ:jK  ã                   ó¬   — d dl mZmZ ddlmZ erddlmZ ddlmZm	Z	m
Z
mZmZ ddlmZ  e
«       rd dlZ ej                   e«      ZdZ G d	„ d
e«      Zy)é    )ÚTYPE_CHECKINGÚOptionalé   )ÚHfQuantizeré   )ÚPreTrainedModel)Úis_accelerate_availableÚis_kernels_availableÚis_torch_availableÚis_triton_availableÚlogging)Úget_module_from_nameNc                   ó   ‡ — e Zd ZdZdZdZdgZˆ fd„Zd„ Zd„ Z	d#d	„Z
d
ddedefd„Zd
ddddeddfd„Zd$d„Zd
ddee   dee   fd„Z	 d%d
ddeee      fd„Zdee   dedee   fd„Zd„ Zd„ Zdedefd„Zd&defd „Zd%d!„Zedefd"„«       Zˆ xZS )'ÚMxfp4HfQuantizerz/
    FP4 quantization using fbgemm kernels
    TFÚ
acceleratec                 óB   •— t        ‰| �  |fi |¤Ž || _        d | _        y ©N)ÚsuperÚ__init__Úquantization_configÚtriton_kernels_hub)Úselfr   ÚkwargsÚ	__class__s      €ú|/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/transformers/quantizers/quantizer_mxfp4.pyr   zMxfp4HfQuantizer.__init__1   s&   ø€ Ü‰ÑÐ,Ñ7°Ò7Ø#6ˆÔ Ø"&ˆÕó    c                 ó¢   — | j                   € 	 ddlm}  |d«      | _         | j                   S | j                   S # t        $ r t        d«      ‚w xY w)z3Lazy import and initialize kernels only when neededr   )Ú
get_kernelz kernels-community/triton_kernelsz2kernels package is required for MXFP4 quantization)r   Úkernelsr   ÚImportError)r   r   s     r   Ú_lazy_import_kernelsz%Mxfp4HfQuantizer._lazy_import_kernels6   s]   € à×"Ñ"Ð*ðXÝ.á*4Ð5WÓ*X�Ô'ð ×&Ñ&Ð&ˆt×&Ñ&Ð&øô ò XÜ!Ð"VÓWÐWðXús	   Ž9 ¹Ac                 óx  — t        «       st        d«      ‚| j                  j                  ry t        j
                  j                  «       s\t        j                  j                  «       s>| j                  r't        j                  d«       d| j                  _        y t        d«      ‚t        «       st        d«      ‚t        j                  j                  «       rd}t        d«      xr
 t        «       }n:t        j
                  j                  «       }|dk\  }t        d«      xr
 t        «       }| j                  rR|s't        j                  d	«       d| j                  _        y |sAt        j                  d
«       d| j                  _        y |st!        d«      ‚|st!        d«      ‚| j                  s| j#                  «        |j%                  d«      }|€t        j                  d«       y |�N| j                  sAt'        |t(        «      r0d|j+                  «       v sd|j+                  «       v rt!        d«      ‚y y y y )NzqUsing mxfp4 quantization requires torchPlease install the latest version of torch ( pip install --upgrade torch )z^Using MXFP4 quantized models requires a GPU, we will default to dequantizing the model to bf16Tz-Quantizing a model using MXFP4 requires a GPUz9Using mxfp4 requires Accelerate: `pip install accelerate`z3.5.0)é   é   z3.4.0uÑ   MXFP4 quantization is only supported on GPUs with compute capability >= 7.5 (e.g T4, A100, L4, H100, or B200) or XPUs (e.g IntelÂ® Data Center GPU Max Series) We will default to dequantizing the model to bf16.z¨MXFP4 quantization requires Triton and kernels installed: CUDA requires Triton >= 3.4.0, XPU requires Triton >= 3.5.0, we will default to dequantizing the model to bf16uŸ   MXFP4 quantization is only supported on GPUs with compute capability >= 7.5 (e.g T4, A100, L4, H100, or B200) or XPUs (e.g IntelÂ® Data Center GPU Max Series) zuMXFP4 quantization requires Triton and kernels installed: CUDA requires Triton >= 3.4.0, XPU requires Triton >= 3.5.0Ú
device_mapzÞYou have loaded an FP4 model on CPU and have a CUDA/XPU device available, make sure to set your model on a GPU/XPU device in order to run your model. To remove this warning, pass device_map = 'cuda' or device_map = 'xpu'. ÚcpuÚdiskzòYou are attempting to load an FP4 model with a device_map that contains a CPU or disk device.This is not supported when the model is quantized on the fly. Please use a quantized checkpoint or remove the CPU or disk device from the device_map.)r   r    r   Ú
dequantizeÚtorchÚcudaÚis_availableÚxpuÚpre_quantizedÚloggerÚwarning_onceÚRuntimeErrorr	   r   r
   Úget_device_capabilityÚ
ValueErrorr!   ÚgetÚ
isinstanceÚdictÚvalues)r   Úargsr   Úgpu_is_supportedÚkernels_availableÚcompute_capabilityr%   s          r   Úvalidate_environmentz%Mxfp4HfQuantizer.validate_environmentA   s#  € Ü!Ô#Üð]óð ð
 ×#Ñ#×.Ò.Øä—
‘
×'Ñ'Ô)¬U¯Y©Y×-CÑ-CÔ-EØ×!Ò!Ü×#Ñ#Øtôð 7;�×(Ñ(Ô3Øä"Ð#RÓSÐSä&Ô(ÜÐYÓZÐZä�9‰9×!Ñ!Ô#Ø#ÐÜ 3°GÓ <Ò WÔAUÓAWÑä!&§¡×!AÑ!AÓ!CÐØ1°VÑ;ÐÜ 3°GÓ <Ò WÔAUÓAWÐà×Òá#Ü×#Ñ#ðIôð 7;�×(Ñ(Ô3Øá$Ü×#Ñ#ð ôð 7;�×(Ñ(Ô3ØÙ!äð róð ñ #äð Hóð ð ×!Ò!Ø×%Ñ%Ô'à—Z‘Z Ó-ˆ
ØÐÜ×ÑðVõð Ð#à×&Ò&Ü˜z¬4Ô0Ø˜j×/Ñ/Ó1Ñ1°V¸z×?PÑ?PÓ?RÑ5Rä ðnóð ð 6Sð 1ð 'ð $r   Úreturnc                 óV   — |€&t         j                  }t        j                  d|«       |S )NzôOverriding dtype=%s with `dtype=torch.bfloat16` due to requirements of `fbgemm-gpu` to enable model loading in fp4. Pass your own dtype to specify the dtype of the remaining non-linear layers or pass dtype=torch.bfloat16 to remove this warning.)r)   Úbfloat16r.   Úinfo)r   Údtypes     r   Úupdate_dtypezMxfp4HfQuantizer.update_dtype�   s.   € Øˆ=Ü—N‘NˆEÜ�K‰Kð@ð ôð ˆr   Úmodelr   Ú
param_namec                 ó  — ddl m} ddlm} | j                  j
                  r%d|v sd|v rt        ||d t        d«        «      \  }}nt        ||«      \  }}t        ||«      s"t        ||«      r| j                  j
                  r|dv ryy	y)
Nr   ©ÚMxfp4GptOssExperts©ÚGptOssExpertsÚblocksÚscalesÚ_blocks)Údown_proj_biasÚgate_up_proj_biasFT)	ÚintegrationsrF   Úmodels.gpt_oss.modeling_gpt_ossrH   r   r(   r   Úlenr4   )r   rB   rC   r   rF   rH   ÚmoduleÚtensor_names           r   Úparam_needs_quantizationz)Mxfp4HfQuantizer.param_needs_quantizationœ   sŽ   € Ý5ÝCð ×#Ñ#×.Ò.°HÀ
Ñ4JÈhÐZdÑNdÜ"6°u¸jÐIZÌCÐPYËNÈ?Ð>[Ó"\ÑˆF‘Kä"6°u¸jÓ"IÑˆF�KÜ�fÐ0Ô1Ü�v˜}Ô-°$×2JÑ2J×2UÒ2UàÐEÑEØØØr   Úparam_valueztorch.TensorÚtarget_deviceztorch.devicec                 ó   — ddl m}m}m}m}	m}
 ddlm} | j                  sü| j                  «       }t        ||«      \  }}t        j                  |«      5  t        ||«      r² |	||«      \  }}|j                  j                  |j                  j                   |j                  j"                  }}} |
|||«      \  }}d|v rdnd}t%        |||«       t%        ||› d� || | |«       ¬«      ¬«      «       t'        ||› d	�«       t'        ||› d
�«       d d d «       y |j)                  d«      }|j)                  d«      }|j)                  d«      }|j)                  d«      }|j)                  d«      }d|v sd|v r3| j*                  j                  rt        ||d t-        d	«        «      \  }}nt        ||«      \  }}||||||dœ}t        ||«      s"t        ||«      rf| j*                  j                  rO| j*                  j                  r|d t-        d	«        } ||||||fi |¤Ž y  |||||| j                  «       fi |¤Ž y y y # 1 sw Y   y xY w)Nr   )rF   r(   Úload_and_swizzle_mxfp4Úquantize_to_mxfp4Úswizzle_mxfp4rG   Úgate_up_projÚ	down_projÚ_precision_config)Úrhs_data)Úweight_scaleÚflex_ctxrK   Ú_scalesÚempty_paramÚcasting_dtypeÚto_contiguousÚrankÚdevice_meshrI   rJ   )ra   rb   rc   rd   re   rB   )rN   rF   r(   rW   rX   rY   rO   rH   r-   r!   r   r)   Údevicer4   Ú
matmul_ogsÚPrecisionConfigÚFlexCtxÚ
InFlexDataÚsetattrÚdelattrr3   r   rP   )r   rB   rT   rC   rU   r   rF   r(   rW   rX   rY   rH   r   rQ   Ú_Útriton_weight_tensorr^   rh   ri   rj   Úprojra   rb   rc   rd   re   Úshard_kwargsÚdq_param_names                               r   Úcreate_quantized_paramz'Mxfp4HfQuantizer.create_quantized_param­   sz  € ÷	
õ 	
õ 	Dà×!Ò!Ø!%×!:Ñ!:Ó!<ÐÜ,¨U°JÓ?‰IˆF�AÜ—‘˜mÓ,ñ 6Ü˜fÐ&8Ô9Ù9JÈ;ÐXjÓ9kÑ6Ð(¨,à*×5Ñ5×EÑEØ*×5Ñ5×=Ñ=Ø*×5Ñ5×@Ñ@ð /9 W�Oñ
 :GØ,¨lÐ<Nó:Ñ6Ð(¨,ð .<¸zÑ-I™>È{�DÜ˜F DÐ*>Ô?ÜØØ˜&Ð 1Ð2Ù'°\ÉGÑ]gÓ]iÔLjÔkôô ˜F t f¨GÐ$4Ô5Ü˜F t f¨GÐ$4Ô5÷+6ð 6ð4 !Ÿ*™* ]Ó3ˆKØ"ŸJ™J Ó7ˆMØ"ŸJ™J Ó7ˆMØ—:‘:˜fÓ%ˆDØ Ÿ*™* ]Ó3ˆKØ˜JÑ&¨(°jÑ*@Àd×F^ÑF^×FiÒFiä0°¸
ÐCTÄcÈ)ÃnÀ_Ð8UÓV‘	�™ä0°¸
ÓC‘	�˜ð  +Ø!.Ø!.ØØ*ØñˆLô ˜&Ð"4Ô5Ü˜6 =Ô1°d×6NÑ6N×6YÒ6Yà×+Ñ+×6Ò6ð %/Ð/@´#°i³.°Ð$A�MÙ˜v z°;ÀÈ}ÑmÐ`lÓmá*ØØ"Ø#Ø%Ø×1Ñ1Ó3ñð 'óð 7ZÐ1÷_6ð 6ús   ÁB?IÉIc                 óF  — | j                   j                  r| j                  |«       t        j                  j                  «       rt        j                  j                  «        y t        j                  j                  «       rt        j                  j                  «        y y r   )r   r(   Úremove_quantization_configr)   r*   r+   Úempty_cacher,   )r   rB   r   s      r   Ú#_process_model_after_weight_loadingz4Mxfp4HfQuantizer._process_model_after_weight_loading  sd   € à×#Ñ#×.Ò.Ø×+Ñ+¨EÔ2ä�:‰:×"Ñ"Ô$Ü�J‰J×"Ñ"Õ$Ü�Y‰Y×#Ñ#Ô%Ü�I‰I×!Ñ!Õ#ð &r   Úexpected_keysÚcheckpoint_keysc                 óœ  — g }|D �]C  }|j                  d«      r8|d t        d«        }|j                  |dz   «       |j                  |dz   «       ŒM|j                  d«      r8|d t        d«        }|j                  |dz   «       |j                  |dz   «       Œ–| j                  s‘|j                  d	«      r$|d t        d«        }|j                  |dz   «       Œ×|j                  d
«      r%|d t        d«        }|j                  |dz   «       �Œ|j                  d«      r�Œ |j                  |«       �Œ3|j                  |«       �ŒF |S )Nz.mlp.experts.gate_up_projrZ   Úgate_up_proj_blocksÚgate_up_proj_scalesz.mlp.experts.down_projr[   Údown_proj_blocksÚdown_proj_scalesz.mlp.experts.down_proj_blocksz .mlp.experts.gate_up_proj_blocksrJ   )ÚendswithrP   Úappendr-   )r   rB   rw   rx   Únew_expected_keysÚkeyÚbases          r   Úupdate_expected_keysz%Mxfp4HfQuantizer.update_expected_keys  sS  € àÐØ ó 	.ˆCØ�|‰|Ð7Ô8ØÐ1œc .Ó1Ð1Ð2�Ø!×(Ñ(¨Ð0EÑ)EÔFØ!×(Ñ(¨Ð0EÑ)EÕFØ—‘Ð6Ô7ØÐ.œc +Ó.Ð.Ð/�Ø!×(Ñ(¨Ð0BÑ)BÔCØ!×(Ñ(¨Ð0BÑ)BÕCØ×'Ò'à—<‘<Ð ?Ô@ØÐ9¤#Ð&8Ó"9Ð!9Ð:�DØ%×,Ñ,¨T°KÑ-?Õ@Ø—\‘\Ð"DÔEØÐ<¤#Ð&;Ó"<Ð!<Ð=�DØ%×,Ñ,¨T°NÑ-BÖCØ—\‘\ (Ô+áà%×,Ñ,¨SÖ1à!×(Ñ(¨Ö-ð/	.ð0 !Ð r   Úkeep_in_fp32_modulesc                 ój  — ddl m} | j                  || j                  j                  |«      | _        |j                  dd«      }|r&t        j                  d«       d| j                  _        |j                  } ||| j                  | j                  |¬«      }| j                  |j                  _        y )Nr   )Úreplace_with_mxfp4_linearÚuse_kernelsFzžYou are using full precision kernels, we will dequantize the model to bf16. To use the quantized model with quantization kernels, please set use_kernels=FalseT)Úmodules_to_not_convertr   Úconfig)
rN   r†   Úget_modules_to_not_convertr   rˆ   r3   r.   r/   r(   r‰   )r   rB   r„   r   r†   r‡   r‰   s          r   Ú$_process_model_before_weight_loadingz5Mxfp4HfQuantizer._process_model_before_weight_loading(  s¦   € õ 	=à&*×&EÑ&EØ�4×+Ñ+×BÑBÐDXó'
ˆÔ#ð —j‘j °Ó6ˆáÜ×Ñðeôð 37ˆD×$Ñ$Ô/à—‘ˆÙ)ØØ#'×#>Ñ#>Ø $× 8Ñ 8Øô	
ˆð ,0×+CÑ+Cˆ�‰Õ(r   Úmissing_keysÚprefixc                 ó$  — ddl m} g }|j                  «       D ]\  \  }}t        ||«      sŒ|D ]E  }||v s
||› d|› �v sŒ|j	                  d«      rŒ#|j	                  d«      rŒ5|j                  |«       ŒG Œ^ |D �	cg c]	  }	|	|vsŒ|	‘Œ c}	S c c}	w )Nr   rE   ú.z.weightz.bias)rN   rF   Únamed_modulesr4   r~   r   )
r   rB   rŒ   r�   rF   Únot_missing_keysÚnamerQ   ÚmissingÚks
             r   Úupdate_missing_keysz$Mxfp4HfQuantizer.update_missing_keysG  s¢   € Ý5àÐØ!×/Ñ/Ó1ò 	9‰LˆD�&Ü˜&Ð"4Õ5Ø+ò 9�Gà ™¨D°v°h¸aÀ¸yÐ4IÒ,IØ '× 0Ñ 0°Õ ;Ø '× 0Ñ 0°Õ 9à(×/Ñ/°Õ8ñ9ð	9ð (ÖE�a¨1Ð4DÒ+D’ÒEÐEùÒEs   Á<	BÂBc                 ó�   — d|j                   j                  v r-t        |dd «      � |j                  j	                  dddddœ«       |S )NÚGptOssConfigÚbase_model_tp_planÚgrouped_gemm©z(layers.*.mlp.experts.gate_up_proj_blocksz(layers.*.mlp.experts.gate_up_proj_scalesz%layers.*.mlp.experts.down_proj_blocksz%layers.*.mlp.experts.down_proj_scales)r   Ú__name__Úgetattrr˜   Úupdate©r   r‰   s     r   Úupdate_tp_planzMxfp4HfQuantizer.update_tp_planV  óR   € Ø˜V×-Ñ-×6Ñ6Ñ6Ü�vÐ3°TÓ:ÐFØ×)Ñ)×0Ñ0àDRØDRØAOØAOñ	ôð ˆr   c                 ó�   — d|j                   j                  v r-t        |dd «      � |j                  j	                  dddddœ«       |S )Nr—   Úbase_model_ep_planr™   rš   )r   r›   rœ   r¢   r�   rž   s     r   Úupdate_ep_planzMxfp4HfQuantizer.update_ep_planc  r    r   c                 ó2  — | j                   j                  r.d|v r|j                  dd«      S d|v r|j                  dd«      S |S | j                  sF|j	                  d«      r|j                  dd«      S |j	                  d«      r|j                  dd«      S |S )NrK   Ú r`   rZ   rz   r[   r|   )r   r(   Úreplacer-   r~   )r   rC   s     r   Úget_param_namezMxfp4HfQuantizer.get_param_namep  s¤   € Ø×#Ñ#×.Ò.Ø˜JÑ&Ø!×)Ñ)¨)°RÓ8Ð8Ø˜jÑ(Ø!×)Ñ)¨)°RÓ8Ð8ð Ðð ×#Ò#Ø×"Ñ" >Ô2Ø!×)Ñ)¨.Ð:OÓPÐPØ×"Ñ" ;Ô/Ø!×)Ñ)¨+Ð7IÓJÐJØÐr   Úsafe_serializationc                 ól  — ddl m} |j                  «       }|j                  «       D �]  \  }}t	        ||«      sŒt        |d«      sŒ!t        |d«      sŒ.|j                  j                  j                  j                  |j                  j                  j                  «      j                  dd«      j                  dddd	«      ||› d
�<   |j                  j                  j                  j                  j                  |j                  j                  j                  j                  «      j                  dd«      ||› d�<   |j                  j                  j                  j                  |j                  j                  j                  «      j                  dd«      j                  dddd«      ||› d�<   |j                   j                  j                  j                  j                  |j                   j                  j                  j                  «      j                  dd«      ||› d�<   �Œ i }||fS )Nr   rE   rZ   r[   éÿÿÿÿéþÿÿÿé    éZ   é   z.gate_up_proj_blocksz.gate_up_proj_scalesi@  z.down_proj_blocksz.down_proj_scales)rN   rF   Ú
state_dictr�   r4   ÚhasattrrZ   ÚstorageÚlayoutÚunswizzle_dataÚdataÚ	transposeÚreshapeÚgate_up_proj_precision_configr^   r[   Údown_proj_precision_config)r   rB   r¨   rF   r¯   r’   rQ   Úmetadatas           r   Úget_state_dict_and_metadataz,Mxfp4HfQuantizer.get_state_dict_and_metadata}  sæ  € Ý5à×%Ñ%Ó'ˆ
à!×/Ñ/Ó1ó 	‰LˆD�&ä˜6Ð#5Õ6Ü˜F NÕ3Ü˜F KÕ0ð ×'Ñ'×/Ñ/×6Ñ6×EÑEÀf×FYÑFY×FaÑFa×FfÑFfÓgß‘Y˜r 2Ó&ß‘W˜R  R¨Ó,ð ˜d˜VÐ#7Ð8Ñ9ð ×8Ñ8×EÑE×MÑM×TÑT×cÑcØ×<Ñ<×IÑI×QÑQ×VÑVóç‘i  BÓ'ð ˜d˜VÐ#7Ð8Ñ9ð ×$Ñ$×,Ñ,×3Ñ3×BÑBÀ6×CSÑCS×C[ÑC[×C`ÑC`Óaß‘Y˜r 2Ó&ß‘W˜R  r¨2Ó.ð ˜d˜VÐ#4Ð5Ñ6ð ×5Ñ5×BÑB×JÑJ×QÑQ×`Ñ`Ø×9Ñ9×FÑF×NÑN×SÑSóç‘i  BÓ'ð ˜d˜VÐ#4Ð5Ó6ð+	ð6 ˆØ˜8Ð#Ð#r   c                  ó   — y)NT© )r   r¨   s     r   Úis_serializablez Mxfp4HfQuantizer.is_serializable   s   € Ør   c                 ó.   — t         j                  d«       y)Nz©MXFP4 quantization don't support training, please consider dequantizing the model first by passing quantization_config=Mxfp4Config(dequantize=True) to .from_pretrained()F)r.   r/   )r   s    r   Úis_trainablezMxfp4HfQuantizer.is_trainable£  s   € ä×Ñð xô	
ð r   )r@   útorch.dtyper<   rÀ   )rB   r   r   )F)r›   Ú
__module__Ú__qualname__Ú__doc__Ú requires_parameters_quantizationÚrequires_calibrationÚrequired_packagesr   r!   r;   rA   ÚstrÚboolrS   rr   rv   Úlistrƒ   r   r‹   r•   rŸ   r£   r§   rº   r½   Úpropertyr¿   Ú__classcell__)r   s   @r   r   r   '   sI  ø„ ñð (,Ð$Ø Ðà%˜Ðô'ò
	'òMó^
ðÐ.?ð ÈSð Ð_có ð"Rà ðRð $ðRð ð	Rð
 &óRóh$ð!Ð*;ð !ÈDÐQTÉIð !ÐhlÐmpÑhqó !ð@ 59ñDà ðDð ' t¨C¡yÑ1óDð>F°t¸C±yð FÈ#ð FÐRVÐWZÑR[ó Fòòð¨ð °ó ñ!$ÀTó !$óFð ð˜dò ó ôr   r   )Útypingr   r   r‚   r   Úmodeling_utilsr   Úutilsr	   r
   r   r   r   Úquantizers_utilsr   r)   Ú
get_loggerr›   r.   r   r   r¼   r   r   ú<module>rÑ      sU   ð÷ +å ñ Ý0÷õ õ 3ñ ÔÛà	ˆ×	Ñ	˜HÓ	%€ØÐ ôA�{õ Ar   