Ë
    HêñiÙ3  ã                   óÀ   — d dl mZ ddlmZ erddlmZ 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dd
lmZ  ej&                  e«      Zd	Z G d„ de«      Zy	)é    )ÚTYPE_CHECKINGé   )ÚHfQuantizeré   )ÚPreTrainedModel)ÚMxfp4Config)Úis_accelerate_availableÚis_kernels_availableÚis_torch_availableÚis_triton_availableÚlogging)Úget_module_from_nameN)ÚWeightConverterc                   ó¨   ‡ — e Zd ZU dZdZded<   ˆ fd„Zd„ Zd„ Zdd	d
e	de
fd„Zdd„Z	 ddd	de
fd„Zd„ Zd„ Zd„ Zd„ Zede
fd„«       Zd„ Zd„ Zˆ xZS )ÚMxfp4HfQuantizerz/
    FP4 quantization using fbgemm kernels
    Fr   Úquantization_configc                 ó4   •— t        ‰| �  |fi |¤Ž d | _        y ©N)ÚsuperÚ__init__Útriton_kernels_hub)Úselfr   ÚkwargsÚ	__class__s      €úi/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/transformers/quantizers/quantizer_mxfp4.pyr   zMxfp4HfQuantizer.__init__2   s   ø€ Ü‰ÑÐ,Ñ7°Ò7Ø"&ˆÕó    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/gpt-oss-triton-kernelsz2kernels package is required for MXFP4 quantization)r   Úintegrations.hub_kernelsr   ÚImportError)r   r   s     r   Ú_lazy_import_kernelsz%Mxfp4HfQuantizer._lazy_import_kernels6   s]   € à×"Ñ"Ð*ðXÝAá*4Ð5_Ó*`�Ô'ð ×&Ñ&Ð&ˆt×&Ñ&Ð&øô ò XÜ!Ð"VÓWÐWðXús	   Ž9 ¹Ac                 ó>  — t        «       st        d«      ‚| j                  j                  ry t	        «       st        d«      ‚t
        j                  j                  «       xs t        j                  d«      }|j                  dvrF| j                  r+t        j                  d|› d�«       d| j                  _        y t        d|› d	�«      ‚t
        j                  j                  «       rd}t!        d
«      }t#        «       }n„t
        j$                  j                  «       r9t
        j$                  j'                  «       }|dk\  }t!        d«      }t#        «       }n-|j                  dk(  rd}t!        d
«      }t#        «       }nd}d}d}| j                  r{|s't        j                  d«       d| j                  _        y |s't        j                  d«       d| j                  _        y |sNt        j                  d«       d| j                  _        y |st)        d«      ‚|st)        d«      ‚|st)        d«      ‚| j                  s| j+                  «        |j-                  d«      }|�<t/        |t0        «      r+| j                  sd|j3                  «       v rt)        d«      ‚y y y y )NzqUsing mxfp4 quantization requires torchPlease install the latest version of torch ( pip install --upgrade torch )z9Using mxfp4 requires Accelerate: `pip install accelerate`Úcpu)ÚcudaÚxpur#   zGUsing MXFP4 quantized models requires model on cuda/xpu/cpu, but found zj, we will default to dequantizing the model to bf16. To use mxfp4, please disable the current accelerator.TzIQuantizing a model using MXFP4 requires model on cuda/xpu/cpu, but found z7. To use mxfp4, please disable the current accelerator.z3.5.0)é   é   z3.4.0FuÒ   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: CUDA requires Triton >= 3.4.0, XPU/CPU requires Triton >= 3.5.0. Please install triton: `pip install triton`. We will default to dequantizing the model to bf16.z„MXFP4 quantization requires the `kernels` package: `pip install kernels>=0.12.0`. We will default to dequantizing the model to bf16.u¥   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) or CPUz�MXFP4 quantization requires Triton: CUDA requires Triton >= 3.4.0, XPU/CPU requires Triton >= 3.5.0. Please install triton: `pip install triton`zPMXFP4 quantization requires the `kernels` package: `pip install kernels>=0.12.0`Ú
device_mapÚdiskzäYou are attempting to load an FP4 model with a device_map that contains a disk device.This is not supported when the model is quantized on the fly. Please use a quantized checkpoint or remove the disk device from the device_map.)r   r    r   Ú
dequantizer	   ÚtorchÚacceleratorÚcurrent_acceleratorÚdeviceÚtypeÚpre_quantizedÚloggerÚwarning_onceÚRuntimeErrorr%   Úis_availabler   r
   r$   Úget_device_capabilityÚ
