Ë
    T^(h¿/  ã                   ó  — d dl Z d dlZd dlZd dlmZmZmZ d dlmZ ddl	m
Z
 ddlmZ erddlmZ d d	lmZmZmZ dd
lmZmZmZ ddlmZ  e«       r
d dlZd dlmZ  ej6                  e«      Zdedee   fd„Zd„ Z d„ Z!d„ Z" G d„ de
«      Z#y)é    N)ÚTYPE_CHECKINGÚOptionalÚUnion)Úversioné   )ÚHfQuantizer)Úget_module_from_nameé   )ÚPreTrainedModel)ÚAnyÚDictÚList)Úis_torch_availableÚis_torchao_availableÚlogging)ÚTorchAoConfigÚconfig_nameÚreturnc                 óv   — | j                  «       } t        j                  d| «      }|r|j                  d«      S y)z†
    Extract the size digit from strings like "4weight", "8weight".
    Returns the digit as an integer if found, otherwise None.
    z
(\d)weightr   N)ÚlowerÚreÚsearchÚgroup)r   Ú	str_matchs     úg/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/quantizers/quantizer_torchao.pyÚfuzzy_match_sizer   )   s7   € ð
 ×#Ñ#Ó%€Kä—	‘	˜-¨Ó5€IáØ�‰˜qÓ!Ð!àó    c                 ó^   — |j                  d«      d d }| }|D ]  }|j                  |   }Œ |S )Nú.éÿÿÿÿ)ÚsplitÚ_modules)ÚmodelÚnameÚmodule_treeÚparentÚms        r   Úfind_parentr(   9   s=   € Ø—*‘*˜S“/ # 2Ð&€KØ€FØò $ˆØ—‘ Ñ#‰ð$à€Mr   c                 ó  — ddl m} ddlm} t	        | |«      r*| j
                  j                  › d| j                  «       › d�S t	        | |«      r<| j
                  j                  › d| j                  › dt        | j                  «      › d�S y )Nr   )ÚAffineQuantizedTensor)ÚLinearActivationQuantizedTensorú(ú)z(activation=ú	, weight=)
Útorchao.dtypesr*   Ú7torchao.quantization.linear_activation_quantized_tensorr+   Ú
isinstanceÚ	__class__Ú__name__Ú_quantization_typeÚinput_quant_funcÚoriginal_weight_tensor)Úweightr*   r+   s      r   r4   r4   A   s¤   € Ý4Ýgä�&Ð/Ô0Ø×"Ñ"×+Ñ+Ð,¨A¨f×.GÑ.GÓ.IÐ-JÈ!ÐLÐLä�&Ð9Ô:Ø×"Ñ"×+Ñ+Ð,¨L¸×9PÑ9PÐ8QÐQZÔ[mÐnt÷  oLñ  oLó  \Mð  [Nð  NOð  Pð  	Pð ;r   c                 ó  — t        | j                  «      }|€7d| j                  j                  d   › d| j                  j                  d   › d�S d| j                  j                  d   › d| j                  j                  d   › d|› �S )Nzin_features=r   z, out_features=r   z, weight=Noner.   )r4   r7   Úshape)Úselfr7   s     r   Ú_linear_extra_reprr;   L   s‰   € Ü §¡Ó,€FØ€~Ø˜dŸk™k×/Ñ/°Ñ2Ð3°?À4Ç;Á;×CTÑCTÐUVÑCWÐBXÐXeÐfÐfà˜dŸk™k×/Ñ/°Ñ2Ð3°?À4Ç;Á;×CTÑCTÐUVÑCWÐBXÐXaÐbhÐaiÐjÐjr   c                   ó2  ‡ — e Zd ZdZdZdZdgZˆ fd„Zd„ Zd„ Z	dd	„Z
d
eeeeef   f   deeeeef   f   fd„Z	 ddddeee      fd„Zdddddedeeef   def
d„Zdddddedddeeef   dee   fd„Zd„ Zddefd„Zedefd„«       Zedefd„«       Zˆ xZS )ÚTorchAoHfQuantizerz?
    Quantizer for torchao: https://github.com/pytorch/ao/
    TFÚtorchaoc                 ó&   •— t        ‰| �  |fi |¤Ž y ©N)ÚsuperÚ__init__)r:   Úquantization_configÚkwargsr2   s      €r   rB   zTorchAoHfQuantizer.__init__]   s   ø€ Ü‰ÑÐ,Ñ7°Ó7r   c                 óú  — t        «       st        d«      ‚d| _        |j                  dd «      }t	        |t
        «      rBd|j                  «       v sd|j                  «       v r| j                  rt        d«      ‚d| _        | j                  ro|j                  dd «      }|rZt        j                  t        j                  j                  d	«      «      }|t        j                  d
