§
    ‚Štjb&  ã                   óÂ   — 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mZmZ ddlmZ  e¦   «         rd d	lZ ej        e¦  «        Z G d
„ de¦  «        Zd	S )é    )Ú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                   ó°   ‡ — e Zd ZU dZdZded<   ˆ fd„Zd„ Zdd
„Zddde	d	e
fd„Zddde	ddd	efˆ fd„Z	 	 dd„Zd„ Zd„ Zd„ Zed	e
fd„¦   «         Zd„ Zˆ xZS )ÚFbgemmFp8HfQuantizerz/
    FP8 quantization using fbgemm kernels
    Fr   Úquantization_configc                 ó<   •—  t          ¦   «         j        |fi |¤Ž d S )N)ÚsuperÚ__init__)Úselfr   ÚkwargsÚ	__class__s      €új/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/quantizers/quantizer_fbgemm_fp8.pyr   zFbgemmFp8HfQuantizer.__init__1   s)   ø€ Ø�‰ŒÔÐ,Ð7Ð7°Ð7Ð7Ð7Ð7Ð7ó    c                 óê  — t          ¦   «         st          ¦   «         st          d¦  «        ‚t          ¦   «         rt          ¦   «         st          d¦  «        ‚t          ¦   «         rt	          ¦   «         st          d¦  «        ‚t          ¦   «         st          d¦  «        ‚t          ¦   «         r8t          j                             ¦   «         }|\  }}|dk     rt          d¦  «        ‚| 
                    d¦  «        }|€t                               d¦  «         d S t          |t          ¦  «        rB| j        s=d	|                     ¦   «         v sd
