Ë
    T^(h•  ã                   óÀ   — d dl Z d dlmZmZmZmZmZ d dlmZ ddl	m
Z
mZmZ ddlmZ ddlmZ  e«       rd dlZerdd	lmZ  ej(                  e«      Z G d
„ de«      Zy)é    N)ÚTYPE_CHECKINGÚAnyÚDictÚListÚOptional)Úversioné   )Úis_accelerate_availableÚis_torch_availableÚloggingé   )ÚHfQuantizer)Úget_module_from_name)ÚPreTrainedModelc                   ó   ‡ — e Zd ZdZdZdZdgZˆ fd„Zd„ Zdd„Z		 dd	d
ddde
dddee
ef   deee
      fd„Zd	d
ddde
dee
ef   fd„Z	 dd	d
deee
      fd„Zdd„Zdee
   de
dee
   fd„Zdd„Zedefd„«       Zˆ xZS )ÚFineGrainedFP8HfQuantizerz†
    FP8 quantization implementation supporting both standard and MoE models.
    Supports both e4m3fn formats based on platform.
    TFÚ
acceleratec                 ó4   •— t        ‰| �  |fi |¤Ž || _        y ©N)ÚsuperÚ__init__Úquantization_config)Úselfr   ÚkwargsÚ	__class__s      €úo/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/quantizers/quantizer_finegrained_fp8.pyr   z"FineGrainedFP8HfQuantizer.__init__   s   ø€ Ü‰ÑÐ,Ñ7°Ò7Ø#6ˆÕ ó    c                 ó  — t        «       rHt        j                  t        j                  j                  d«      «      t        j                  d«      k  rt        d«      ‚t        «       st        d«      ‚|j                  dd«      s|j                  dd«      rt        d«      ‚t        j                  j                  «       st        d	«      ‚t        j                  j                  «       }|\  }}|d
k  s
|d
k(  r|dk  rt        d|› d|› d�«      ‚|j                  dd «      }|€t        j                  d«       y |�N| j                   sAt#        |t$        «      r0d|j'                  «       v sd|j'                  «       v rt        d«      ‚y y y y )NÚtorchz2.1.0zxUsing fp8 quantization requires torch >= 2.1.0Please install the latest version of torch ( pip install --upgrade torch )zMLoading an FP8 quantized model requires accelerate (`pip install accelerate`)Úfrom_tfFÚ	from_flaxz€Converting into FP8 weights from tf/flax weights is currently not supported, please make sure the weights are in PyTorch format.z3No GPU found. A GPU is needed for FP8 quantization.é   é	   ziFP8 quantized models is only supported on GPUs with compute capability >= 8.9 (e.g 4090/H100), actual = `ú.ú`Ú
device_mapzÀYou have loaded an FP8 model on CPU and have a CUDA device available, make sure to set your model on a GPU device in order to run your model. To remove this warning, pass device_map = 'cuda'. Ú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   r   ÚparseÚ	importlibÚmetadataÚImportErrorr
   ÚgetÚ
ValueErrorr   ÚcudaÚis_availableÚRuntimeErrorÚget_device_capabilityÚloggerÚwarning_onceÚpre_quantizedÚ
isinstanceÚdictÚvalues)r   Úargsr   Úcompute_capabilityÚmajorÚminorr&   s          r   Úvalidate_environmentz.FineGrainedFP8HfQuantizer.validate_environment"   s�  € Ü!Ô#¤w§}¡}´Y×5GÑ5G×5OÑ5OÐPWÓ5XÓ'YÔ\c×\iÑ\iÐjqÓ\rÒ'rÜð]óð ô
 'Ô(ÜÐmÓnÐnà�:‰:�i Ô'¨6¯:©:°kÀ5Ô+IÜðFóð ô
 �z‰z×&Ñ&Ô(ÜÐTÓUÐUä"ŸZ™Z×=Ñ=Ó?ÐØ)‰ˆˆuØ�AŠI˜5 Aš:¨%°!ª)ÜðØ$˜g Q u g¨Qð0óð ð
 —Z‘Z ¨dÓ3ˆ
ØÐÜ×Ñð|õð Ð#à×&Ò&Ü˜z¬4Ô0Ø˜j×/Ñ/Ó1Ñ1°V¸z×?PÑ?PÓ?RÑ5Rä ðkóð ð 6Sð 1ð 'ð $r   Úreturnc                 óT   — |€%t         j                  d«       t        j                  }|S )NzWSetting torch_dtype to torch.float32 as no torch_dtype was specified in from_pretrained)r3   Úinfor   Úfloat32)r   Útorch_dtypes     r   Úupdate_torch_dtypez,FineGrainedFP8HfQuantizer.update_torch_dtypeO   s$   € ØÐÜ�K‰KÐqÔrÜŸ-™-ˆKØÐr   Úmodelr   Úparam_valueztorch.TensorÚ
param_nameÚtarget_deviceztorch.deviceÚ
state_dictÚunexpected_keysc                 óV  — ddl m}  |||||«       t        ||«      \  }}	t        j                  t        j
                  «      j                  }
