Ë
    T^(hs  ã                   ó    — 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  e«       rd dlZ ej                  e«      Z G d„ d	e«      Zy)
é    )ÚTYPE_CHECKINGÚDictÚListÚOptionalÚUnioné   )ÚHfQuantizeré   )ÚPreTrainedModel)Úis_accelerate_availableÚis_torch_availableÚloggingNc                   ó¼   ‡ — e Zd ZdZdZdZdgZˆ fd„Zd„ Zdd	„Z		 dddd
e
ee      fd„Zdeeeeef   f   deeeeef   f   fd„Zdd„Zdd„Zedefd„«       Zˆ xZS )ÚBitNetHfQuantizerzë
    1.58-bit quantization from BitNet quantization method:
    Before loading: it converts the linear layers into BitLinear layers during loading.

    Checkout the paper introducing this method : https://arxiv.org/pdf/2402.17764
    FTÚ
acceleratec                 ó4   •— t        ‰| �  |fi |¤Ž || _        y ©N)ÚsuperÚ__init__Úquantization_config)Úselfr   ÚkwargsÚ	__class__s      €úf/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/quantizers/quantizer_bitnet.pyr   zBitNetHfQuantizer.__init__-   s   ø€ Ü‰ÑÐ,Ñ7°Ò7Ø#6ˆÕ ó    c                 óÎ  — t        «       st        d«      ‚|j                  dd«      s|j                  dd«      rt        d«      ‚t        j
                  j                  «       st        j                  d«       y |j                  dd «      }|€t        j                  d«       y |�At        |t        «      r0d	|j                  «       v sd
|j                  «       v rt        d«      ‚y y y )NzOLoading a BitNet quantized model requires accelerate (`pip install accelerate`)Úfrom_tfFÚ	from_flaxztLoading ternary weights from tf/flax is currently not supported, please make sure the weights are in PyTorch format.zhYou don't have a GPU available to load the model, the inference will be slow because of weight unpackingÚ
device_mapz�You have loaded a BitNet 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 a BitNet 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   ÚImportErrorÚgetÚ
ValueErrorÚtorchÚcudaÚis_availableÚloggerÚwarning_onceÚ
isinstanceÚdictÚvalues)r   Úargsr   r   s       r   Úvalidate_environmentz&BitNetHfQuantizer.validate_environment1   sé   € Ü&Ô(ÜÐoÓpÐpà�:‰:�i Ô'¨6¯:©:°kÀ5Ô+IÜð;óð ô
 �z‰z×&Ñ&Ô(Ü×ÑØzôð à—Z‘Z ¨dÓ3ˆ
ØÐÜ×ÑðIõð Ð#Ü˜*¤dÔ+°¸*×:KÑ:KÓ:MÑ1MÐQWÐ[e×[lÑ[lÓ[nÑQnÜ ðgóð ð RoÐ+ð $r   Úmodelr   c                 ó   — |S r   © )r   r/   r   s      r   Ú#_process_model_after_weight_loadingz5BitNetHfQuantizer._process_model_after_weight_loadingN   s   € Øˆr   Úkeep_in_fp32_modulesc                 ó¼   — ddl m} | j                  || j                  j                  |«      | _         ||| j                  | j                  | j
                  ¬«      }y )Nr
   )Úreplace_with_bitnet_linear)Úmodules_to_not_convertr   Úpre_quantized)Úintegrationsr5   Úget_modules_to_not_convertr   r6   r7   )r   r/   r3   r   r5   s        r   Ú$_process_model_before_weight_loadingz6BitNetHfQuantizer._process_model_before_weight_loadingQ   sX   € õ 	>à&*×&EÑ&EØ�4×+Ñ+×BÑBÐDXó'
ˆÔ#ñ +ØØ#'×#>Ñ#>Ø $× 8Ñ 8Ø×,Ñ,ô	
‰r   Ú
max_memoryÚreturnc                 ó^   — |j                  «       D ��ci c]  \  }}||dz  “Œ }}}|S c c}}w )NgÍÌÌÌÌÌì?)Úitems)r   r;   ÚkeyÚvals       r   Úadjust_max_memoryz#BitNetHfQuantizer.adjust_max_memoryd   s6   € Ø6@×6FÑ6FÓ6H×I©(¨#¨s�c˜3 ™:‘oÐIˆ
ÑIØÐùó Js   ”)c                 ó&   — t         j                  }|S r   )r%   Úint8)r   Útarget_dtypes     r   Úadjust_target_dtypez%BitNetHfQuantizer.adjust_target_dtypeh   s   € Ü—z‘zˆØÐr   c                  ó   — y)NTr1   )r   Úsafe_serializations     r   Úis_serializablez!BitNetHfQuantizer.is_serializablel   s   € Ør   c                  ó   — y)NFr1   )r   s    r   Úis_trainablezBitNetHfQuantizer.is_trainableo   s   € àr   )r/   r   r   )rD   útorch.dtyper<   rK   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__Ú requires_parameters_quantizationÚrequires_calibrationÚrequired_packagesr   r.   r2   r   r   Ústrr:   r   r   ÚintrA   rE   rH   ÚpropertyÚboolrJ   Ú__classcell__)r   s   @r   r   r       s´   ø„ ñð (-Ð$ØÐà%˜Ðô7òó:ð 59ñ
à ð
ð ' t¨C¡yÑ1ó
ð&¨D°°e¸CÀ¸H±oÐ1EÑ,Fð È4ÐPSÐUZÐ[^Ð`cÐ[cÑUdÐPdÑKeó óóð ð˜dò ó ôr   r   )Útypingr   r   r   r   r   Úbaser	   Úmodeling_utilsr   Úutilsr   r   r   r%   Ú
get_loggerrL   r(   r   r1   r   r   ú<module>r]      sK   ð÷ >Õ =å ñ Ý0ç HÑ Hñ ÔÛð 
ˆ×	Ñ	˜HÓ	%€ôQ˜õ Qr   