«      k  rt        d|› d�«      ‚y y y )NzSLoading an torchao quantized model requires torchao library (`pip install torchao`)FÚ
device_mapÚcpuÚdiskz§You are attempting to perform cpu/disk offload with a pre-quantized torchao model This is not supported yet . Please remove the CPU or disk device from the device_map.TÚweights_onlyÚtorchz2.5.0zlIn order to use torchao pre-quantized model, you need to have torch>=2.5.0. However, the current version is zc. You can also set with `weights_only=False` in `from_pretrained` if you don't want to update torch)r   ÚImportErrorÚoffloadÚgetr1   ÚdictÚvaluesÚpre_quantizedÚ
ValueErrorr   ÚparseÚ	importlibÚmetadataÚRuntimeError)r:   ÚargsrD   rF   rI   Útorch_versions         r   Úvalidate_environmentz'TorchAoHfQuantizer.validate_environment`   s  € Ü#Ô%ÜÐsÓtÐtàˆŒØ—Z‘Z ¨dÓ3ˆ
Ü�j¤$Ô'Ø˜
×)Ñ)Ó+Ñ+¨v¸×9JÑ9JÓ9LÑ/LØ×%Ò%Ü$ðpóð ð
 $(�D”LØ×ÒØ!Ÿ:™: n°dÓ;ˆLÙÜ '§¡¬i×.@Ñ.@×.HÑ.HÈÓ.QÓ R�Ø ¤7§=¡=°Ó#9Ò9Ü&ð Gð  HUð  GVð V}ð ~óð ð :ð ð r   c                 ób  — | j                   j                  dk(  rU|�,|t        j                  k7  rt        j                  d|› d�«       |€%t        j                  d«       t        j                  }| j                   j                  dk(  r'|€%t        j                  d«       t        j                  }|S )NÚint4_weight_onlyzSetting torch_dtype to zu for int4_weight_only quantization, but only bfloat16 is supported right now. Please set the torch_dtype to bfloat16.z±Setting torch_dtype to torch.bfloat16 for int4_weight_only quantization since only bfloat16 is supported right now. Please set torch_dtype=torch.bfloat16 to remove this warning.Ú#int8_dynamic_activation_int8_weightzŒSetting torch_dtype to torch.float32 for int8_dynamic_activation_int8_weight quantization as no torch_dtype was specified in from_pretrained)rC   Ú
