Ë
    T^(hT!  ã                   óÌ   — d dl mZmZ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mZmZ dd	lmZ  e«       rd d
lZ ej.                  e«      Z G d„ de	«      Zy
)é    )ÚTYPE_CHECKINGÚAnyÚDictÚListÚOptionalé   )Útqdmé   )ÚHfQuantizer)Úget_module_from_name)ÚPreTrainedModel)Úis_accelerate_availableÚis_flute_availableÚis_hadamard_availableÚis_torch_availableÚlogging)ÚQuantizationConfigMixinNc                   ó  ‡ — e Zd ZdZdZdZddgZdefˆ 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„Zd d„Zdee   ded	ee   fd„Zedded   fd„«       Zdd„Zdddddedeeef   d	ef
d„Zd„ Zˆ xZS )!ÚHiggsHfQuantizerzˆ
    Quantizer of the HIGGS method. Enables the loading of prequantized models and in-flight quantization of full-precision models.
    FTzflute-kernelÚfast_hadamard_transformÚquantization_configc                 ó4   •— t        ‰| �  |fi |¤Ž || _        y ©N)ÚsuperÚ__init__r   )Úselfr   ÚkwargsÚ	__class__s      €úe/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/quantizers/quantizer_higgs.pyr   zHiggsHfQuantizer.__init__+   s   ø€ Ü‰ÑÐ,Ñ7°Ò7Ø#6ˆÕ ó    c                 ón  — t         j                  j                  «       st        d«      ‚t	        «       st        d«      ‚t        «       st        d«      ‚t        «       st        d«      ‚|€t        d«      ‚t        |t        «      r0d|j                  «       v sd|j                  «       v rt        d«      ‚y y )	NzNHIGGS quantization is only supported on GPU. Please use a different quantizer.zHUsing `higgs` quantization requires Accelerate: `pip install accelerate`zLUsing `higgs` quantization requires FLUTE: `pip install flute-kernel>=0.3.0`zbUsing `higgs` quantization requires fast_hadamard_transform: `pip install fast_hadamard_transform`zwYou are attempting to load a HIGGS model without setting device_map. Please set device_map comprised of 'cuda' devices.ÚcpuÚdiskz¯You are attempting to load a HIGGS 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.)ÚtorchÚcudaÚis_availableÚNotImplementedErrorr   ÚImportErrorr   r   Ú
ValueErrorÚ
isinstanceÚdictÚvalues)r   Ú
device_mapr   s      r   Úvalidate_environmentz%HiggsHfQuantizer.validate_environment/   s½   € Ü�z‰z×&Ñ&Ô(Ü%Ð&vÓwÐwä&Ô(ÜÐhÓiÐiä!Ô#ÜÐlÓmÐmä$Ô&ÜØtóð ð ÐÜðFóð ô ˜
¤DÔ)¨u¸
×8IÑ8IÓ8KÑ/KÈvÐYc×YjÑYjÓYlÑOlÜðdóð ð PmÐ)r    Úreturnc                 óÂ   — |€'t         j                  d«       t        j                  }|S |t        j                  k7  r"|t        j                  k7  rt        d|› d�«      ‚|S )NzS`torch_dtype` is None. Setting `torch_dtype=torch.float16` for FLUTE compatibility.zInvalid `torch_dtype` z_. HIGGS quantization only supports `torch_dtype=torch.float16` or `torch_dtype=torch.bfloat16`.)ÚloggerÚinfor$   Úfloat16Úbfloat16r)   )r   Útorch_dtypes     r   Úupdate_torch_dtypez#HiggsHfQuantizer.update_torch_dtypeI   sg   € ØÐÜ�K‰KÐmÔnÜŸ-™-ˆKð Ðð œEŸM™MÒ)¨k¼U¿^¹^Ò.KÜØ(¨¨ð  6Uð  Vóð ð Ðr    Úmodelr   Úparam_valueztorch.TensorÚ
param_nameÚtarget_deviceztorch.deviceÚ
state_dictÚunexpected_keysc                 ó(  — ddl m} 	  ||j                  |«      | j                  j                  | j                  j
                  | j                  j                  | j                  j                  «      }~t        ||«      \  }	}
