Ë
    Gêñi!)  ã                   óà   — d dl Z d dlZd dlZd dlmZ d dlmZmZ  e«       rddlm	Z	 ddl
mZ  e«       rd dlmZ d dlmZ  ej                   e«      Zd	„ Zd
„ Z G d„ de	«      Z G d„ de	«      Zy)é    N)Úlogging)Úis_torch_availableÚis_torchao_availableé   )ÚConversionOps)Úget_module_from_name)Úunflatten_tensor_state_dict)Úis_metadata_torchaoc                 ó  — ddl m} ddlm} t	        | |«      r*| j
                  j                  › d| j                  «       › d�S t	        | |«      r<| j
                  j                  › d| j                  › dt        | j                  «      › d�S y )Nr   )ÚAffineQuantizedTensor)ÚLinearActivationQuantizedTensorú(ú)z(activation=ú	, weight=)
Útorchao.dtypesr   Ú7torchao.quantization.linear_activation_quantized_tensorr   Ú
isinstanceÚ	__class__Ú__name__Ú_quantization_typeÚinput_quant_funcÚoriginal_weight_tensor)Úweightr   r   s      úc/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/transformers/integrations/torchao.pyr   r   &   s¤   € Ý4Ýgä�&Ð/Ô0Ø×"Ñ"×+Ñ+Ð,¨A¨f×.GÑ.GÓ.IÐ-JÈ!ÐLÐLä�&Ð9Ô:Ø×"Ñ"×+Ñ+Ð,¨L¸×9PÑ9PÐ8QÐQZÔ[mÐnt÷  oLñ  oLó  \Mð  [Nð  NOð  Pð  	Pð ;ó    c                 ó  — t        | j                  «      }|€7d| j                  j                  d   › d| j                  j                  d   › d�S d| j                  j                  d   › d| j                  j                  d   › d|› �S )Nzin_features=é   z, out_features=r   z, weight=Noner   )r   r   Úshape)Úselfr   s     r   Ú_linear_extra_reprr    1   s‰   € Ü §¡Ó,€FØ€~Ø˜dŸk™k×/Ñ/°Ñ2Ð3°?À4Ç;Á;×CTÑCTÐUVÑCWÐBXÐXeÐfÐfà˜dŸk™k×/Ñ/°Ñ2Ð3°?À4Ç;Á;×CTÑCTÐUVÑCWÐBXÐXaÐbhÐaiÐjÐjr   c                   ó¨   — e Zd Zd„ Zd„ Z	 	 	 d	deeej                  f   dej                  j                  dz  dedz  deeej                  f   fd„Zy)
ÚTorchAoQuantizec                 ó   — || _         y ©N©Úhf_quantizer©r   r&   s     r   Ú__init__zTorchAoQuantize.__init__:   ó
   € Ø(ˆÕr   c                 ó  — ddl m} t        |j                  «       «      j                  }| j
                  j                  r?|j                  dk(  r0|j                  d«        |||g|¢­i |¤Ž |j                  d«       y |||g|¢­i |¤Ž y)a7  Run quantize_, moving to CUDA first if CPU offloading is active.

        Some torchao quantization ops (e.g. int4 packing) only have CUDA kernels.
        When a layer is destined for CPU (e.g. CPU offloading), we temporarily move
        it to CUDA for quantization, then move the result back to CPU.
        r   )Ú	quantize_ÚcpuÚcudaN)	Útorchao.quantizationr+   ÚnextÚ
parametersÚdevicer&   Úoffload_to_cpuÚtypeÚto)r   ÚmoduleÚconfigÚargsÚkwargsr+   Útarget_devices          r   Ú	_quantizezTorchAoQuantize._quantize=   s|   € õ 	3ä˜V×.Ñ.Ó0Ó1×8Ñ8ˆØ×Ñ×+Ò+°×0BÑ0BÀeÒ0KØ�I‰I�fÔÙ�f˜fÐ6 tÒ6¨vÒ6Ø�I‰I�eÕá�f˜fÐ6 tÒ6¨vÓ6r   NÚ
input_dictÚmodelÚfull_layer_nameÚreturnc                 óD  — t        |j                  «       «      d   \  }}t        |t        «      r|d   n|}t	        ||«      \  }}	t
        j                  j                  ||j                  ¬«      |j                  |	<   |j                  «       }
