Ë
    îÍ:jz  ã                   ó¤   — d dl 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Ú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e
defd„Zd
dddde
dd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      €ú{/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/transformers/quantizers/quantizer_eetq.pyr   zEetqHfQuantizer.__init__-   s   ø€ Ü‰ÑÐ,Ñ7°Ò7Ø#6ˆÕ ó    c                 ó@  — 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«      }|€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 Ó-ˆ
ØÐÜ×ÑðIõð Ð#Ü˜*¤dÔ+°¸*×:KÑ:KÓ:MÑ1MÐQWÐ[e×[lÑ[lÓ[nÑQnÜ ðhóð ð RoÐ+ð $øô= ò 
	Ø!¤S¨£XÑ-ô "ðnóð ðð
 ûð
	ús   —C5 Ã5	DÃ>DÄDÚreturnc                 óª   — |€(t         j                  }t        j                  d|«       |S |t         j                  k7  rt        j                  d«       |S )NzîOverriding dtype=%s with `dtype=torch.float16` due to requirements of `eetq` to enable model loading in 8-bit. Pass your own dtype to specify the dtype of the remaining non-linear layers or pass dtype=torch.float16 to remove this warning.zLWe suggest you to set `dtype=torch.float16` for better efficiency with EETQ.)r'   Úfloat16r+   Úinfo)r   Údtypes     r   Úupdate_dtypezEetqHfQuantizer.update_dtype_   sM   € Øˆ=Ü—M‘MˆEÜ�K‰Kð?ð ôð ˆð ”e—m‘mÒ#Ü�K‰KÐfÔgØˆr   Úmodelr   Ú
param_namec                 ól   — ddl m} t        ||«      \  }}t        ||«      r| j                  s|dk(  ryyy)Nr   )Ú
EetqLinearÚbiasFT)r   r<   r   r-   Úpre_quantized)r   r9   r:   r   r<   ÚmoduleÚtensor_names          r   Úparam_needs_quantizationz(EetqHfQuantizer.param_needs_quantizationm   s9   € Ý#ä2°5¸*ÓEÑˆ�ä�f˜jÔ)Ø×!Ò! [°FÒ%:ØàØr   Úparam_valueztorch.TensorÚtarget_deviceztorch.devicec                 óz  — ddl m}m} t        ||«      \  }}	 ||«      \  }
}t	        ||«      rN| j
                  s|	dk(  r-|	dk(  r8|j                  t        j                  k7  rt        d«      ‚|	dk(  rt        d«      ‚|
j                  |«      |j                  |	<   |j                  d|j                  |«      «       y )	Nr   )r<   Úquantize_and_preprocess_weightsr=   Úweightz6Expect quantized weights but got an unquantized weightÚweight_scalez;Expect unquantized weights but got a quantized weight_scaleÚweight_scales)r   r<   rE   r   r-   r>   r7   r'   Úint8r&   ÚtoÚ_buffersÚregister)r   r9   rB   r:   rC   r   r<   rE   r?   r@   Ú	new_valuerG   s               r   Úcreate_quantized_paramz&EetqHfQuantizer.create_quantized_paramy   s«   € ÷ 	Eä2°5¸*ÓEÑˆ�Ù"AÀ+Ó"NÑˆ	�<ô �f˜jÔ)Ø×!Ò! [°FÒ%:Ø (Ò*¨{×/@Ñ/@ÄEÇJÁJÒ/NÜ$Ð%]Ó^Ð^à .Ò0Ü$Ð%bÓcÐcà'0§|¡|°MÓ'Bˆ�‰˜Ñ$Ø�‰˜¨¯©¸Ó)GÕHr   c                 ó   — |S r   © )r   r9   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   r>   )ÚintegrationsrT   Úget_modules_to_not_convertr   rU   r>   Úconfig)r   r9   rR   r   rT   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©NTrP   )r   Úsafe_serializations     r   Úis_serializablezEetqHfQuantizer.is_serializableª   s   € Ør   c                  ó   — yr[   rP   )r   s    r   Úis_trainablezEetqHfQuantizer.is_trainable­   s   € àr   )r7   útorch.dtyper3   r`   )r9   r   r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__Ú requires_parameters_quantizationÚrequires_calibrationÚrequired_packagesr   r2   r8   r$   ÚboolrA   rN   rQ   r   ÚlistrY   r]   Úpropertyr_   Ú__classcell__)r   s   @r   r   r   !   sÊ   ø„ ñð (,Ð$Ø Ðà Ð.Ðô7ò,ó\ð
Ð.?ð 
ÈSð 
Ð_có 
ðIà ðIð $ðIð ð	Ið
 &óIó2ð 59ñDà ðDð ' t¨C¡yÑ1óDó*ð ð˜dò ó ôr   r   )Útypingr   r   Úbaser   Úmodeling_utilsr   Úutilsr	   r
   r   r   Úquantizers_utilsr   r'   Ú
get_loggerra   r+   r   rP   r   r   ú<module>rr      sK   ð÷ +å ñ Ý0ç [Ó [Ý 2ñ ÔÛð 
ˆ×	Ñ	˜HÓ	%€ôN�kõ Nr   