ó
    qyüia&  ã                   óÌ   • S SK Jr  SSKJr  \(       a  SSKJr  SSKJr  SSKJ	r	J
r
JrJrJrJrJr  SSKJr  \" 5       (       a  S S	Kr\R&                  " \5      r " S
 S\5      rg	)é    )ÚTYPE_CHECKINGé   )ÚHfQuantizeré   )ÚPreTrainedModel)ÚFbgemmFp8Config)Úis_accelerate_availableÚis_fbgemm_gpu_availableÚis_kernels_availableÚis_torch_availableÚis_torch_cuda_availableÚis_torch_xpu_availableÚlogging)Úget_module_from_nameNc                   óÀ   ^ • \ rS rSr% SrSrS\S'   U 4S jrS rSS	 j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 r\S\4S j5       rS rSrU =r$ )ÚFbgemmFp8HfQuantizeré)   z'
FP8 quantization using fbgemm kernels
Fr   Úquantization_configc                 ó(   >• [         TU ]  " U40 UD6  g )N)ÚsuperÚ__init__)Úselfr   ÚkwargsÚ	__class__s      €Úi/home/mande/repo/quber/.venv/lib/python3.13/site-packages/transformers/quantizers/quantizer_fbgemm_fp8.pyr   ÚFbgemmFp8HfQuantizer.__init__1   s   ø€ Ü‰ÒÐ,Ñ7°Ó7ó    c                 ó¼  • [        5       (       d  [        5       (       d  [        S5      e[        5       (       a  [        5       (       d  [        S5      e[        5       (       a  [	        5       (       d  [        S5      e[        5       (       d  [        S5      e[        5       (       a3  [        R                  R                  5       nUu  pEUS:  a  [        S5      eUR                  S5      nUc  [        R                  S5        g [        U[        5      (       aF  U R                  (       d4  S	UR!                  5       ;   d  S
UR!                  5       ;   a  [        S5      eg g g )Nz3Using fbgemm fp8 quantization requires a GPU or XPUz@Using FP8 fbgemm on XPU requires kernels (`pip install kernels`)züLoading an FP8 fbgemm quantized model on CUDA requires fbgemm-gpu libraryPlease install the latest version of fbgemm-gpu library by following : https://pytorch.org/FBGEMM/fbgemm_gpu-development/InstallationInstructions.html#fbgemm-gpu-install-librarieszWLoading an FP8 quantized model requires accelerate (`pip install --upgrade accelerate`)é	   zXFP8 quantized models is only supported on GPUs with compute capability >= 9.0 (e.g H100)Ú
device_mapzÛYou have loaded an FP8 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 'xpu' or 'auto'. ÚcpuÚdiskzòYou are attempting to load an FP8 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   ÚImportErrorr   r
   r	   ÚtorchÚcudaÚget_device_capabilityÚ
ValueErrorÚgetÚloggerÚwarning_onceÚ
isinstanceÚdictÚpre_quantizedÚvalues)r   Úargsr   Úcompute_capabilityÚmajorÚ_r    s          r   Úvalidate_environmentÚ)FbgemmFp8HfQuantizer.validate_environment4   s?  € Ü&×(Ñ(Ô1G×1IÑ1IÜÐSÓTÐTÜ!×#Ñ#Ô,@×,BÑ,BÜÐ`ÓaÐaÜ"×$Ñ$Ô-D×-FÑ-FÜðFóð ô '×(Ñ(ÜØióð ô #×$Ñ$Ü!&§¡×!AÑ!AÓ!CÐØ)‰HˆEØ�q‹yÜ Ønóð ð —Z‘Z Ó-ˆ
ØÑÜ×ÑðSõô ˜
¤D×)Ñ)Ø×%×%¨5°J×4EÑ4EÓ4GÓ+GÈ6ÐU_×UfÑUfÓUhÓKhÜ ðnóð ð LiÐ%ð *r   Úreturnc                 ó€   • U[         R                  :w  a)  [        R                  SU S35        [         R                  nU$ )NzSetting dtype to zP, but only bfloat16 is supported right now. Overwriting torch_dtype to bfloat16.)r$   Úbfloat16r)   r*   )r   Údtypes     r   Úupdate_dtypeÚ!FbgemmFp8HfQuantizer.update_dtypeX   s9   € Ø”E—N‘NÓ"Ü×ÑØ# E 7Ð*zÐ{ôô —N‘NˆEØˆr   Úmodelr   Ú
param_namec                 óÒ   • SSK JnJn  [        X5      u  pg[	        Xd5      (       a  U R
                  (       d  US:X  a  gg[	        Xe5      (       a  U R
                  (       d  US:X  a  ggg)Nr   ©ÚFbgemmFp8LinearÚFbgemmFp8Llama4TextExpertsÚbiasFT)Úintegrationsr?   r@   r   r+   r-   )r   r;   r<   r   r?   r@   ÚmoduleÚtensor_names           r   Úparam_needs_quantizationÚ-FbgemmFp8HfQuantizer.param_needs_quantization`   sW   € ßNä2°5ÓEÑˆä�f×.Ñ.Ø×!×! [°FÓ%:ØàÜ�f×9Ñ9Ø×!×! [°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   )rE   r   Úparam_element_size)r   r;   r<   rG   r   s       €r   rI   Ú'FbgemmFp8HfQuantizer.param_element_sizeq   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                  UR                  S9ng )Nr   )Úreplace_with_fbgemm_fp8_linear)Úmodules_to_not_convertr   r-   Útp_plan)rB   rL   Úget_modules_to_not_convertr   rM   Ú_keep_in_fp32_modulesr-   Ú_tp_plan)r   r;   r   rL   s       r   Ú$_process_model_before_weight_loadingÚ9FbgemmFp8HfQuantizer._process_model_before_weight_loadingx   sc   € õ
 	Bà&*×&EÑ&EØ×+Ñ+×BÑBÀE×D_ÑD_ó'
ˆÔ#ñ /ØØ#'×#>Ñ#>Ø $× 8Ñ 8Ø×,Ñ,Ø—N‘Nñ
‰r   c                 óð   • SSK JnJn  UR                  5        HY  n[	        XSU45      (       d  M  [        US5      (       d  M*  UR                  R                  U R                  R                  5        M[     U$ )zÉ
Force update the input scale upper bound after weight loading and device dispatch are complete.
This resolves issues where persistent buffers are zeroed out or overwritten during the loading process.
r   r>   Úinput_scale_ub)
Úintegrations.fbgemm_fp8r?   r@   Úmodulesr+   ÚhasattrrU   Úfill_r   Úactivation_scale_ub)r   r;   r   r?   r@   Úms         r   Ú#_process_model_after_weight_loadingÚ8FbgemmFp8HfQuantizer._process_model_after_weight_loading‹   s^   € ÷
 	Zà—‘–ˆAÜ˜!Ð/IÐJ×KÓKÜ˜1Ð.×/Ó/à×$Ñ$×*Ñ*¨4×+CÑ+C×+WÑ+WÖXñ	 !ð
 ˆr   c                 ó  • SUR                   R                  ;   am  0 SS_SS_SS_SS_SS_SS_S	S
_SS_SS_SS_SS_SS_SS_SS_SS
_SS_SS_SSS
SSSS.EnUR                  5       b  X!R                  5       l        U$ X!l        U$ U$ )NÚLlama4z layers.*.self_attn.q_proj.weightÚcolwisez&layers.*.self_attn.q_proj.weight_scalez layers.*.self_attn.k_proj.weightz&layers.*.self_attn.k_proj.weight_scalez layers.*.self_attn.v_proj.weightz&layers.*.self_attn.v_proj.weight_scalez layers.*.self_attn.o_proj.weightÚrowwisezlayers.*.input_layernorm.weightÚsequence_parallelz(layers.*.post_attention_layernorm.weightznorm.weightz4layers.*.feed_forward.shared_expert.gate_proj.weightz:layers.*.feed_forward.shared_expert.gate_proj.weight_scalez2layers.*.feed_forward.shared_expert.up_proj.weightz8layers.*.feed_forward.shared_expert.up_proj.weight_scalez4layers.*.feed_forward.shared_expert.down_proj.weightz0layers.*.feed_forward.experts.*.gate_proj.weightz6layers.*.feed_forward.experts.*.gate_proj.weight_scaleÚpacked_rowwise)z.layers.*.feed_forward.experts.*.up_proj.weightz4layers.*.feed_forward.experts.*.up_proj.weight_scalez0layers.*.feed_forward.experts.*.down_proj.weightz*layers.*.feed_forward.experts.gate_up_projz0layers.*.feed_forward.experts.gate_up_proj_scalez'layers.*.feed_forward.experts.down_proj)r   Ú__name__Úget_text_configÚbase_model_tp_plan)r   ÚconfigÚ	text_plans      r   Úupdate_tp_planÚ#FbgemmFp8HfQuantizer.update_tp_plan™   sL  € Ø�v×'Ñ'×0Ñ0Ó0ð!ð 3°Ið	!ð
 9¸)ð!ð 3°Ið!ð 9¸)ð!ð 3°Ið!ð 9¸)ð!ð 3°Ið!ð 2Ð3Fð!ð ;Ð<Oð!ð Ð2ð!ð$ GÈ	ð%!ð& MÈið'!ð( EÀið)!ð* KÈIð+!ð, GÈ	ð-!ð. CÀIð/!ð0 IÈ)ð1!ð2 CLØHQØDMð ?OØDTØ;DòA!ˆIðD ×%Ñ%Ó'Ñ3Ø>G×&Ñ&Ó(Ô;ð ˆMð -6Ô)ØˆMàˆr   c                 ó   • g)NT© ©r   s    r   Úis_serializableÚ$FbgemmFp8HfQuantizer.is_serializableÅ   s   € Ør   c                 ó   • g)NFrl   rm   s    r   Úis_trainableÚ!FbgemmFp8HfQuantizer.is_trainableÈ   s   € àr   c                 ó   • SSK Jn  U" U 5      $ )Nr   )ÚFbgemmFp8Quantize)rV   rt   )r   rt   s     r   Úget_quantize_opsÚ%FbgemmFp8HfQuantizer.get_quantize_opsÌ   s   € Ý?á  Ó&Ð&r   )rM   )r8   útorch.dtyper5   rw   )r;   r   )rd   Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__Úrequires_calibrationÚ__annotations__r   r3   r9   ÚstrÚboolrE   ÚfloatrI   rR   r\   ri   rn   Úpropertyrq   ru   Ú__static_attributes__Ú__classcell__)r   s   @r   r   r   )   s­   ø‡ ñð !ÐØ*Ó*õ8ò"ôHðÐ.?ð ÈSð Ð_cô ð"DÐ(9ð DÀsð DÐSað DÐfk÷ Dð
à ô
ò&ò*òXð ð˜dó ó ð÷'ð 'r   r   )Útypingr   Úbaser   Úmodeling_utilsr   Úutils.quantization_configr   Úutilsr	   r
   r   r   r   r   r   Úquantizers_utilsr   r$   Ú
get_loggerrd   r)   r   rl   r   r   Ú<module>r‹      sX   ðõ !å ö Ý0Ý;÷÷ ñ õ 3ñ ×ÑÛà	×	Ò	˜HÓ	%€ôf'˜;õ f'r   