|                     ¦   «         v rt          d¦  «        ‚d S d S d S )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_environmentz)FbgemmFp8HfQuantizer.validate_environment4   s¿  € Ý&Ñ(Ô(ð 	UÕ1GÑ1IÔ1Ið 	UÝÐSÑTÔTÐTÝ!Ñ#Ô#ð 	bÕ,@Ñ,BÔ,Bð 	bÝÐ`ÑaÔaÐaÝ"Ñ$Ô$ð 	Õ-DÑ-FÔ-Fð 	ÝðFñô ð õ 'Ñ(Ô(ð 	ÝØiñô ð õ #Ñ$Ô$ð 	Ý!&¤×!AÒ!AÑ!CÔ!CÐØ)‰HˆE�1Ø�qŠyˆyÝ Ønñô ð ð —Z’Z Ñ-Ô-ˆ
ØÐÝ×ÒðSñô ð ð ð õ ˜
¥DÑ)Ô)ð 	ØÔ%ð ¨5°J×4EÒ4EÑ4GÔ4GÐ+GÐ+GÈ6ÐU_×UfÒUfÑUhÔUhÐKhÐKhÝ ðnñô ð ð	ð 	ðð ÐKhÐKhr   Údtypeútorch.dtypeÚreturnc                 óz   — |t           j        k    r*t                               d|› d�¦  «         t           j        }|S )NzSetting dtype to zP, but only bfloat16 is supported right now. Overwriting torch_dtype to bfloat16.)r"   Úbfloat16r'   r(   )r   r2   s     r   Úupdate_dtypez!FbgemmFp8HfQuantizer.update_dtypeX   sC   € Ø•E”NÒ"Ð"Ý×ÒØ{ EÐ{Ð{Ð{ñô ð õ ”NˆEØˆr   Úmodelr   Ú
param_namec                 óÀ   — ddl m}m} t          ||¦  «        \  }}t	          ||¦  «        r| j        s|dk    rdS dS t	          ||¦  «        r| j        s|dk    rdS dS dS )Nr   ©ÚFbgemmFp8LinearÚFbgemmFp8Llama4TextExpertsÚbiasFT)Úintegrationsr<   r=   r   r)   r+   )r   r8   r9   r   r<   r=   ÚmoduleÚtensor_names           r   Úparam_needs_quantizationz-FbgemmFp8HfQuantizer.param_needs_quantization`   s–   € ØNÐNÐNÐNÐNÐNÐNÐNå2°5¸*ÑEÔEÑˆ�å�f˜oÑ.Ô.ð 	ØÔ!ð  [°FÒ%:Ð%:Ø�uà�tÝ�fÐ8Ñ9Ô9ð 	ØÔ!ð  [°FÒ%:Ð%:Ø�uà�tØˆur   Úparamztorch.Tensorc                 óz   •— |                       ||¦  «        rdS t          ¦   «                              |||¦  «        S )z4Return the element size (in bytes) for `param_name`.r   )rB   r   Úparam_element_size)r   r8   r9   rC   r   s       €r   rE   z'FbgemmFp8HfQuantizer.param_element_sizeq   s<   ø€ à×(Ò(¨°
Ñ;Ô;ð 	à�1Ý‰wŒw×)Ò)¨%°¸UÑCÔCÐCr   c                 ó°   — ddl m} |                      || j        j        |j        ¦  «        | _         ||| j        | j        | j        |j        ¬¦  «        }d S )Nr   )Úreplace_with_fbgemm_fp8_linear)Úmodules_to_not_convertr   r+   Útp_plan)r?   rG   Úget_modules_to_not_convertr   rH   Ú_keep_in_fp32_modulesr+   Ú_tp_plan)r   r8   r   rG   s       r   Ú$_process_model_before_weight_loadingz9FbgemmFp8HfQuantizer._process_model_before_weight_loadingx   sv   € ð
 	BÐAÐAÐAÐAÐAà&*×&EÒ&EØ�4Ô+ÔBÀEÔD_ñ'
ô '
ˆÔ#ð /Ð.ØØ#'Ô#>Ø $Ô 8ØÔ,Ø”Nð
ñ 
ô 
ˆˆˆr   c                 óÐ   — ddl m}m} |                     ¦   «         D ]H}t	          |||f¦  «        r4t          |d¦  «        r$|j                             | j        j	        ¦  «         ŒI|S )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)   ÚhasattrrO   Úfill_r   Úactivation_scale_ub)r   r8   r   r<   r=   Úms         r   Ú#_process_model_after_weight_loadingz8FbgemmFp8HfQuantizer._process_model_after_weight_loading‹   s†   € ð
 	ZÐYÐYÐYÐYÐYÐYÐYà—’‘”ð 	Yð 	YˆAÝ˜!˜oÐ/IÐJÑKÔKð YÝ˜1Ð.Ñ/Ô/ð YàÔ$×*Ò*¨4Ô+CÔ+WÑXÔXÐXøØˆr   c                 ó  — d|j         j        v rui dd“dd“dd“dd“dd“dd“d	d
“dd“dd“dd“dd“dd“dd“dd“dd
“dd“dd“ddd
ddddœ¥}|                     ¦   «         �||                     ¦   «         _        n||_        |S |S )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_planz#FbgemmFp8HfQuantizer.update_tp_plan™   sW  € Ø�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�×&Ò&Ñ(Ô(Ô;Ð;à,5�Ô)ØˆMàˆr   c                 ó   — dS )NT© ©r   s    r   Úis_serializablez$FbgemmFp8HfQuantizer.is_serializableÅ   s   € Øˆtr   c                 ó   — dS )NFrd   re   s    r   Úis_trainablez!FbgemmFp8HfQuantizer.is_trainableÈ   s   € àˆur   c                 ó$   — ddl m}  || ¦  «        S )Nr   )ÚFbgemmFp8Quantize)rP   rj   )r   rj   s     r   Úget_quantize_opsz%FbgemmFp8HfQuantizer.get_quantize_opsÌ   s%   € Ø?Ð?Ð?Ð?Ð?Ð?à Ð  Ñ&Ô&Ð&r   )r2   r3   r4   r3   )r8   r   )r]   Ú
__module__Ú__qualname__Ú__doc__Úrequires_calibrationÚ__annotations__r   r1   r7   ÚstrÚboolrB   ÚfloatrE   rM   rV   rb   rf   Úpropertyrh   rk   Ú__classcell__)r   s   @r   r   r   )   sn  ø€ € € € € € ðð ð !ÐØ*Ð*Ð*Ñ*ð8ð 8ð 8ð 8ð 8ð"ð "ð "ðHð ð ð ðÐ.?ð ÈSð Ð_cð ð ð ð ð"DÐ(9ð DÀsð DÐSað DÐfkð Dð Dð Dð Dð Dð Dð
à ð
ð 
ð 
ð 
ð&ð ð ð*ð *ð *ðXð ð ð ð˜dð ð ð ñ „Xðð'ð 'ð 'ð 'ð 'ð 'ð 'r   r   )Útypingr   Úbaser   Úmodeling_utilsr   Úutils.quantization_configr   Úutilsr	   r
   r   r   r   r   r   Úquantizers_utilsr   r"   Ú
get_loggerr]   r'   r   rd   r   r   ú<module>r}      s3  ðð !Ð  Ð  Ð  Ð  Ð  à Ð Ð Ð Ð Ð ð ð <Ø0Ð0Ð0Ð0Ð0Ð0Ø;Ð;Ð;Ð;Ð;Ð;ðð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð 3Ð 2Ð 2Ð 2Ð 2Ð 2ð ÐÑÔð Ø€L€L€Là	ˆÔ	˜HÑ	%Ô	%€ðf'ð f'ð f'ð f'ð f'˜;ñ f'ô f'ð f'ð f'ð f'r   