t        |«      t        |
«      k(  }| j                  j                  j                  }|r)|r't        |j                   j#                  d¬«      dd«       ddlm} | j                  j                  j)                  «       }t        ||«      �ré|j+                  dd	«      \  }}d }||j,                  v r(|j/                  d
«      rJ d«       ‚|j0                  |   }nÉ||j,                  v r(|j/                  d
«      rJ d«       ‚|j0                  |   }n“|j,                  D ]h  }|j/                  d
«      sŒt3        j4                  |dd  |«      r|j0                  |   } nHt3        j4                  |dd  |«      sŒY|j0                  |   } n |j0                  j7                  dd «      }|�Í|dk(  rr|r|r|j8                  j;                  «       }| j=                  ||d„ «       |j?                  |«       d|_         |jC                  d¬«      D ]	  }d|_         Œ |r|rdiS i S  |||i«      }| j=                  ||d ¬«       |j?                  |«       d|_         |jC                  d¬«      D ]	  }d|_         Œ i S ||iS |r|r|j8                  j;                  «       }| j=                  || j                  j                  j)                  «       «       |j?                  |«       d|_         |jC                  d¬«      D ]	  }d|_         Œ |r|rdiS i S )Nr   )Úrequires_gradT)ÚdecoderÚtie_word_embeddingsF)ÚFqnToConfigú.r   zre:zHparam fqn should not start with`re:`, which is used for specifying regexzImodule fqn should not start with`re:`, which is used for specifying regexé   Ú_defaultr   c                  ó   — y)NT© )ÚxÚfqns     r   ú<lambda>z)TorchAoQuantize.convert.<locals>.<lambda>Ž   s   � r   )Úrecursezlm_head.weight)Ú	filter_fn)"ÚtupleÚitemsr   Úlistr   ÚtorchÚnnÚ	Parameterr@   Ú_parametersÚget_input_embeddingsÚidr&   Úquantization_configÚuntie_embedding_weightsÚsetattrr6   Úget_text_configr.   rC   Úget_apply_tensor_subclassÚrsplitÚfqn_to_configÚ
startswithÚmodule_fqn_to_configÚreÚ	fullmatchÚgetr   Úcloner:   ÚdiscardÚ_is_hf_initializedr0   )r   r;   r<   r=   Úmissing_keysr8   Ú_Úvaluer5   Útensor_nameÚinput_embedÚis_embedding_paramrX   rC   r6   Ú
module_fqnÚtop_level_param_nameÚcÚmaybe_module_fqn_patternÚlm_headÚparamÚcustom_param_fqn_configs                         r   ÚconvertzTorchAoQuantize.convertN   s³  € ô ˜×)Ñ)Ó+Ó,¨QÑ/‰ˆˆ5Ü& u¬dÔ3��a’¸ˆä2°5¸/ÓJÑˆ�ä*/¯(©(×*<Ñ*<¸UÐRW×ReÑReÐ*<Ó*fˆ×Ñ˜;Ñ'ð ×0Ñ0Ó2ˆÜ ›Z¬2¨k«?Ñ:ÐØ"&×"3Ñ"3×"GÑ"G×"_Ñ"_Ðá"Ñ'9Ü�E—L‘L×0Ñ0¸Ð0Ó>Ð@UÐW\Ô]å4à×"Ñ"×6Ñ6×PÑPÓRˆÜ�f˜kÕ*Ø/>×/EÑ/EÀcÈ1Ó/MÑ,ˆJÐ,ØˆAØ &×"6Ñ"6Ñ6Ø%×0Ñ0°Ô7ð Ø^óÐ7ð ×/Ñ/°Ñ@‘Ø˜v×3Ñ3Ñ3Ø%×0Ñ0°Ô7ð Ø_óÐ7ð ×/Ñ/°
Ñ;‘ð 17×0DÑ0Dò JÐ,à3×>Ñ>¸uÔEØ äŸ™Ð&>¸q¸rÐ&BÀOÔTØ"×7Ñ7Ð8PÑQ˜ÙÜŸ™Ð&>¸q¸rÐ&BÀJÕOà"×7Ñ7Ð8PÑQ˜ÙðJð ×3Ñ3×7Ñ7¸
ÀDÓI�Aàˆ}Ø'¨8Ò3Ù)Ñ.EØ"(§-¡-×"5Ñ"5Ó"7˜à—N‘N 6¨1Ñ/BÔDØ ×(Ñ(¨Ô9Ø04�FÔ-ð
 "(×!2Ñ!2¸5Ð!2Ó!Aò 8˜Ø37˜Õ0ð8á:LÑQhÐ,¨gÐ6ÐpÐnpÐpñ /:Ð;OÐQRÐ:SÓ.TÐ+Ø—N‘N 6Ð+BÈd�NÔSØ ×(Ñ(¨Ô9Ø04�FÔ-Ø!'×!2Ñ!2¸5Ð!2Ó!Aò 8˜Ø37˜Õ0ð8à�IØ# UÐ+Ð+áÑ"9Ø—m‘m×)Ñ)Ó+ˆGØ�‰�v˜t×0Ñ0×DÑD×^Ñ^Ó`ÔaØ×Ñ˜_Ô-Ø$(ˆÔ!Ø×&Ñ&¨uÐ&Ó5ò 	,ˆEØ'+ˆEÕ$ð	,á.@ÑE\Ð  'Ð*ÐdÐbdÐdr   )NNN)r   Ú
__module__Ú__qualname__r(   r:   ÚdictÚstrrQ   ÚTensorrR   ÚModulers   rH   r   r   r"   r"   9   sy   „ ò)ò7ð( )-Ø&*Øñ\eà˜˜eŸl™lÐ*Ñ+ð\eð �x‰x�‰ Ñ%ð\eð ˜t™ð	\eð 
ˆc�5—<‘<ÐÑ	 ô\er   r"   c                   ó´   — e Zd Zd„ Z	 	 	 	 d	deeej                  f   dee   dz  dej                  j                  dz  dedz  deeej                  f   f
d„Zy)
ÚTorchAoDeserializec                 ó   — || _         y r$   r%   r'   s     r   r(   zTorchAoDeserialize.__init__®   r)   r   Nr;   Úsource_patternsr<   r=   r>   c           
      óÜ  — t        |j                  «       «      d   |v}i }dj                  |j                  d«      dd «      }	|r"t	        |d   t         «      r	|d   d   }
nZ|d   }
nT|j                  «       D ]A  }t        ||   «      dk7  rt        d|› dt        ||   «      › d	�«      ‚||   d   ||	› d|› �<   ŒC |r|
iS t        | j                  j                  «      st        d
«      ‚t        || j                  j                  «      \  }}|rJ ‚||   }t        ||«      \  }}t	        |t        j                  j                  «      rt        j                   t"        |«      |_        ||iS )a&  
        Consolidates tensor subclass components before reconstructing the object

        For example:
            input_dict: {
                "_weight_qdata": torch.Tensor,
                "_weight_scale": torch.Tensor,
            }
            full_layer_name: "model.layers.0.self_attn.k_proj.weight"

            Given this, we reconstruct a Float8Tensor instance using the qdata and scale
            and return it as a dictionary with the full_layer_name as the key and the recovered
            Float8Tensor instance as the value.
        r   rD   Néÿÿÿÿr   r   zExpected a single tensor for z	 but got z tensors insteadz$Invalid torchao safetensors metadata)rP   ÚkeysÚjoinÚsplitr   ÚlenÚ
ValueErrorr
   r&   Úmetadatar	   r   rQ   rR   ÚLinearÚtypesÚ
MethodTyper    Ú
extra_repr)r   r;   r}   r<   r=   rf   r8   Úis_unsafe_serializationÚ
param_dataÚ
layer_namer   ÚsuffixÚunflattened_state_dictÚleftover_state_dictÚ	new_paramr5   rg   s                    r   rs   zTorchAoDeserialize.convert±   s’  € ô. #' z§¡Ó'8Ó"9¸!Ñ"<ÀOÐ"SÐàˆ
Ø—X‘X˜o×3Ñ3°CÓ8¸¸"Ð=Ó>ˆ
Ù"Ü˜* XÑ.´Ô5Ø# HÑ-¨aÑ0‘à# HÑ-‘à$Ÿ/™/Ó+ò M�Ü�z &Ñ)Ó*¨aÒ/Ü$Ø7¸°x¸yÌÈZÐX^ÑM_ÓI`ÐHaÐaqÐróð ð 8BÀ&Ñ7IÈ!Ñ7L�
˜j˜\¨¨6¨(Ð3Ò4ðMñ #Ø# VÐ,Ð,Ü$ T×%6Ñ%6×%?Ñ%?Ô@ÜÐCÓDÐDä6QØ˜×)Ñ)×2Ñ2ó7
Ñ3ÐÐ 3ñ 'Ð&Ð&Ø*¨?Ñ;ˆ	ä(¨°Ó@‰	ˆ�ä�fœeŸh™hŸo™oÔ.Ü %× 0Ñ 0Ô1CÀVÓ LˆFÔà Ð+Ð+r   )NNNN)r   rt   ru   r(   rv   rw   rQ   rx   rP   rR   ry   rs   rH   r   r   r{   r{   ­   s€   „ ò)ð -1Ø(,Ø&*Øñ9,à˜˜eŸl™lÐ*Ñ+ð9,ð ˜c™ TÑ)ð9,ð �x‰x�‰ Ñ%ð	9,ð
 ˜t™ð9,ð 
ˆc�5—<‘<ÐÑ	 ô9,r   r{   )r`   r‡   rQ   Útransformers.utilsr   Útransformers.utils.import_utilsr   r   Úcore_model_loadingr   Úquantizers.quantizers_utilsr   Ú1torchao.prototype.safetensors.safetensors_supportr	   Ú/torchao.prototype.safetensors.safetensors_utilsr
   Ú
get_loggerr   Úloggerr   r    r"   r{   rH   r   r   ú<module>r™      sr   ðó 
Û ã å &ß Tñ ÔÝ2Ý >ñ Ôõõ Tà	ˆ×	Ñ	˜HÓ	%€òPòkôqe�mô qeôh=,˜õ =,r   