dj                  |j                  d«      d d «      }|j                  «       D ]Á  \  }}||	j                  v r/t        j                  j                  |d¬«      |	j                  |<   ŒC||	j                   v r-t        j                  j#                  |«      |	j                   |<   Œ~|dk(  r/||	_        |j'                  «       | j                  j$                  |<   Œ²t)        d|› d	|	› �«      ‚ |�||v r|j+                  |«       y y y )
Nr   )Úquantize_with_higgsú.éÿÿÿÿF)Úrequires_gradÚtune_metadatazUnexpected key z in module )Úintegrationsr>   Útor   ÚbitsÚpÚ
group_sizeÚhadamard_sizer   ÚjoinÚsplitÚitemsÚ_parametersr$   ÚnnÚ	ParameterÚ_buffersÚBufferrB   Úto_dictr)   Úremove)r   r7   r8   r9   r:   r;   r<   r>   Ú
flute_dictÚmoduleÚ_Úmodule_nameÚkeyÚvalues                 r   Úcreate_quantized_paramz'HiggsHfQuantizer.create_quantized_paramT   sw  € õ 	7ð	ñ )Ø�N‰N˜=Ó)Ø×$Ñ$×)Ñ)Ø×$Ñ$×&Ñ&Ø×$Ñ$×/Ñ/Ø×$Ñ$×2Ñ2ó
ˆ
ð ä(¨°
Ó;‰	ˆ�Ø—h‘h˜z×/Ñ/°Ó4°S°bÐ9Ó:ˆØ$×*Ñ*Ó,ò 		M‰JˆC�Ø�f×(Ñ(Ñ(Ü*/¯(©(×*<Ñ*<¸UÐRWÐ*<Ó*X�×"Ñ" 3Ò'Ø˜Ÿ™Ñ'Ü',§x¡x§¡°uÓ'=�—‘ Ò$Ø˜Ò'Ø',�Ô$ØFKÇmÁmÃo�×(Ñ(×6Ñ6°{ÒCä  ?°3°%°{À6À(Ð!KÓLÐLð		Mð Ð&¨:¸Ñ+HØ×"Ñ" :Õ.ð ,IÐ&r    c                 ón   — ddl m}  ||| j                  ¬«       | j                  |j                  _        y )Nr   )Úreplace_with_higgs_linear)r   )rC   r[   r   Úconfig)r   r7   r   r[   s       r   Ú$_process_model_before_weight_loadingz5HiggsHfQuantizer._process_model_before_weight_loading{   s/   € õ
 	=á!ØØ $× 8Ñ 8õ	
ð ,0×+CÑ+Cˆ�‰Õ(r    c                 ó   — ddl m}m} ddlm} ddlm} i }|j                  «       D ��	ci c]  \  }}	t        |	|«      sŒ||	“Œ }
}}	t        |
j                  «       dd¬«      D �]"  \  }}	|	j                  j                  |vr4 ||	j                  j                  ¬	«      ||	j                  j                  <   ||	j                  j                     |	_        |j                  | j                  j                   |   «      |	_         ||	j                  j"                  |	j$                  j"                  |	j                   ¬
«      \  |	j                  _        |	_        |	j                   j'                  «       | j                  j                   |<   �Œ% y c c}	}w )Nr   )ÚTuneMetaDataÚmaybe_tune_and_repack)Úmake_workspace_streamkr   ©ÚHiggsLinearzRepacking HIGGS modulesF)ÚdescÚleave)Údevice)ÚweightÚscalesÚmetadata)Ú
flute.tuner_   r`   Úflute.utilsra   rC   rc   Únamed_modulesr*   r	   rK   rg   rf   Ú	workspaceÚ	from_dictr   rB   Údatarh   rQ   )r   r7   r   r_   r`   ra   rc   Úflute_workspacesÚnamerT   Úflute_moduless              r   Ú#_process_model_after_weight_loadingz4HiggsHfQuantizer._process_model_after_weight_loadingˆ   sN  € ßBÝ6å.àÐØ:?×:MÑ:MÓ:O×s©,¨$°ÔS]Ð^dÐfqÕSr˜˜v™ÐsˆÑsÜ  ×!4Ñ!4Ó!6Ð=VÐ^cÔdó 	Z‰LˆD�&ð �}‰}×#Ñ#Ð+;Ñ;Ù9OÐW]×WdÑWd×WkÑWkÔ9lÐ  §¡×!5Ñ!5Ñ6Ø/°·±×0DÑ0DÑEˆFÔð $0×#9Ñ#9¸$×:RÑ:R×:`Ñ:`ÐaeÑ:fÓ#gˆFÔ Ù7LØ—}‘}×)Ñ)Ø—}‘}×)Ñ)Ø×-Ñ-ô8Ñ4ˆF�M‰MÔ Ô 4ð
 <B×;OÑ;O×;WÑ;WÓ;YˆD×$Ñ$×2Ñ2°4Ó8ñ	Zùó ts
   ªF
