Ë
    T^(h©0  ã                   óØ   — d dl mZmZmZmZ ddlmZ ddlmZm	Z	m
Z
mZ ddlmZ ddlmZ erddlmZ  e«       rd d	lmZ  e
«       rd d
lZ ej*                  e«      Zd„ Z G d„ de«      Zy
)é    )ÚTYPE_CHECKINGÚAnyÚDictÚListé   )Úprepare_for_hqq_linear)Úis_accelerate_availableÚis_hqq_availableÚis_torch_availableÚloggingé   )ÚHfQuantizer)Úget_module_from_name)ÚPreTrainedModel)Úremove_hook_from_moduleNc                 ó^   — |j                  d«      d d }| }|D ]  }|j                  |   }Œ |S )Nú.éÿÿÿÿ)ÚsplitÚ_modules)ÚmodelÚnameÚmodule_treeÚparentÚms        úc/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/quantizers/quantizer_hqq.pyÚfind_parentr   %   s=   € Ø—*‘*˜S“/ # 2Ð&€KØ€FØò $ˆØ—‘ Ñ#‰ð$à€Mó    c                   ó  ‡ — e Zd ZdZdZdZdZdgZˆ fd„Zd„ Z	ddd	e
e   d
ede
e   fd„Zddde
e   de
e   d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„Zdd„Zdd„Zedefd„«       Zˆ xZS ) ÚHqqHfQuantizerzä
    HQQ quantizer base HF class.
    nn.Linear modules are first tagged with quant_config in _process_model_before_weight_loading().
    The actual quantization and offloading to the GPU is done in check_quantized_param().
    FTÚhqqc                 óB   •— t        ‰| �  |fi |¤Ž d | _        d| _        y )NF)ÚsuperÚ__init__Útorch_dtypeÚusing_multi_gpu)ÚselfÚquantization_configÚkwargsÚ	__class__s      €r   r$   zHqqHfQuantizer.__init__9   s&   ø€ Ü‰ÑÐ,Ñ7°Ò7ØˆÔØ$ˆÕr   c                 ó`  — t        «       st        d«      ‚|j                  dd«      s|j                  dd«      rt        d«      ‚t        j
                  j                  «       st        d«      ‚| j                  €9d|v r|d   | _        n*t        j                  | _        t        j                  d«       |j                  d	d «      }t        |t        «      rZd
|j                  «       v sd|j                  «       v rt        d«      ‚t        t!        |j                  «       «      «      dkD  | _        y y )Nz�A valid HQQ version (>=0.2.1) is not available. Please follow the instructions to install it: `https://github.com/mobiusml/hqq/`.Úfrom_tfFÚ	from_flaxzwConverting weights from tf/flax weights is currently not supported, please make sure the weights are in PyTorch format.z/No GPU found. A GPU is needed for quantization.r%   zUSetting torch_dtype to torch.float32 as the default value since it was not specified.Ú
device_mapÚcpuÚdiskz­You are attempting to use an HQQ model with a device_map that contains a CPU or disk device. This is not supported. Please remove the CPU or disk device from the device_map.r   )r
   ÚImportErrorÚgetÚ
ValueErrorÚtorchÚcudaÚis_availableÚRuntimeErrorr%   Úfloat32ÚloggerÚinfoÚ
isinstanceÚdictÚvaluesÚlenÚsetr&   )r'   Úargsr)   r.   s       r   Úvalidate_environmentz#HqqHfQuantizer.validate_environment>   s  € Ü Ô"Üð Tóð ð �:‰:�i Ô'¨6¯:©:°kÀ5Ô+IÜð;óð ô
 �z‰z×&Ñ&Ô(ÜÐPÓQÐQà×ÑÐ#Ø Ñ&Ø#)¨-Ñ#8�Õ ä#(§=¡=�Ô Ü—‘ÐsÔtà—Z‘Z ¨dÓ3ˆ
Ü�j¤$Ô'Ø˜
×)Ñ)Ó+Ñ+¨v¸×9JÑ9JÓ9LÑ/LÜ ðhóð ô
 (+¬3¨z×/@Ñ/@Ó/BÓ+CÓ'DÀqÑ'H�Õ$ð (r   r   r   Úmissing_keysÚprefixÚreturnc                 óR   — | j                   r|D �cg c]	  }d|vsŒ|‘Œ c}S |S c c}w )NÚweight)Úpre_quantized)r'   r   rB   rC   r)   Úkeys         r   Úupdate_missing_keysz"HqqHfQuantizer.update_missing_keys^   s1   € ð ×ÒØ#/ÖI˜C°HÀCÒ4G’CÒIÐIàÐùò Js   ‘	$›$Úexpected_keysÚloaded_keysc                 ó  ‡‡— | j                   s|S ˆfd„Št        |«      }t        «       �rNddlm} |j                  «       D ]  \  }}||_        Œ t        «       } ‰||«       t        «       }	|D ]6  }