ValueErrorr!   ÚgetÚ
isinstanceÚdictÚvalues)	r   Úargsr   r.   Úis_device_supported_mxfp4Útriton_availableÚkernels_installedÚcompute_capabilityr(   s	            r   Úvalidate_environmentz%Mxfp4HfQuantizer.validate_environmentA   s¯  € Ü!Ô#Üð]óð ð
 ×#Ñ#×.Ò.Øä&Ô(ÜÐYÓZÐZä×"Ñ"×6Ñ6Ó8ÒO¼E¿L¹LÈÓ<OˆØ�;‰;Ð4Ñ4Ø×!Ò!Ü×#Ñ#Ø]Ð^dÐ]eð  fPð  Qôð 7;�×(Ñ(Ô3Øä"Ø_Ð`fÐ_gð  h_ð  `óð ô �9‰9×!Ñ!Ô#Ø(,Ð%Ü2°7Ó;ÐÜ 4Ó 6ÑÜ�Z‰Z×$Ñ$Ô&Ü!&§¡×!AÑ!AÓ!CÐØ(:¸fÑ(DÐ%Ü2°7Ó;ÐÜ 4Ó 6ÑØ�[‰[˜EÒ!Ø(,Ð%Ü2°7Ó;ÐÜ 4Ó 6Ñà(-Ð%Ø$ÐØ %Ðà×ÒÙ,Ü×#Ñ#ðIôð
 7;�×(Ñ(Ô3Øá#Ü×#Ñ#ðIôð
 7;�×(Ñ(Ô3Øá$Ü×#Ñ#ðIôð
 7;�×(Ñ(Ô3ØÙ*Üðlóð ñ "Üð`óð ñ #ÜÐoÓpÐpà×!Ò!Ø×%Ñ%Ô'à—Z‘Z Ó-ˆ
ØÐ!¤j°¼TÔ&BØ×%Ò%¨&°J×4EÑ4EÓ4GÑ*GÜ ðgóð ð +HÐ%ð 'CÐ!r   Úmodelr   Ú
param_nameÚreturnc                 óR   — ddl m} t        ||«      \  }}t        ||«      r|dv ryyy)Nr   ©ÚMxfp4GptOssExperts)Údown_proj_biasÚgate_up_proj_biasFT)ÚintegrationsrF   r   r8   )r   rA   rB   r   rF   ÚmoduleÚtensor_names          r   Úparam_needs_quantizationz)Mxfp4HfQuantizer.param_needs_quantization¡   s3   € Ý5ä2°5¸*ÓEÑˆ�Ü�fÐ0Ô1ØÐEÑEØØØr   c                 óø   — t         j                  j                  «       rt         j                  j                  «        y t         j                  j                  «       rt         j                  j                  «        y y r   )r+   r$   r4   Úempty_cacher%   )r   rA   r   s      r   Ú#_process_model_after_weight_loadingz4Mxfp4HfQuantizer._process_model_after_weight_loading«   sG   € ä�:‰:×"Ñ"Ô$Ü�J‰J×"Ñ"Õ$Ü�Y‰Y×#Ñ#Ô%Ü�I‰I×!Ñ!Õ#ð &r   Úuse_kernelsc                 óü  — ddl m} t        j                  j	                  «       xs t        j
                  d«      }|r4|j                  dvr&t        j                  d«       d| j                  _
        |s4|j                  dv r&t        j                  d«       d| j                  _
        | j                  || j                  j                  |j                  «      | _         ||| j                  | j                  ¬«      }y )	Nr   )Úreplace_with_mxfp4_linearr#   )r#   zž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=FalseTz¯MXFP4 inference on CPU requires use_kernels=True, but use_kernels is disabled. We will dequantize the model to bf16. To run MXFP4 natively on CPU, please set use_kernels=True.)Úmodules_to_not_convertr   )rI   rR   r+   r,   r-   r.   r/   r1   r2   r   r*   Úget_modules_to_not_convertrS   Ú_keep_in_fp32_modules)r   rA   rP   r   rR   r.   s         r   Ú$_process_model_before_weight_loadingz5Mxfp4HfQuantizer._process_model_before_weight_loading²   sÛ   € õ 	=ô ×"Ñ"×6Ñ6Ó8ÒO¼E¿L¹LÈÓ<OˆÙ˜6Ÿ;™;¨gÑ5Ü×Ñðeôð 37ˆD×$Ñ$Ô/á˜vŸ{™{¨gÑ5Ü×Ñðsôð 37ˆD×$Ñ$Ô/à&*×&EÑ&EØ�4×+Ñ+×BÑBÀE×D_ÑD_ó'