quant_typerJ   Úbfloat16ÚloggerÚwarning_onceÚinfoÚfloat32)r:   Útorch_dtypes     r   Úupdate_torch_dtypez%TorchAoHfQuantizer.update_torch_dtypey   s®   € Ø×#Ñ#×.Ñ.Ð2DÒDØÐ&¨;¼%¿.¹.Ò+HÜ×#Ñ#Ø-¨k¨]ð  ;pð  qôð Ð"Ü×#Ñ#ð Hôô $Ÿn™n�Ø×#Ñ#×.Ñ.Ð2WÒWØÐ"Ü—‘ð côô $Ÿm™m�ØÐr   r   c                 ót  — t        j                  t        j                  j                  d«      «      t        j                  d«      kD  ræddlm} | j                  j                  «       t        j                  d«      kD  rjddl	m
} | j                  j                  }t        ||«      rB|j                  j                  }t        |«      }|dk(  r|j                   S t"        j$                  S |j                   t"        j$                  t"        j$                  d dœ}|| j                  j                     S t'        d	«      ‚)
NÚ
acceleratez0.19.0r   )ÚCustomDtypez0.9.0)ÚAOBaseConfigÚ4)rZ   Úint8_weight_onlyr[   Ú	autoquantzÉYou are using `device_map='auto'` on a torchao quantized model. To automatically compute the appropriate device map, you should upgrade your `accelerate` library with `pip install --upgrade accelerate`)r   rR   rS   rT   Úaccelerate.utilsrf   rC   Ú_get_ao_versionÚVersionÚtorchao.core.configrg   r\   r1   r2   r3   r   ÚINT4rJ   Úint8rQ   )r:   rb   rf   rg   r\   r   Ú
size_digitÚmap_to_target_dtypes           r   Úadjust_target_dtypez&TorchAoHfQuantizer.adjust_target_dtype�   sý   € Ü�=‰=œ×+Ñ+×3Ñ3°LÓAÓBÄWÇ]Á]ÐS[ÓE\Ò\Ý4ð ×'Ñ'×7Ñ7Ó9¼G¿O¹OÈGÓ<TÒTÝ<à!×5Ñ5×@Ñ@�
Ü˜j¨,Ô7à",×"6Ñ"6×"?Ñ"?�KÜ!1°+Ó!>�Jð " SÒ(Ø*×/Ñ/Ð/ô  %Ÿz™zÐ)ð %0×$4Ñ$4Ü$)§J¡JÜ7<·z±zØ!ñ	#Ðð ' t×'?Ñ'?×'JÑ'JÑKÐKäð5óð r   Ú
max_memoryc                 ó^   — |j                  «       D ��ci c]  \  }}||dz  “Œ }}}|S c c}}w )NgÍÌÌÌÌÌì?)Úitems)r:   rt   ÚkeyÚvals       r   Úadjust_max_memoryz$TorchAoHfQuantizer.adjust_max_memory±   s6   € à5?×5EÑ5EÓ5G×H©¨¨c�c˜3 ™9‘nÐHˆ
ÑHØÐùó Is   ”)r#   r   Úkeep_in_fp32_modulesc                 ó\   — | j                  || j                  j                  |«      | _        y r@   )Úget_modules_to_not_convertrC   Úmodules_to_not_convert)r:   r#   rz   rD   s       r   Ú$_process_model_before_weight_loadingz7TorchAoHfQuantizer._process_model_before_weight_loading¶   s0   € ð '+×&EÑ&EØ�4×+Ñ+×BÑBÐDXó'
ˆÔ#ð 	r   Úparam_valueztorch.TensorÚ
param_nameÚ
state_dictc                 ó2  ‡— | j                   j                  dk(  ry|j                  dd «      }t        ˆfd„| j                  D «       «      ry|dk(  r| j
                  ryt        |‰«      \  }}t        |t        j                  j                  «      xr |dk(  S )Nrj   FÚparam_devicec              3   ó:   •K  — | ]  }|d z   ‰v xs |‰k(  –— Œ y­w)r   N© )Ú.0rw   r€   s     €r   ú	<genexpr>z;TorchAoHfQuantizer.check_quantized_param.<locals>.<genexpr>Ë   s'   øè ø€ ÒgÀC��c‘	˜ZÐ'Ò?¨S°JÑ->Ó?Ñgùs   ƒrG   r7   )rC   r\   ÚpopÚanyr}   rL   r	   r1   rJ   ÚnnÚLinear)	r:   r#   r   r€   r�   rD   rƒ   ÚmoduleÚtensor_names	      `     r   Úcheck_quantized_paramz(TorchAoHfQuantizer.check_quantized_param¾   s†   ø€ ð ×#Ñ#×.Ñ.°+Ò=Øà—z‘z .°$Ó7ˆäÓgÈ4×KfÑKfÔgÔgØØ˜UÒ" t§|¢|àô #7°u¸jÓ"IÑˆF�KÜ˜f¤e§h¡h§o¡oÓ6ÒT¸KÈ8Ñ<SÐTr   Útarget_deviceztorch.deviceÚunexpected_keysc                 óZ  — | j                   j                  dk(  ryddlm} t	        ||«      \  }}	| j
                  rwt        j                  j                  |j                  |¬«      «      |j                  |	<   t        |t        j                  «      r t        j                  t        |«      |_        yyt        | j                   t"        «      sJ ‚t        j                  j                  |«      j                  |¬«      |j                  |	<    ||| j                   j%                  «       «       y)zÏ
        Each nn.Linear layer that needs to be quantized is processsed here.
        First, we set the value the weight tensor, then we move it to the target device. Finally, we quantize the module.
        rj   Nr   )Ú	quantize_)Údevice)rC   r\   Útorchao.quantizationr’   r	   rP   rJ   rŠ   Ú	ParameterÚtoÚ_parametersr1   r‹   ÚtypesÚ
MethodTyper;   Ú
extra_reprr   Úget_apply_tensor_subclass)
r:   r#   r   r€   r�   r�   r�   r’   rŒ   r�   s
             r   Úcreate_quantized_paramz)TorchAoHfQuantizer.create_quantized_paramÕ   së   € ð ×#Ñ#×.Ñ.°+Ò=Øå2ä2°5¸*ÓEÑˆ�Ø×ÒÜ.3¯h©h×.@Ñ.@ÀÇÁÐWdÀÓAeÓ.fˆF×Ñ˜{Ñ+Ü˜&¤"§)¡)Ô,Ü$)×$4Ñ$4Ô5GÈÓ$P�Õ!ð -ô ˜d×6Ñ6¼ÔFÐFÐFÜ.3¯h©h×.@Ñ.@ÀÓ.M×.PÑ.PÐXeÐ.PÓ.fˆF×Ñ˜{Ñ+Ù�f˜d×6Ñ6×PÑPÓRÕSr   c                 óÀ   — | j                   j                  dk(  rEddlm} ddlm} t        j                  |d¬«      } ||f|ddœ| j                   j                  ¤Ž}|S y	)