t        j                  t        j
                  «      j                  }| j                  j                  \  }}|j                  dd \  }}||z  dk7  s||z  dk7  rt        d|› d|› d|› d|› d�	«      ‚|j                  }|j                  d	||z  |||z  |«      j                  dd
ddd«      }t        j                  t        j                  |«      d¬«      }||z  }|j                  }|j!                  d	«      j!                  d	«      }t        j"                  ||z  |
|¬«      j%                  t        j
                  «      }|j                  dd
ddd«      }|j                  |«      }|j                  |«      j'                  «       j)                  «       }|j%                  |«      |j*                  |	<   |j%                  |«      |j*                  d<   y)zO
        Quantizes weights to FP8 format using Block-wise quantization
        r   )Úset_module_tensor_to_deviceéþÿÿÿNzMatrix dimensions (z, z$) must be divisible by block sizes (ú)éÿÿÿÿr   é   r	   é   )rN   rL   )Údim)ÚminÚmaxÚweight_scale_inv)Úaccelerate.utilsrK   r   r   ÚfinfoÚfloat8_e4m3fnrR   rS   r   Úweight_block_sizeÚshaper.   ÚreshapeÚpermuteÚamaxÚabsÚ	unsqueezeÚclampÚtoÚsqueezeÚ
reciprocalÚ_buffers)r   rD   rE   rF   rG   rH   rI   rK   ÚmoduleÚtensor_nameÚfp8_minÚfp8_maxÚblock_size_mÚblock_size_nÚrowsÚcolsÚparam_value_orig_shapeÚmax_absÚscaleÚscale_orig_shapeÚquantized_params                        r   Úcreate_quantized_paramz0FineGrainedFP8HfQuantizer.create_quantized_paramU   s  € õ 	Aá# E¨:°}ÀkÔRä2°5¸*ÓEÑˆ�ô —+‘+œe×1Ñ1Ó2×6Ñ6ˆÜ—+‘+œe×1Ñ1Ó2×6Ñ6ˆà%)×%=Ñ%=×%OÑ%OÑ"ˆ�là ×&Ñ& r sÐ+‰
ˆˆdà�,Ñ !Ò# t¨lÑ':¸aÒ'?ÜØ% d V¨2¨d¨VÐ3WÐXdÐWeÐegÐhtÐguÐuvÐwóð ð "-×!2Ñ!2Ðà!×)Ñ)Ø�˜Ñ$ l°D¸LÑ4HÈ,ó
ç
‰'�!�Q˜˜1˜aÓ
 ð 	ô
 —*‘*œUŸY™Y {Ó3¸ÔBˆØ˜'Ñ!ˆØ Ÿ;™;ÐØ—‘ Ó#×-Ñ-¨bÓ1ˆô  Ÿ+™+ k°EÑ&9¸wÈGÔT×WÑWÔX]×XkÑXkÓlˆà)×1Ñ1°!°Q¸¸1¸aÓ@ˆà)×1Ñ1Ð2HÓIˆð —‘Ð.Ó/×7Ñ7Ó9×DÑDÓFˆà'6×'9Ñ'9¸-Ó'Hˆ�‰˜Ñ$Ø.3¯h©h°}Ó.Eˆ�‰Ð*Ò+r   c                 óæ   — ddl m} t        ||«      \  }}t        ||«      rP| j                  s|dk(  r.|dk(  r(|j
                  t        j                  k7  rt        d«      ‚y|dk(  rt        d«      ‚y	y)
Nr	   ©Ú	FP8LinearÚbiasÚweightz6Expect quantized weights but got an unquantized weightFrT   z;Expect unquantized weights but got a quantized weight_scaleT)	Úintegrations.finegrained_fp8rt   r   r6   r5   Údtyper   rW   r.   )	r   rD   rE   rF   rH   r   rt   rd   re   s	            r   Úcheck_quantized_paramz/FineGrainedFP8HfQuantizer.check_quantized_paramŒ   sw   € õ 	=ä2°5¸*ÓEÑˆ�ä�f˜iÔ(Ø×!Ò! [°FÒ%:Ø (Ò*¨{×/@Ñ/@ÄE×DWÑDWÒ/WÜ$Ð%]Ó^Ð^ØàÐ"4Ò4Ü$Ð%bÓcÐcØØr   Úkeep_in_fp32_modulesc                 óÜ   — ddl m} | j                  || j                  j                  |«      | _         ||| j                  | j                  ¬«      }| j                  |j
                  _        y )Nr	   )Úreplace_with_fp8_linear)Úmodules_to_not_convertr   )rw   r|   Úget_modules_to_not_convertr   r}   Úconfig)r   rD   rz   r   r|   s        r   Ú$_process_model_before_weight_loadingz>FineGrainedFP8HfQuantizer._process_model_before_weight_loading£   sd   € õ 	Kà&*×&EÑ&EØ�4×+Ñ+×BÑBÐDXó'
