Ë
    T^(h9  ã                   ó°   — d dl mZmZmZmZ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 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é   )ÚHfQuantizeré   )ÚPreTrainedModel)Úis_accelerate_availableÚis_eetq_availableÚis_torch_availableÚlogging)Úget_module_from_nameNc                   óâ   ‡ — e Zd ZdZdZdZddgZˆ fd„Zd„ Zdd	„Z	d
dddde
dee
ef   f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
ddeee
      fd„Zdd„Zedefd„«       Zˆ xZS )ÚEetqHfQuantizera  
    8-bit quantization from EETQ quantization method:
        before loading: converts transformer layers into W8A16Linear during loading: load 16bit weight and pass to the
        layer object after: quantizes individual weights in Linear8bitLt into 8bit at first .cuda() call
    TFÚeetqÚ
acceleratec                 ó4   •— t        ‰| �  |fi |¤Ž || _        y ©N)ÚsuperÚ__init__Úquantization_config)Úselfr   ÚkwargsÚ	__class__s      €úd/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/quantizers/quantizer_eetq.pyr   zEetqHfQuantizer.__init__-   s   ø€ Ü‰ÑÐ,Ñ7°Ò7Ø#6ˆÕ ó    c                 óB  — t        «       st        d«      ‚	 dd l}t	        «       st        d«      ‚|j                  dd«      s|j                  dd«      rt        d	«      ‚t        j                  j                  «       st        d
«      ‚|j                  dd «      }|€t        j                  d«       y |�At        |t        «      r0d|j                  «       v sd|j                  «       v rt        d«      ‚y y y # t        $ r}dt        |«      v rt        d«      |‚‚ d }~ww xY w)NzƒUsing `eetq` 8-bit quantization requires eetq.Please install the latest version of eetq from : https://github.com/NetEase-FuXi/EETQr   Úshard_checkpointz³You are using a version of EETQ that is incompatible with the current transformers version. Either downgrade transformers to <= v4.46.3 or, if available, upgrade EETQ to > v1.0.0.zNLoading an EETQ quantized model requires accelerate (`pip install accelerate`)Úfrom_tfFÚ	from_flaxz‚Converting into 8-bit 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.Ú
device_mapzŽYou have loaded an EETQ 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.ÚcpuÚdiskz¯You are attempting to load an EETQ 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   ÚImportErrorr   Ústrr   ÚgetÚ
ValueErrorÚtorchÚcudaÚis_availableÚRuntimeErrorÚloggerÚwarning_onceÚ
isinstanceÚdictÚvalues)r   Úargsr   r   Úexcr#   s         r   Úvalidate_environmentz$EetqHfQuantizer.validate_environment1   s?  € Ü Ô"Üðhóð ð
	Ûô 'Ô(ÜÐnÓoÐoà�:‰:�i Ô'¨6¯:©:°kÀ5Ô+IÜð;óð ô
 �z‰z×&Ñ&Ô(ÜÐPÓQÐQà—Z‘Z ¨dÓ3ˆ
ØÐÜ×ÑðIõð Ð#Ü˜*¤dÔ+°¸*×:KÑ:KÓ:MÑ1MÐQWÐ[e×[lÑ[lÓ[nÑQnÜ ðhóð ð RoÐ+ð $øô= ò 
	Ø!¤S¨£XÑ-ô "ðnóð ðð
 ûð
	ús   —C6 Ã6	DÃ?DÄDÚreturnc                 óª   — |€(t         j                  }t        j                  d|«       |S |t         j                  k7  rt        j                  d«       |S )Na  Overriding torch_dtype=%s with `torch_dtype=torch.float16` due to requirements of `eetq` to enable model loading in 8-bit. Pass your own torch_dtype to specify the dtype of the remaining non-linear layers or pass torch_dtype=torch.float16 to remove this warning.zRWe suggest you to set `torch_dtype=torch.float16` for better efficiency with EETQ.)r*   Úfloat16r.   Úinfo)r   Útorch_dtypes     r   Úupdate_torch_dtypez"EetqHfQuantizer.update_torch_dtype_   sQ   € ØÐÜŸ-™-ˆKÜ�K‰KðEð ôð Ðð œEŸM™MÒ)Ü�K‰KÐlÔmØÐr   Úmodelr   Úparam_valueztorch.TensorÚ