|j                  j                  d   D ]  }||
v sŒ|	j                  |
«       Œ Œ8 ||	z  } |d d t        j                  d¬«      j                  «       dhz
  }t        «       }|D ](  Št        ˆfd„|D «       «      sŒ|j                  ‰«       Œ* ||z  }|D ]_  }
|
d	z   |v r|j                  |
d	z   «       n%|j                  |D �ch c]
  }|
d
z   |z   ’Œ c}«       |
dz   |v sŒL|j                  |
dz   «       Œa t        |«      S c c}w )Nc                 óÆ   •— | j                  «       D ]M  \  }}t        |t        j                  j                  «      r|j                  |j                  «        ‰||«       ŒO y ©N)Únamed_childrenr;   r4   ÚnnÚLinearÚaddr   )r   Úlayersr   ÚmoduleÚ_find_hqq_quantizable_layerss       €r   rU   zIHqqHfQuantizer.update_expected_keys.<locals>._find_hqq_quantizable_layersn   sK   ø€ Ø %× 4Ñ 4Ó 6ò =‘��fÜ˜f¤u§x¡x§¡Ô8Ø—J‘J˜vŸ{™{Ô+Ù,¨V°VÕ<ñ=r   r   ©Ú	HQQLinearÚskip_modulesr/   ©Úlinear_layerÚquant_configÚcompute_dtypeÚdeviceÚbiasc              3   ó&   •K  — | ]  }|‰v –— Œ
 y ­wrN   © )Ú.0Ú_modulerH   s     €r   ú	<genexpr>z6HqqHfQuantizer.update_expected_keys.<locals>.<genexpr>�   s   øè ø€ ÒD¨'�w #”~ÑDùs   ƒz.weightr   z.bias)rG   r?   r
   Úhqq.core.quantizerW   Únamed_modulesr   Úconfigr(   rR   r4   Úfloat16Ústate_dict_keysÚanyÚupdateÚlist)r'   r   rJ   rK   Únew_keysrW   r   rT   Ú_valid_modulesÚ_skipped_modulesrb   Ú_skip_moduleÚ	_ref_keysÚ_rm_keysÚ_ref_keyrU   rH   s                  @@r   Úupdate_expected_keysz#HqqHfQuantizer.update_expected_keysg   s¨  ù€ ð ×!Ò!Ø Ð ô	=ô �}Ó%ˆÜÕÝ3ð !&× 3Ñ 3Ó 5ò #‘��fØ"�•ð#ô !›UˆNÙ(¨°Ô?ô  #›uÐØ)ò 6�Ø$)§L¡L×$DÑ$DÀ^Ñ$Tò 6�LØ# wÒ.Ø(×,Ñ,¨WÕ5ñ6ð6ð Ð.Ñ.ˆNñ "Ø!°ÄEÇMÁMÐZ_ôç‰oÓ 6 (ñ+ˆIô
 “uˆHØò &�ÜÓD°^ÔDÕDØ—L‘L Õ%ð&ð ˜Ñ ˆHð *ò 4�Ø˜YÑ&¨+Ñ5Ø—L‘L ¨9Ñ!4Õ5à—O‘OÈiÖ$XÀ( W¨s¡]°XÓ%=Ò$XÔYØ˜WÑ$¨Ò3Ø—L‘L ¨7Ñ!2Õ3ð4ô �H‹~Ðùò	 %Ys   ÅF
Úparam_valueztorch.TensorÚ
param_nameÚ
state_dictc                 óX  — t        «       rddlm} t        ||«      \  }}| j                  r@t        |t        j                  j                  «      xs t        |«      xr |dk7  xr |dk7  S t        |t        j                  j                  «      xr |dk(  xs t        |«      xr |dk(  S )Nr   rV   rF   r^   )	r
   rd   rW   r   rG   r;   r4   rP   rQ   )	r'   r   rt   ru   rv   r)   rW   rT   Útensor_names	            r   Úcheck_quantized_paramz$HqqHfQuantizer.check_quantized_param    s¤   € ô ÔÝ3Ü2°5¸*ÓEÑˆ�à×Òä˜F¤E§H¡H§O¡OÓ4ÒU¼