z/No process required for torchao quantized modelrj   r   )rj   )ÚALL_AUTOQUANT_CLASS_LISTzmax-autotune)ÚmodeF)Úqtensor_class_listÚset_inductor_configN)	rC   r\   r>   rj   r”   rž   rJ   ÚcompileÚquant_type_kwargs)r:   r#   rD   rj   rž   s        r   Ú#_process_model_after_weight_loadingz6TorchAoHfQuantizer._process_model_after_weight_loadingñ   sd   € à×#Ñ#×.Ñ.°+Ò=Ý)ÝEä—M‘M %¨nÔ=ˆEÙØðà#;Ø$)ñð ×*Ñ*×<Ñ<ñ	ˆEð ˆLØr   c                 ód  — |rt         j                  d«       yt        j                  t        j
                  j                  d«      «      t        j                  d«      k\  }|st         j                  d«       | j                  r,| j                  j                  €t         j                  d«       y|S )Nzetorchao quantized model does not support safe serialization, please set `safe_serialization` to FalseFÚhuggingface_hubz0.25.0zMtorchao quantized model is only serializable after huggingface_hub >= 0.25.0 a  The model contains offloaded modules and these modules are not quantized. We don't recommend saving the model as we won't be able to reload them.If you want to specify modules to not quantize, please specify modules_to_not_convert in the quantization_config.)	r^   Úwarningr   rR   rS   rT   rL   rC   r}   )r:   Úsafe_serializationÚ_is_torchao_serializables      r   Úis_serializablez"TorchAoHfQuantizer.is_serializable  s˜   € ÙÜ�N‰NØwôð Ü#*§=¡=´×1CÑ1C×1KÑ1KÐL]Ó1^Ó#_Ôcj×cpÑcpØód
ñ $
Ð ñ (Ü�N‰NÐjÔkØ�<Š<˜D×4Ñ4×KÑKÐSÜ�N‰NðDôð Ø'Ð'r   c                 ó:   — ddg}| j                   j                  |v S )Nri   r[   )rC   r\   )r:   Ú"supported_quant_types_for_trainings     r   Úis_trainablezTorchAoHfQuantizer.is_trainable  s,   € ð Ø1ð.
Ð*ð ×'Ñ'×2Ñ2Ð6XÐXÐXr   c                  ó   — y)NTr…   )r:   s    r   Úis_compileablez!TorchAoHfQuantizer.is_compileable  s   € àr   )rb   útorch.dtyper   r°   r@   )r3   Ú
__module__Ú__qualname__Ú__doc__Ú requires_parameters_quantizationÚrequires_calibrationÚrequired_packagesrB   rX   rc   rs   r   Ústrr   Úintry   r   r   r~   r   ÚboolrŽ   rœ   r¤   rª   Úpropertyr­   r¯   Ú__classcell__)r2   s   @r   r=   r=   T   st  ø„ ñð (,Ð$Ø ÐØ"˜Ðô8òò2ó("ðH¨D°°e¸CÀ¸H±oÐ1EÑ,Fð È4ÐPSÐUZÐ[^Ð`cÐ[cÑUdÐPdÑKeó ð UYñØ&ðØ>FÀtÈCÁyÑ>QóðUà ðUð $ðUð ð	Uð
 ˜˜c˜‘NðUð 
óUð.Tà ðTð $ðTð ð	Tð
 &ðTð ˜˜c˜‘NðTð ˜c™óTò8ñ (¸$ó (ð& ðY˜dò Yó ðYð ð ò ó ôr   r=   )$rS   r   r˜   Útypingr   r   r   Ú	packagingr   Úbaser   Úquantizers_utilsr	   Úmodeling_utilsr   r   r   r   Úutilsr   r   r   Úutils.quantization_configr   rJ   Útorch.nnrŠ   Ú
get_loggerr3   r^   r·   r   r(   r4   r;   r=   r…   r   r   ú<module>rÅ      sŒ   ðó Û 	Û ß 1Ñ 1å å Ý 2ñ Ý0ç "Ñ "ç EÑ EÝ 5ñ ÔÛÝà	ˆ×	Ñ	˜HÓ	%€ð #ð ¨(°3©-ó ò òPòkôJ˜õ Jr   