¿F
Úmissing_keysÚprefixc                 óà   ‡‡	— ddl m} |j                  «       D ��ch c]  \  }}t        ||«      sŒ|’Œ c}}Š	dt        dt
        fˆ	ˆfd„}|D �cg c]  } ||«      rŒ|‘Œ c}S c c}}w c c}w )Nr   rb   rW   r/   c                 ó†   •‡ ‡— ‰ j                  d«      s‰ j                  d«      ry‰› d‰ › �Št        ˆˆ fd„‰D «       «      S )Nz.weightz.biasFr?   c              3   ó2   •K  — | ]  }|‰v xs |‰v –— Œ y ­wr   © )Ú.0rq   Úfull_keyrW   s     €€r   ú	<genexpr>zNHiggsHfQuantizer.update_missing_keys.<locals>.should_update.<locals>.<genexpr>ª   s"   øè ø€ ÒO¸4�t˜s�{Ò6 d¨hÐ&6Ó6ÑOùs   ƒ)ÚendswithÚany)rW   r{   Úhiggs_namesru   s   `@€€r   Úshould_updatez;HiggsHfQuantizer.update_missing_keys.<locals>.should_update¦   s>   ú€ Ø�|‰|˜IÔ&¨#¯,©,°wÔ*?ØØ ˜  3 %Ð(ˆHÜÔOÀ;ÔOÓOÐOr    )rC   rc   rl   r*   ÚstrÚbool)
r   r7   rt   ru   rc   rq   rT   r€   rW   r   s
      `     @r   Úupdate_missing_keysz$HiggsHfQuantizer.update_missing_keys¡   sj   ù€ Ý.à05×0CÑ0CÓ0E×i¡  fÌÐTZÐ\gÕIh’tÓiˆð	Pœsð 	P¤tö 	Pð  ,ÖF˜±=ÀÕ3E’ÒFÐFùó jùò Gs   œA%±A%ÁA+ÁA+c                  ó   — y)NFry   )r   r7   s     r   Úis_trainablezHiggsHfQuantizer.is_trainable®   s   € àr    c                  ó   — y)NTry   )r   Úsafe_serializations     r   Úis_serializablez HiggsHfQuantizer.is_serializable²   s   € Ør    c                 óŒ   — ddl m} t        ||«      \  }}t        ||«      r#|dk(  r|j                  t
        j                  k7  ryy)Nr   rb   rg   TF)rC   rc   r   r*   Údtyper$   Úint16)	r   r7   r8   r9   r;   r   rc   rT   Útensor_names	            r   Úcheck_quantized_paramz&HiggsHfQuantizer.check_quantized_paramµ   sC   € õ 	/ä2°5¸*ÓEÑˆ�Ü�f˜kÔ*¨{¸hÒ/FÈ;×K\ÑK\Ô`e×`kÑ`kÒKkààr    c                 ó"   — ddl m}  ||«      }|S )Nr   )Údequantize_higgs)rC   r�   )r   r7   r�   s      r   Ú_dequantizezHiggsHfQuantizer._dequantizeÆ   s   € Ý3á  Ó'ˆØˆr    )r5   útorch.dtyper/   r‘   r   )r7   r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__Úrequires_calibrationÚ requires_parameters_quantizationÚrequired_packagesr   r   r.   r6   r�   r   r   r   r   rY   r]   rs   rƒ   Úpropertyr…   rˆ   r‚   r�   r�   Ú__classcell__)r   s   @r   r   r   "   s5  ø„ ñð !ÐØ'+Ð$Ø'Ð)BÐCÐð7Ð,Cõ 7òó4	ð$ 04ñ%/à ð%/ð $ð%/ð ð	%/ð
 &ð%/ð ˜˜c˜‘Nð%/ð " $ s¡)Ñ,ó%/ðNDà óDóZð2G°t¸C±yð GÈ#ð GÐRVÐWZÑR[ó Gð ñ (Ð+<Ñ"=ò ó ðóðà ðð $ðð ð	ð
 ˜˜c˜‘Nðð 
óö"r    r   )Útypingr   r   r   r   r   Úutils.loggingr	   Úbaser   Úquantizers_utilsr   Úmodeling_utilsr   Úutilsr   r   r   r   r   Úutils.quantization_configr   r$   Ú
get_loggerr’   r1   r   ry   r    r   ú<module>r£      sR   ð÷ <Õ ;å  Ý Ý 2ñ Ý0ç sÕ sÝ ?ñ ÔÛà	ˆ×	Ñ	˜HÓ	%€ôh�{õ hr    