ó
    qyüif)  ã                   óÀ   • S SK Jr  SSKJrJrJrJr  SSKJr  SSK	J
r
  \" 5       (       a  S SKr\(       a  SSKJr  SS	KJr  \R                   " \5      r " S
 S\5      rg)é    )ÚTYPE_CHECKINGé   )Úis_accelerate_availableÚis_torch_availableÚis_torch_xpu_availableÚloggingé   )ÚHfQuantizer)Úget_module_from_nameN)ÚPreTrainedModel)ÚFineGrainedFP8Configc                   óÔ   ^ • \ rS rSr% SrSrS\S'   U 4S jrS rSS	S
\	S\
4S jrSS	S
\	SSS\4U 4S jjr  SS jrS rS r\S\
4S j5       r\S\
4S j5       rS rS rS rSrU =r$ )ÚFineGrainedFP8HfQuantizeré   zz
FP8 quantization implementation supporting both standard and MoE models.
Supports both e4m3fn formats based on platform.
Fr   Úquantization_configc                 ó(   >• [         TU ]  " U40 UD6  g )N)ÚsuperÚ__init__)Úselfr   ÚkwargsÚ	__class__s      €Ún/home/mande/repo/quber/.venv/lib/python3.13/site-packages/transformers/quantizers/quantizer_finegrained_fp8.pyr   Ú"FineGrainedFP8HfQuantizer.__init__   s   ø€ Ü‰ÒÐ,Ñ7°Ó7ó    c                 óŠ  • [        5       (       d  [        S5      eU R                  R                  (       a  g [        R
                  R                  5       (       dR  [        5       (       dC  U R                  (       a'  [        R                  S5        SU R                  l        g [        S5      e[        R
                  R                  5       (       ab  [        R
                  R                  5       nUu  pEUS:  d  US:X  a4  US:  a.  [        R                  SU SU S	35        SU R                  l        g UR                  S
5      nUc  [        R                  S5        g [        U[        5      (       aT  U R                  (       d#  [!        U5      S:”  a  SUR#                  5       ;   d  SUR#                  5       ;   a  [%        S5      eg g )NzMLoading an FP8 quantized model requires accelerate (`pip install accelerate`)z„Using FP8 quantized models requires a GPU or XPU, we will default to dequantizing the model to bf16 since no GPU or XPU is availableTzANo GPU or XPU found. A GPU or XPU is needed for FP8 quantization.é   é	   ziFP8 quantized models is only supported on GPUs with compute capability >= 8.9 (e.g 4090/H100), actual = `Ú.zƒ`. We will default to dequantizing the model to bf16. Feel free to use a different quantization method like bitsandbytes or torchaoÚ
device_mapz×You have loaded an FP8 model on CPU and have a CUDA or XPU device available, make sure to set your model on a GPU or XPU device in order to run your model. To remove this warning, pass device_map = 'cuda' or 'xpu'. r	   ÚcpuÚdiskzìYou are attempting to load an FP8 model with a device_map that contains a cpu/disk device.This is not supported when the model is quantized on the fly. Please use a quantized checkpoint or remove the cpu/disk device from the device_map.)r   ÚImportErrorr   Ú
dequantizeÚtorchÚcudaÚis_availabler   Úpre_quantizedÚloggerÚwarning_onceÚRuntimeErrorÚget_device_capabilityÚgetÚ
isinstanceÚdictÚlenÚvaluesÚ
ValueError)r   Úargsr   Úcompute_capabilityÚmajorÚminorr   s          r   Úvalidate_environmentÚ.FineGrainedFP8HfQuantizer.validate_environment   sˆ  € Ü&×(Ñ(ÜÐmÓnÐnà×#Ñ#×.×.Øä�z‰z×&Ñ&×(Ñ(Ô1G×1IÑ1IØ×!×!Ü×#Ñ#ð [ôð 7;�×(Ñ(Ô3Øä"Ð#fÓgÐgä�:‰:×"Ñ"×$Ñ$Ü!&§¡×!AÑ!AÓ!CÐØ-‰LˆEØ˜“	˜u¨›z¨e°a«iÜ×#Ñ#ð#Ø#( '¨¨5¨'ð 2Zð[ôð
 7;�×(Ñ(Ô3Øà—Z‘Z Ó-ˆ
ØÑÜ×Ñð6õô
 ˜
¤D×)Ñ)à×&×&Ü˜
“O aÓ'Ø˜Z×.Ñ.Ó0Ó0Ø˜Z×.Ñ.Ó0Ó0ä ðkóð ð 1ð *r   Úmodelr   Ú
param_nameÚreturnc                 ó„   • SSK JnJn  [        X5      u  pg[	        XeU45      (       a  U R
                  (       d  US:X  a  ggg)Nr   )Ú
FP8ExpertsÚ	FP8LinearÚbiasFT)Úintegrations.finegrained_fp8r<   r=   r   r-   r'   )r   r8   r9   r   r<   r=   ÚmoduleÚtensor_names           r   Úparam_needs_quantizationÚ2FineGrainedFP8HfQuantizer.param_needs_quantizationO   s;   € ßHä2°5ÓEÑˆÜ�f¨*Ð5×6Ñ6Ø×!×! [°FÓ%:ØàØr   Úparamztorch.Tensorc                 óR   >• U R                  X5      (       a  g[        TU ]	  XU5      $ )z4Return the element size (in bytes) for `param_name`.r	   )rB   r   Úparam_element_size)r   r8   r9   rD   r   s       €r   rF   Ú,FineGrainedFP8HfQuantizer.param_element_sizeZ   s)   ø€ à×(Ñ(¨×;Ñ;àÜ‰wÑ)¨%¸UÓCÐCr   c                 óÈ   • SSK Jn  U R                  XR                  R                  UR
                  5      U l        U" UU R                  U R                  U R                  S9ng )Nr   )Úreplace_with_fp8_linear)Úmodules_to_not_convertr   r'   )r?   rI   Úget_modules_to_not_convertr   rJ   Ú_keep_in_fp32_modulesr'   )r   r8   r   rI   s       r   Ú$_process_model_before_weight_loadingÚ>FineGrainedFP8HfQuantizer._process_model_before_weight_loadinga   s\   € õ
 	Kà&*×&EÑ&EØ×+Ñ+×BÑBÀE×D_ÑD_ó'
ˆÔ#ñ (ØØ#'×#>Ñ#>Ø $× 8Ñ 8Ø×,Ñ,ñ	
‰r   c                 óh   • SUR                   R                  ;   a  SSSSSSSSSSSSSSS.nX!l        U$ )NÚQwen3ÚcolwiseÚrowwise)z layers.*.self_attn.q_proj.weightz*layers.*.self_attn.q_proj.weight_scale_invz layers.*.self_attn.k_proj.weightz*layers.*.self_attn.k_proj.weight_scale_invz layers.*.self_attn.v_proj.weightz*layers.*.self_attn.v_proj.weight_scale_invz layers.*.self_attn.o_proj.weightz*layers.*.self_attn.o_proj.weight_scale_invzlayers.*.mlp.gate_proj.weightz'layers.*.mlp.gate_proj.weight_scale_invzlayers.*.mlp.up_proj.weightz%layers.*.mlp.up_proj.weight_scale_invzlayers.*.mlp.down_proj.weightz'layers.*.mlp.down_proj.weight_scale_inv)r   Ú__name__Úbase_model_tp_plan)r   ÚconfigÚ	text_plans      r   Úupdate_tp_planÚ(FineGrainedFP8HfQuantizer.update_tp_plans   sT   € Ø�f×&Ñ&×/Ñ/Ó/à4=Ø>GØ4=Ø>GØ4=Ø>GØ4=Ø>GØ1:Ø;DØ/8Ø9BØ1:Ø;DñˆIð" )2Ô%àˆr   c                 ó   • g©NT© ©r   s    r   Úis_serializableÚ)FineGrainedFP8HfQuantizer.is_serializableŠ   s   € Ør   c                 ó   • g)NFr[   r\   s    r   Úis_trainableÚ&FineGrainedFP8HfQuantizer.is_trainable�   s   € àr   c                 ó   • grZ   r[   r\   s    r   Úis_compileableÚ(FineGrainedFP8HfQuantizer.is_compileable‘   s   € àr   c                 ó   • SSK Jn  U" U 5      $ )Nr   )ÚFp8Quantize)r?   rf   )r   rf   s     r   Úget_quantize_opsÚ*FineGrainedFP8HfQuantizer.get_quantize_ops•   s   € Ý>á˜4Ó Ð r   c                 óš   • SSK Jn  SSKJn  U R                  (       a-  U R
                  R                  (       a  U" / SQSU" U 5      /S9/$ / $ )Nr   )ÚWeightConverter©ÚFp8Dequantize)zweight$Úweight_scale_invÚactivation_scaleÚweight©Úsource_patternsÚtarget_patternsÚ
operations)Úcore_model_loadingrj   r?   rl   r'   r   r#   )r   rj   rl   s      r   Úget_weight_conversionsÚ0FineGrainedFP8HfQuantizer.get_weight_conversionsš   sK   € Ý8Ý@à×× $×":Ñ":×"E×"Eñ  Ú$WØ$,Ù -¨dÓ 3Ð4ñðð ð ˆ	r   c           	      óf  • U R                   (       a  U R                  R                  (       d  XR                  5       -   $ SSKJnJn  SSKJn  U" SSS9nU/[        U5      -   n/ nU GH  n[        Xr5      (       d  UR                  U5        M'  UR                   Vs/ s H  oˆR                  S5      (       d  M  UPM     n	nU	(       a   U	 Vs/ s H  oˆS-   PM	     n
nU	 Vs/ s H  oˆS	[        S5      *  S
-   PM     nnUR                   Vs/ s H  oˆR                  S5      (       a  M  UPM     nnX«-   U-   nU" U 5      /[        UR                  5      -   nU" UUR                   US9nUR                  U5        GM     UR#                  U R                  5       5        U$ s  snf s  snf s  snf s  snf )u]  When loading with ``dequantize=True``, attach an :class:`Fp8Dequantize` op to
every existing :class:`WeightConverter` so that per-block scales are folded into
the weight *before* any later merge/concat ops collapse the per-expert structure.

For each model-supplied converter that has a ``.weight`` source, we:
  1. anchor the existing weight patterns with ``$`` so they don't accidentally
     also match the ``.weight_scale_inv`` keys (the regex is searched, so the
     unanchored prefix would match both, sending scales to the wrong bucket);
  2. add anchored ``*.weight_scale_inv`` sources next to each weight pattern so
     the loader collects scale tensors alongside the weight tensors into the
     *same* converter bucket (both keys rewrite to the same target);
  3. prepend a fresh :class:`Fp8Dequantize` op so dequant runs first, before
     any merge/concat collapses the per-expert structure.

The generic ``weight$ + weight_scale_inv â†’ weight`` converter from
:meth:`get_weight_conversions` is still appended at the end as a fallback for
plain ``nn.Linear`` weights with no model-specific converter.
r   )rj   ÚWeightRenamingrk   z^(.+)\.scale$z\1.weight_scale_inv)rq   rr   z.weightÚ$Nz.weight_scale_inv$rp   )r'   r   r#   ru   rt   rj   rx   r?   rl   Úlistr-   Úappendrq   Úendswithr/   rs   Ú_original_target_patternsÚextend)r   Úweight_conversionsrj   rx   rl   Úscale_renameÚupdatedÚconvÚpÚweight_sourcesÚanchored_weightÚscale_sourcesÚotherÚnew_sourcesÚnew_opss                  r   Úupdate_weight_conversionsÚ3FineGrainedFP8HfQuantizer.update_weight_conversionsª   s‡  € ð& ×"×" t×'?Ñ'?×'J×'JØ%×(CÑ(CÓ(EÑEÐEçHÝ@ñ &Ð6FÐXnÑoˆØ*˜^¬dÐ3EÓ.FÑFÐàˆÜ&ˆDô ˜d×4Ñ4Ø—‘˜tÔ$ÙØ)-×)=Ò)=ÓWÒ)= AÇÁÈI×AVŸaÑ)=ˆNÐWÞÙ4BÓ"C²N¨q s¤7±N�Ð"CÙVdÓ eÒVdÐQRÐ#4¤c¨)£n _Ð!5Ð8LÔ!LÑVd�Ð eØ$(×$8Ò$8ÓVÒ$8˜qÇ
Á
È9×@UŸÑ$8�ÐVØ-Ñ=ÀÑE�Ù(¨Ó.Ð/´$°t·±Ó2GÑG�Ù&Ø$/Ø$(×$BÑ$BØ&ñ�ð
 �N‰N˜4× ñ% 'ð( 	�‰�t×2Ñ2Ó4Ô5Øˆùò Xùâ"CùÚ eùÚVs$   ÂFÂ9FÃF$Ã F)ÄF.Ä(F.)rJ   )r8   r   )rS   Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__Úrequires_calibrationÚ__annotations__r   r6   ÚstrÚboolrB   ÚfloatrF   rM   rW   r]   Úpropertyr`   rc   rg   ru   rŠ   Ú__static_attributes__Ú__classcell__)r   s   @r   r   r      sÇ   ø‡ ñð
 !ÐØ/Ó/õ8ò/ðb	Ð.?ð 	ÈSð 	Ð_cô 	ðDÐ(9ð DÀsð DÐSað DÐfk÷ Dð
à ô
ò$ò.ð ð˜dó ó ðð ð ó ó ðò!ò
÷ 8ð 8r   r   )Útypingr   Úutilsr   r   r   r   Úbaser
   Úquantizers_utilsr   r$   Úmodeling_utilsr   Úutils.quantization_configr   Ú
get_loggerrS   r(   r   r[   r   r   Ú<module>rŸ      sI   ðÝ  ç `Ó `Ý Ý 2ñ ×ÑÛæÝ0Ý@à	×	Ò	˜HÓ	%€ôP õ Pr   