ˆÔ#ñ *Ø¨$×*EÑ*EÐ[_×[sÑ[sô
‰r   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__ÚgetattrrY   Úupdate©r   Úconfigs     r   Úupdate_tp_planzMxfp4HfQuantizer.update_tp_planÓ   ó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 )NrX   Úbase_model_ep_planrZ   r[   )r   r\   r]   rd   r^   r_   s     r   Úupdate_ep_planzMxfp4HfQuantizer.update_ep_planà   rb   r   c                 ó0  — ddl m} |j                  «       }t        |j                  dd«      }t        |j                  dd«      }|j                  «       D �]9  \  }}t        ||«      rt        |d«      rt        |d«      sŒ,d	D �]  }t        ||«      }	t        ||› d
�«      }
|	j                  j                  j                  |	j                  j                  «      j                  dd«      }|dk(  r|j                  |ddd«      }n|j                  ||dd«      }|
j                  j                  j                  j                  |
j                  j                  j                  «      j                  dd«      }|||› d|› d�<   |||› d|› d�<   �Œ �Œ< i }||fS )Nr   rE   Únum_local_expertsé    Úhidden_sizei@  Úgate_up_projÚ	down_proj)rj   rk   Ú_precision_configéÿÿÿÿéþÿÿÿéZ   é   ú.Ú_blocksÚ_scales)rI   rF   Ú
state_dictr]   r`   Únamed_modulesr8   ÚhasattrÚstorageÚlayoutÚunswizzle_dataÚdataÚ	transposeÚreshapeÚweight_scale)r   rA   rF   rt   rg   ri   ÚnamerJ   ÚprojÚtriton_tensorÚprecision_configÚblocksÚscalesÚmetadatas                 r   Úget_state_dict_and_metadataz,Mxfp4HfQuantizer.get_state_dict_and_metadataí   s™  € Ý5à×%Ñ%Ó'ˆ
Ü# E§L¡LÐ2EÀrÓJÐÜ˜eŸl™l¨M¸4Ó@ˆà!×/Ñ/Ó1ó 	=‰LˆD�&ä˜6Ð#5Ô6Ü˜F NÔ3Ü˜F KÔ0àà5ó =�Ü '¨°Ó 5�Ü#*¨6°d°VÐ;LÐ3MÓ#NÐ à&×.Ñ.×5Ñ5×DÑDÀ]×EZÑEZ×E_ÑE_Ó`×jÑjÐkmÐoqÓr�Ø˜>Ò)Ø#Ÿ^™^Ð,=¸rÀ2ÀrÓJ‘Fà#Ÿ^™^Ð,=¸{ÈBÐPRÓS�Fà)×6Ñ6×>Ñ>×EÑE×TÑTØ$×1Ñ1×9Ñ9×>Ñ>óç‘)˜B Ó#ð ð 7=�
˜d˜V 1 T F¨'Ð2Ñ3Ø6<�
˜d˜V 1 T F¨'Ð2Ó3ò=ð	=ð2 ˆØ˜8Ð#Ð#r   c                  ó   — y)NT© ©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)r1   r2   rˆ   s    r   Úis_trainablezMxfp4HfQuantizer.is_trainable  s   € ä×Ñð xô	
ð r   c                 ó   — ddl m}  || «      S )Nr   )ÚMxfp4Quantize)Úintegrations.mxfp4r�   )r   r�   s     r   Úget_quantize_opsz!Mxfp4HfQuantizer.get_quantize_ops  s   € Ý6á˜TÓ"Ð"r   c                 ó  — ddl m}m} | j                  rE| j                  j
                  r/t        ddgd || «      g¬«      t        ddgd	g || «      g¬«      gS t        ddgd	 || «      g¬«      t        ddgd || «      g¬«      gS )
Nr   )ÚMxfp4DequantizeÚMxfp4DeserializeÚdown_proj_blocksÚdown_proj_scalesz
down_proj$)Úsource_patternsÚtarget_patternsÚ
operationsÚgate_up_proj_blocksÚgate_up_proj_scaleszgate_up_proj$)rŽ   r‘   r’   r0   r   r*   r   )r   r‘   r’   s      r   Úget_weight_conversionsz'Mxfp4HfQuantizer.get_weight_conversions  sµ   € ßJà×Ò $×":Ñ":×"EÒ"EäØ%7Ð9KÐ$LØ$1Ù /°Ó 5Ð6ôô
  Ø%:Ð<QÐ$RØ%4Ð$5Ù /°Ó 5Ð6ôðð ô Ø!6Ð8MÐ NØ 0Ù,¨TÓ2Ð3ôô
 Ø!3Ð5GÐ HØ -Ù,¨TÓ2Ð3ôð
ð 	
r   )rA   r   )F)r\   Ú
__module__Ú__qualname__Ú__doc__Úrequires_calibrationÚ__annotations__r   r!   r@   ÚstrÚboolrL   rO   rV   ra   re   r…   r‰   Úpropertyr‹   r�   rš   Ú__classcell__)r   s   @r   r   r   *   sŸ   ø… ñð !ÐØ&Ó&ô'ò	'ò^ð@Ð.?ð ÈSð Ð_có ó$ð "ñ
à ð
ð ó
òBòò!$òFð ð˜dò ó ðò#ö

r   r   )Útypingr   Úbaser   Úmodeling_utilsr   Úutils.quantization_configr   Úutilsr	   r
   r   r   r   Úquantizers_utilsr   r+   Úcore_model_loadingr   Ú
get_loggerr\   r1   r   r   r‡   r   r   ú<module>r¬      s[   ðõ !å ñ Ý0Ý7÷õ õ 3ñ ÔÛå4à	ˆ×	Ñ	˜HÓ	%€ØÐ ôQ
�{õ Q
r   