§
    ‚ŠtjÞ  ã                   ód  — d dl mZ d dlZd dlmZ ddlmZ ddlmZ ddl	m
Z
mZmZmZmZ ddlmZ d	d
lmZ  ej        e¦  «        ZdZ G d„ de¦  «        Z G d„ de¦  «        Zd„ Zdd„Ze G d„ de
¦  «        ¦   «         Z G d„ de¦  «        Z G d„ de¦  «        Z G d„ de¦  «        Zg d¢ZdS )é    )ÚOptionalNé   )Úlogging)Úno_inherit_decoratoré   )ÚLlamaAttentionÚLlamaForCausalLMÚLlamaForSequenceClassificationÚLlamaForTokenClassificationÚLlamaRotaryEmbedding)ÚPhi3MLPé   )Ú	GlmConfigzTHUDM/glm-4-9bc                   ó   — e Zd ZdS )ÚGlmMLPN©Ú__name__Ú
__module__Ú__qualname__© ó    úa/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/glm/modular_glm.pyr   r   &   ó   € € € € € Ø€Dr   r   c                   óf   — e Zd Ze	 	 	 d	dedz  ded         dedz  dedef         fd„¦   «         Z	dS )
ÚGlmRotaryEmbeddingNÚconfigÚdeviceztorch.deviceÚseq_lenÚreturnztorch.Tensorc                 óV  — | j         d         }| j                              dd¦  «        }t          | dd¦  «        p| j        | j        z  }t          ||z  ¦  «        }d}d|t          j        d|dt          j        ¬¦  «         	                    |t          j
        ¬	¦  «        |z  z  z  }||fS )
a¨  
        Computes the inverse frequencies according to the original RoPE implementation
        Args:
            config ([`~transformers.PreTrainedConfig`]):
                The model configuration.
            device (`torch.device`):
                The device to use for initialization of the inverse frequencies.
            seq_len (`int`, *optional*):
                The current sequence length. Unused for this type of RoPE.
        Returns:
            Tuple of (`torch.Tensor`, `float`), containing the inverse frequencies for the RoPE embeddings and the
            post-processing scaling factor applied to the computed cos/sin (unused in this type of RoPE).
        Ú