À6È9Ó8Uò *Ø 8Ñ+ò*à 6Ñ)ðô ˜6¤5§8¡8§?¡?Ó3ò ,Ø 8Ñ+òMä˜v yÓ1ÒK°kÀVÑ6Kðr   Útarget_deviceztorch.deviceÚunexpected_keysc           	      óœ  — t        «       rddlm} t        ||«      \  }}	dj	                  |j                  d«      dd «      }
t        ||
«      }|
j                  d«      d   }|	dk(  ryi }|j                  «       D ]=  \  }}|
dz   |v sŒ|||j                  d«      d   <   |€Œ(||v sŒ-|j                  |«       Œ? | j                  rÞt        |«      ry |dd| j                  |¬«      }|j                  |«       |j                  �Rt        |j                  t        j                  «      r.t        j                   j#                  |j                  «      |_        | j$                  r| j'                  |«      }t)        |||«       |`~t        j,                  j/                  «        y|D ]/  }t)        ||t        j                   j#                  ||   «      «       Œ1 |j0                  j2                  d   }|j0                  j2                  d	   }dj	                  |j4                  j                  d«      d
d «      }d}d|v r|}n	||v r||   }|D ]  }||j4                  v sŒd} n |�  ||| j                  |d¬«      }|j                  �Rt        |j                  t        j                  «      r.t        j                   j#                  |j                  «      |_        | j$                  r| j'                  |«      }t)        |||«       n*|j7                  | j                  |¬«      }t)        |||«       t        j,                  j/                  «        y)a  
        Each nn.Linear layer is processed here.
        We first check if the corresponding module state_dict contains already HQQ quantized parameters.
        If not, we create a temp linear layer with the module state_dict params and use it for quantization
        r   rV   r   Nr   r^   rY   r[   rX   éþÿÿÿÚweight_quant_paramsT)r[   r\   r]   Údel_orig)Údtyper]   )r
   rd   rW   r   Újoinr   r   ÚitemsÚremoverG   r;   r%   Úload_state_dictr^   r4   ÚTensorrP   Ú	Parameterr&   Ú_patch_layer_for_multigpuÚsetattrÚ__dict__r5   Úempty_cacherf   r(   r   Úto)r'   r   rt   ru   rz   rv   r{   rW   rT   rx   Ú
layer_nameÚparent_moduleÚnodeÚmodule_state_dictÚkÚvÚ	hqq_layerrH   r[   rX   Ú
module_tagÚmodule_quant_configÚskip_modules                          r   Úcreate_quantized_paramz%HqqHfQuantizer.create_quantized_paramº   s  € ô ÔÝ3ä2°5¸*ÓEÑˆ�Ø—X‘X˜j×.Ñ.¨sÓ3°C°RÐ8Ó9ˆ
Ü# E¨:Ó6ˆØ×Ñ Ó$ RÑ(ˆà˜&Ò àð ÐØ×$Ñ$Ó&ò 	.‰DˆAˆqØ˜CÑ 1Ò$Ø67Ð! !§'¡'¨#£,¨rÑ"2Ñ3Ø"Ñ.°1¸Ò3GØ#×*Ñ*¨1Õ-ð		.ð ×ÒÜ˜& )Ô,Øá%Ø!%Ø!%Ø"&×"2Ñ"2Ø(ô	�	ð ×%Ñ%Ð&7Ô8à�~‰~Ð)¬j¸¿¹ÌÏÉÔ.VÜ!&§¡×!3Ñ!3°I·N±NÓ!C�	”à×#Ò#Ø ×:Ñ:¸9ÓE�	ä�M 4¨Ô3ð � Ü�J‰J×"Ñ"Ô$Øð %ò 	MˆCÜ�F˜C¤§¡×!3Ñ!3Ð4EÀcÑ4JÓ!KÕLð	Mð
 —|‘|×7Ñ7¸ÑGˆØ—|‘|×7Ñ7¸ÑGˆØ—X‘X˜fŸk™k×/Ñ/°Ó4°R°SÐ9Ó:ˆ
Ø"ÐØ  LÑ0Ø".ÑØ˜<Ñ'Ø".¨zÑ":Ðà'ò 	ˆKØ˜fŸk™kÒ)Ø&*Ð#Ùð	ð
 Ð*Ù!ØØ0Ø"×.Ñ.Ø$ØôˆIð �~‰~Ð)¬j¸¿¹ÌÏÉÔ.VÜ!&§¡×!3Ñ!3°I·N±NÓ!C�	”à×#Ò#Ø ×:Ñ:¸9ÓE�	ä�M 4¨Õ3ð —Y‘Y T×%5Ñ%5¸m�YÓLˆFÜ�M 4¨Ô0ä�
‰
×ÑÕ r   c                 ó<   ‡‡— t        ‰«      Šd„ Šˆˆfd„‰_        ‰S )Nc                 óÒ   — t        j                  |j                  | j                  «      | j	                  «       j                  «       «      }| j                  �|| j                  z  }|S rN   )r4   Úmatmulr‹   r]   Ú
dequantizeÚtr^   )r'   ÚxÚouts      r   Úforward_with_devicezEHqqHfQuantizer._patch_layer_for_multigpu.<locals>.forward_with_device&  sL   € Ü—,‘,˜qŸt™t D§K¡KÓ0°$·/±/Ó2C×2EÑ2EÓ2GÓHˆCØ�y‰yÐ$Ø�t—y‘yÑ �ØˆJr   c                 ó   •—  ‰‰| «      S rN   r`   )rœ   rž   r’   s    €€r   ú<lambda>z:HqqHfQuantizer._patch_layer_for_multigpu.<locals>.<lambda>,  s   ø€ Ñ&9¸)ÀQÓ&G€ r   )r   Úforward)r'   r’   rž   s    `@r   r‡   z(HqqHfQuantizer._patch_layer_for_multigpu#  s#   ù€ Ü+¨IÓ6ˆ	ò	ô Hˆ	ÔØÐr   c                 ó2   — t        || j                  ¬«      }y )N)r(   )r   r(   ©r'   r   r)   s      r   Ú$_process_model_before_weight_loadingz3HqqHfQuantizer._process_model_before_weight_loading/  s   € ô ' uÀ$×BZÑBZÔ[‰r   c                 ó>   — d|_         | j                  «       |_        |S ©NT)Úis_hqq_quantizedÚis_serializableÚis_hqq_serializabler£   s      r   Ú#_process_model_after_weight_loadingz2HqqHfQuantizer._process_model_after_weight_loading8  s    € Ø!%ˆÔØ$(×$8Ñ$8Ó$:ˆÔ!Øˆr   c                  ó   — yr¦   r`   )r'   Úsafe_serializations     r   r¨   zHqqHfQuantizer.is_serializable=  s   € Ør   c                  ó   — yr¦   r`   )r'   s    r   Úis_trainablezHqqHfQuantizer.is_trainable@  s   € àr   )r   r   rN   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__Úuse_keep_in_fp32_modulesÚ requires_parameters_quantizationÚrequires_calibrationÚrequired_packagesr$   rA   r   ÚstrrI   rs   r   r   Úboolry   r–   r‡   r¤   rª   r¨   Úpropertyr®   Ú__classcell__)r*   s   @r   r    r    -   sY  ø„ ñð  %ÐØ'+Ð$Ø ÐØ˜Ðô%ò
Ið@ Ø&ð Ø6:¸3±ið ØILð à	ˆc‰ó ð7Ø&ð7Ø7;¸C±yð7ØOSÐTWÉyð7à	ˆc‰ó7ðrà ðð $ðð ð	ð
 ˜˜c˜‘Nðð 
óð4f!à ðf!ð $ðf!ð ð	f!ð
 &ðf!ð ˜˜c˜‘Nðf!ð ˜c™óf!òR
ð\à ó\óó
ð ð˜dò ó ôr   r    )Útypingr   r   r   r   Úintegrationsr   Úutilsr	   r
   r   r   Úbaser   Úquantizers_utilsr   Úmodeling_utilsr   Úaccelerate.hooksr   r4   Ú
get_loggerr¯   r9   r   r    r`   r   r   ú<module>rÃ      s]   ð÷ 2Ó 1å 1ß ZÓ ZÝ Ý 2ñ Ý0ñ ÔÝ8áÔÛà	ˆ×	Ñ	˜HÓ	%€òôU�[õ Ur   