param_nameÚ
state_dictc                 óæ   — 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   )Ú
EetqLinearÚbiasÚweightz6Expect quantized weights but got an unquantized weightFÚweight_scalez;Expect unquantized weights but got a quantized weight_scaleT)	r   rA   r   r0   Úpre_quantizedÚdtyper*   Úint8r)   )	r   r<   r=   r>   r?   r   rA   ÚmoduleÚtensor_names	            r   Úcheck_quantized_paramz%EetqHfQuantizer.check_quantized_paramm   st   € õ 	$ä2°5¸*ÓEÑˆ�ä�f˜jÔ)Ø×!Ò! [°FÒ%:Ø (Ò*¨{×/@Ñ/@ÄEÇJÁJÒ/NÜ$Ð%]Ó^Ð^Øà .Ò0Ü$Ð%bÓcÐcØØr   Útarget_deviceztorch.deviceÚunexpected_keysc                 óÂ   — ddl m} t        ||«      \  }}	 ||«      \  }
}|
j                  |«      |j                  |	<   |j                  d|j                  |«      «       y)zB
        quantizes weights into qweight and weight_scales
        r   )Úquantize_and_preprocess_weightsÚweight_scalesN)r   rN   r   ÚtoÚ_buffersÚregister)r   r<   r=   r>   rK   r?   rL   rN   rH   rI   Ú	new_valuerD   s               r   Úcreate_quantized_paramz&EetqHfQuantizer.create_quantized_param„   sU   € õ 	9ä2°5¸*ÓEÑˆ�Ù"AÀ+Ó"NÑˆ	�<à'0§|¡|°MÓ'Bˆ�‰˜Ñ$Ø�‰˜¨¯©¸Ó)GÕHr   c                 ó   — |S r   © )r   r<   r   s      r   Ú#_process_model_after_weight_loadingz3EetqHfQuantizer._process_model_after_weight_loading˜   s   € Øˆr   Úkeep_in_fp32_modulesc                 óò   — ddl m} | j                  || j                  j                  |«      | _         ||| j                  | j                  | j
                  ¬«      }| j                  |j                  _        y )Nr
   )Úreplace_with_eetq_linear)Úmodules_to_not_convertr   rE   )ÚintegrationsrZ   Úget_modules_to_not_convertr   r[   rE   Úconfig)r   r<   rX   r   rZ   s        r   Ú$_process_model_before_weight_loadingz4EetqHfQuantizer._process_model_before_weight_loading›   sl   € õ 	<à&*×&EÑ&EØ�4×+Ñ+×BÑBÐDXó'
ˆÔ#ñ )ØØ#'×#>Ñ#>Ø $× 8Ñ 8Ø×,Ñ,ô	
ˆð ,0×+CÑ+Cˆ�‰Õ(r   c                  ó   — y©NTrV   )r   Úsafe_serializations     r   Úis_serializablezEetqHfQuantizer.is_serializable°   s   € Ør   c                  ó   — yra   rV   )r   s    r   Úis_trainablezEetqHfQuantizer.is_trainable³   s   € àr   )r:   útorch.dtyper6   rf   r   )r<   r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__Ú requires_parameters_quantizationÚrequires_calibrationÚrequired_packagesr   r5   r;   r'   r   r   rJ   r   r   rT   rW   r_   rc   ÚpropertyÚboolre   Ú__classcell__)r   s   @r   r   r   !   s  ø„ ñð (,Ð$Ø Ðà Ð.Ðô7ò,ó\ðà ðð $ðð ð	ð
 ˜˜c˜‘Nóð< 04ñIà ðIð $ðIð ð	Ið
 &ðIð ˜˜c˜‘NðIð " $ s¡)Ñ,óIó(ð 59ñDà ðDð ' t¨C¡yÑ1óDó*ð ð˜dò ó ôr   r   )Útypingr   r   r   r   r   Úbaser	   Úmodeling_utilsr   Úutilsr   r   r   r   Úquantizers_utilsr   r*   Ú
get_loggerrg   r.   r   rV   r   r   ú<module>rw      sN   ð÷ <Õ ;å ñ Ý0ç [Ó [Ý 2ñ ÔÛð 
ˆ×	Ñ	˜HÓ	%€ôT�kõ Tr   