rope_thetaÚpartial_rotary_factorg      ð?Úhead_dimNr   r   )Údtype)r   r$   )Úrope_parametersÚgetÚgetattrÚhidden_sizeÚnum_attention_headsÚintÚtorchÚarangeÚint64ÚtoÚfloat)	r   r   r   Úbaser"   r#   ÚdimÚattention_factorÚinv_freqs	            r   Úcompute_default_rope_parametersz2GlmRotaryEmbedding.compute_default_rope_parameters+   sº   € ð& Ô% lÔ3ˆØ &Ô 6× :Ò :Ð;RÐTWÑ XÔ XÐÝ˜6 :¨tÑ4Ô4Ðh¸Ô8JÈfÔNhÑ8hˆÝ�(Ð2Ñ2Ñ3Ô3ˆàÐð Ø•U”\ ! S¨!µ5´;Ð?Ñ?Ô?×BÒBÈ&ÕX]ÔXcÐBÑdÔdÐgjÑjÑkñ
ˆð Ð)Ð)Ð)r   )NNN)
r   r   r   Ústaticmethodr   r   r*   Útupler/   r4   r   r   r   r   r   *   s|   € € € € € Øà#'Ø+/Ø"ð*ð *Ø˜DÑ ð*à˜Ô(ð*ð �t‘ð*ð 
ˆ~˜uÐ$Ô	%ð	*ð *ð *ñ „\ð*ð *ð *r   r   c                 óŽ   — | dddd…f         }| dddd…f         }t          j        | |fd¬¦  «                             d¦  «        S )	z*Rotates half the hidden dims of the input..r   Nr   r   éÿÿÿÿ©r1   éþÿÿÿ)r+   ÚstackÚflatten)ÚxÚx1Úx2s      r   Úrotate_halfr@   L   sQ   € à	
ˆ3���1�ˆ9Œ€BØ	
ˆ3���1�ˆ9Œ€BÝŒ;˜˜˜R�y bÐ)Ñ)Ô)×1Ò1°"Ñ5Ô5Ð5r   c                 óT  — |                      |¦  «        }|                      |¦  «        }|dd|j        d         dz  …f                              dd¬¦  «        }|dd|j        d         dz  …f                              dd¬¦  «        }|j        d         }| dd|…f         | d|d…f         }}|dd|…f         |d|d…f         }	}||z  t          |¦  «        |z  z   }
||z  t          |¦  «        |z  z   }t	          j        |
|gd¬¦  «        }
t	          j        ||	gd¬¦  «        }|
|fS )a…  Applies Rotary Position Embedding to the query and key tensors.

    Args:
        q (`torch.Tensor`): The query tensor.
        k (`torch.Tensor`): The key tensor.
        cos (`torch.Tensor`): The cosine part of the rotary embedding.
        sin (`torch.Tensor`): The sine part of the rotary embedding.
        unsqueeze_dim (`int`, *optional*, defaults to 1):
            The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and
            sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note
            that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and
            k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes
            cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have
            the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2.
    Returns:
        `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.
    .Nr8   r   r9   )Ú	unsqueezeÚshapeÚrepeat_interleaver@   r+   Úcat)ÚqÚkÚcosÚsinÚunsqueeze_dimÚ
rotary_dimÚq_rotÚq_passÚk_rotÚk_passÚq_embedÚk_embeds               r   Úapply_rotary_pos_embrR   S   s\  € ð$ �-Š-˜Ñ
&Ô
&€CØ
�-Š-˜Ñ
&Ô
&€Cð ˆcÐ'�S”Y˜r”] aÑ'Ð'Ð'Ô
(×
:Ò
:¸1À"Ð
:Ñ
EÔ
E€CØ
ˆcÐ'�S”Y˜r”] aÑ'Ð'Ð'Ô
(×
:Ò
:¸1À"Ð
:Ñ
EÔ
E€Cð ”˜2”€JØ�c˜;˜J˜;Ð&Ô'¨¨3°
°°Ð+;Ô)<ˆ6€EØ�c˜;˜J˜;Ð&Ô'¨¨3°
°°Ð+;Ô)<ˆ6€Eð �s‰{�{¨5Ñ1Ô1°CÑ7Ñ8€GØ�s‰{�{¨5Ñ1Ô1°CÑ7Ñ8€Gõ Œi˜ &Ð)¨rÐ2Ñ2Ô2€GÝŒi˜ &Ð)¨rÐ2Ñ2Ô2€GØ�GÐÐr   c                   ó0   ‡ — e Zd Zddededz  fˆ fd„Zˆ xZS )ÚGlmAttentionNr   Ú	layer_idxc                 ó¨   •— t          ¦   «                              ||¦  «         t          j        |j        | j        z  |j        d¬¦  «        | _        d S )NF)Úbias)ÚsuperÚ__init__ÚnnÚLinearr)   r#   r(   Úo_proj)Úselfr   rU   Ú	__class__s      €r   rY   zGlmAttention.__init__}   sG   ø€ Ý‰Œ×Ò˜ Ñ+Ô+Ð+Ý”i Ô :¸T¼]Ñ JÈFÔL^ÐejÐkÑkÔkˆŒˆˆr   )N)r   r   r   r   r*   rY   Ú__classcell__)r^   s   @r   rT   rT   {   sa   ø€ € € € € ðlð l˜yð l°S¸4±Zð lð lð lð lð lð lð lð lð lð lr   rT   c                   ó   — e Zd ZdS )ÚGlmForCausalLMNr   r   r   r   ra   ra   ‚   r   r   ra   c                   ó   — e Zd ZdS )ÚGlmForSequenceClassificationNr   r   r   r   rc   rc   †   r   r   rc   c                   ó   — e Zd ZdS )ÚGlmForTokenClassificationNr   r   r   r   re   re   Š   r   r   re   )ÚGlmPreTrainedModelÚGlmModelra   rc   re   )r   ) Útypingr   r+   Útorch.nnrZ   Úutilsr   Úutils.genericr   Úllama.modeling_llamar   r	   r
   r   r   Úphi3.modeling_phi3r   Úconfiguration_glmr   Ú
get_loggerr   ÚloggerÚ_CHECKPOINT_FOR_DOCr   r   r@   rR   rT   ra   rc   re   Ú__all__r   r   r   ú<module>rs      s-  ðð Ð Ð Ð Ð Ð à €€€Ø Ð Ð Ð Ð Ð à Ð Ð Ð Ð Ð Ø 1Ð 1Ð 1Ð 1Ð 1Ð 1ðð ð ð ð ð ð ð ð ð ð ð ð ð ð )Ð (Ð (Ð (Ð (Ð (Ø (Ð (Ð (Ð (Ð (Ð (ð 
ˆÔ	˜HÑ	%Ô	%€à&Ð ð	ð 	ð 	ð 	ð 	ˆWñ 	ô 	ð 	ð*ð *ð *ð *ð *Ð-ñ *ô *ð *ðD6ð 6ð 6ð%ð %ð %ð %ðP ðlð lð lð lð l�>ñ lô lñ Ôðlð	ð 	ð 	ð 	ð 	Ð%ñ 	ô 	ð 	ð	ð 	ð 	ð 	ð 	Ð#Añ 	ô 	ð 	ð	ð 	ð 	ð 	ð 	Ð ;ñ 	ô 	ð 	ðð ð €€€r   