ˆÔ#ñ (ØØ#'×#>Ñ#>Ø $× 8Ñ 8ô
ˆð ,0×+CÑ+Cˆ�‰Õ(r   c                 ó   — |S r   © )r   rD   r   s      r   Ú#_process_model_after_weight_loadingz=FineGrainedFP8HfQuantizer._process_model_after_weight_loading·   s   € Øˆ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	   rs   r$   z.weightz.bias)Úintegrationsrt   Únamed_modulesr6   ÚendswithÚappend)
r   rD   r„   r…   rt   Únot_missing_keysÚnamerd   ÚmissingÚks
             r   Úupdate_missing_keysz-FineGrainedFP8HfQuantizer.update_missing_keysº   s¡   € Ý,àÐØ!×/Ñ/Ó1ò 	9‰LˆD�&Ü˜& )Õ,Ø+ò 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                  ó   — y)NTr‚   )r   Úsafe_serializations     r   Úis_serializablez)FineGrainedFP8HfQuantizer.is_serializableÉ   s   € Ør   c                  ó   — y)NFr‚   )r   s    r   Úis_trainablez&FineGrainedFP8HfQuantizer.is_trainableÌ   s   € àr   )rB   útorch.dtyper>   r•   r   )rD   r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__Ú requires_parameters_quantizationÚrequires_calibrationÚrequired_packagesr   r=   rC   Ústrr   r   r   r   rq   ry   r€   rƒ   r�   r’   ÚpropertyÚboolr”   Ú__classcell__)r   s   @r   r   r      s4  ø„ ñð
 (,Ð$Ø ÐØ%˜Ðô7ò+óZð 04ñ5Fà ð5Fð $ð5Fð ð	5Fð
 &ð5Fð ˜˜c˜‘Nð5Fð " $ s¡)Ñ,ó5Fðnà ðð $ðð ð	ð
 ˜˜c˜‘Nóð4 59ñDà ðDð ' t¨C¡yÑ1óDó(ðF°t¸C±yð FÈ#ð FÐRVÐWZÑR[ó Fóð ð˜dò ó ôr   r   )r*   Útypingr   r   r   r   r   Ú	packagingr   Úutilsr
   r   r   Úbaser   Úquantizers_utilsr   r   Úmodeling_utilsr   Ú
get_loggerr–   r3   r   r‚   r   r   ú<module>r¨      sN   ðÛ ß ;Õ ;å ç HÑ HÝ Ý 2ñ ÔÛáÝ0à	ˆ×	Ñ	˜HÓ	%€ôz õ zr   