§
    ‚ŠtjÕ  ã                   óì  — d dl 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 ddlmZ ddlmZmZmZmZmZ d	d
lmZ  ej        e¦  «        Z G d„ dej        ¦  «        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!e¦  «        Z" G d„ de	¦  «        Z# G d„ de
¦  «        Z$ G d„ d e¦  «        Z%g d!¢Z&dS )#é    Né   )Úlogging)Úno_inherit_decoratoré   )ÚGemmaForCausalLMÚGemmaForSequenceClassificationÚGemmaForTokenClassification)ÚGraniteAttention)ÚLlamaDecoderLayerÚLlamaMLPÚ
LlamaModelÚLlamaPreTrainedModelÚLlamaRotaryEmbeddingé   )ÚHeliumConfigc                   ó,   ‡ — e Zd Zdˆ fd„	Zd„ Zd„ Zˆ xZS )ÚHeliumRMSNormç�íµ ÷Æ°>c                 ó¬   •— t          ¦   «                              ¦   «          t          j        t	          j        |¦  «        ¦  «        | _        || _        d S ©N)ÚsuperÚ__init__ÚnnÚ	ParameterÚtorchÚonesÚweightÚvariance_epsilon)ÚselfÚhidden_sizeÚepsÚ	__class__s      €úg/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/helium/modular_helium.pyr   zHeliumRMSNorm.__init__    sB   ø€ Ý‰Œ×ÒÑÔÐÝ”l¥5¤:¨kÑ#:Ô#:Ñ;Ô;ˆŒØ #ˆÔÐÐó    c                 óT  — |j         }|                     t          j        ¦  «        }|                     d¦  «                             dd¬¦  «        }|t          j        || j        z   ¦  «        z  }| j                             t          j        ¦  «        |z                       |¦  «        S )Nr   éÿÿÿÿT)Úkeepdim)	ÚdtypeÚtor   Úfloat32ÚpowÚmeanÚrsqrtr   r   )r   Úhidden_statesÚinput_dtypeÚvariances       r#   ÚforwardzHeliumRMSNorm.forward%   sŠ   € Ø#Ô)ˆØ%×(Ò(­¬Ñ7Ô7ˆØ ×$Ò$ QÑ'Ô'×,Ò,¨R¸Ð,Ñ>Ô>ˆØ%­¬°H¸tÔ?TÑ4TÑ(UÔ(UÑUˆØ”—’�uœ}Ñ-Ô-°Ñ=×AÒAÀ+ÑNÔNÐNr$   c                 óH   — t          | j        j        ¦  «        › d| j        › �S )Nz, eps=)Útupler   Úshaper   )r   s    r#   Ú
extra_reprzHeliumRMSNorm.extra_repr,   s&   € Ý˜œÔ)Ñ*Ô*ÐIÐI°$Ô2GÐIÐIÐIr$   )r   )Ú__name__Ú
__module__Ú__qualname__r   r1   r5   Ú__classcell__©r"   s   @r#   r   r      se   ø€ € € € € ð$ð $ð $ð $ð $ð $ð
Oð Oð OðJð Jð Jð Jð Jð Jð Jr$   r   c                   ó   — e Zd ZdS )ÚHeliumRotaryEmbeddingN©r6   r7   r8   © r$   r#   r<   r<   0   ó   € € € € € Ø€Dr$   r<   c                   ó   — e Zd ZdS )Ú	HeliumMLPNr=   r>   r$   r#   rA   rA   4   r?   r$   rA   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   r&   ©Údiméþÿÿÿ)r   ÚstackÚflatten)ÚxÚx1Úx2s      r#   Úrotate_halfrK   8   sQ   € à	
ˆ3���1�ˆ9Œ€BØ	
ˆ3���1�ˆ9Œ€BÝŒ;˜˜˜R�y bÐ)Ñ)Ô)×1Ò1°"Ñ5Ô5Ð5r$   c                 óz  — |                      |¦  «        }|                      |¦  «        }|dd|j        d         dz  …f                              dd¬¦  «        }|dd|j        d         dz  …f                              dd¬¦  «        }| |z  t          | ¦  «        |z  z   }||z  t          |¦  «        |z  z   }||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.
    .Nr&   r   rC   )Ú	unsqueezer4   Úrepeat_interleaverK   )ÚqÚkÚcosÚsinÚunsqueeze_dimÚq_embedÚk_embeds          r#   Úapply_rotary_pos_embrV   ?   sË   € ð$ �-Š-˜Ñ
&Ô
&€CØ
�-Š-˜Ñ
&Ô
&€Cð ˆcÐ'�S”Y˜r”] aÑ'Ð'Ð'Ô
(×
:Ò
:¸1À"Ð
:Ñ
EÔ
E€CØ
ˆcÐ'�S”Y˜r”] aÑ'Ð'Ð'Ô
(×
:Ò
:¸1À"Ð
:Ñ
EÔ
E€Cà�3‰w�; q™>œ>¨CÑ/Ñ0€GØ�3‰w�; q™>œ>¨CÑ/Ñ0€Gà�GÐÐr$   c                   ó0   ‡ — e Zd Zddededz  fˆ fd„Zˆ xZS )ÚHeliumAttentionNÚconfigÚ	layer_idxc                 óÚ   •— t          ¦   «                              ||¦  «         t          j        |j        |j        d¬¦  «        | _        dt          j        | j        ¦  «        z  | _	        d S )NF)Úbiasr   )
r   r   r   ÚLinearr    Úo_projÚmathÚsqrtÚhead_dimÚscaling©r   rY   rZ   r"   s      €r#   r   zHeliumAttention.__init__`   sW   ø€ Ý‰Œ×Ò˜ Ñ+Ô+Ð+Ý”i Ô 2°FÔ4FÈUÐSÑSÔSˆŒØ�4œ9 T¤]Ñ3Ô3Ñ3ˆŒˆˆr$   r   ©r6   r7   r8   r   Úintr   r9   r:   s   @r#   rX   rX   ^   sT   ø€ € € € € ð4ð 4˜|ð 4¸¸d¹
ð 4ð 4ð 4ð 4ð 4ð 4ð 4ð 4ð 4ð 4r$   rX   c                   ó0   ‡ — e Zd Zddededz  fˆ fd„Zˆ xZS )ÚHeliumDecoderLayerNrY   rZ   c                 óô   •— t          ¦   «                              ||¦  «         t          |¦  «        | _        t	          |j        |j        ¬¦  «        | _        t	          |j        |j        ¬¦  «        | _        d S )N©r!   )	r   r   rA   Úmlpr   r    Úrms_norm_epsÚinput_layernormÚpost_attention_layernormrc   s      €r#   r   zHeliumDecoderLayer.__init__g   sh   ø€ Ý‰Œ×Ò˜ Ñ+Ô+Ð+å˜VÑ$Ô$ˆŒÝ,¨VÔ-?ÀVÔEXÐYÑYÔYˆÔÝ(5°fÔ6HÈfÔNaÐ(bÑ(bÔ(bˆÔ%Ð%Ð%r$   r   rd   r:   s   @r#   rg   rg   f   sa   ø€ € € € € ðcð c˜|ð c¸¸d¹
ð cð cð cð cð cð cð cð cð cð cr$   rg   c                   ó   — e Zd ZdS )ÚHeliumPreTrainedModelNr=   r>   r$   r#   ro   ro   o   r?   r$   ro   c                   ó$   ‡ — e Zd Zdefˆ fd„Zˆ xZS )ÚHeliumModelrY   c                 ó0  •‡— t          ¦   «                              ‰¦  «         t          j        ˆfd„t	          ‰j        ¦  «        D ¦   «         ¦  «        | _        t          ‰j        ‰j	        ¬¦  «        | _
        d| _        |                      ¦   «          d S )Nc                 ó0   •— g | ]}t          ‰|¦  «        ‘ŒS r>   )rg   )Ú.0rZ   rY   s     €r#   ú
<listcomp>z(HeliumModel.__init__.<locals>.<listcomp>w   s$   ø€ ÐdÐdÐd°yÕ ¨	Ñ2Ô2ÐdÐdÐdr$   ri   F)r   r   r   Ú
ModuleListÚrangeÚnum_hidden_layersÚlayersr   r    rk   ÚnormÚgradient_checkpointingÚ	post_init)r   rY   r"   s    `€r#   r   zHeliumModel.__init__t   s�   øø€ Ý‰Œ×Ò˜Ñ Ô Ð Ý”mØdÐdÐdÐdÅEÈ&ÔJbÑDcÔDcÐdÑdÔdñ
ô 
ˆŒõ " &Ô"4¸&Ô:MÐNÑNÔNˆŒ	Ø&+ˆÔ#ð 	�ŠÑÔÐÐÐr$   )r6   r7   r8   r   r   r9   r:   s   @r#   rq   rq   s   sD   ø€ € € € € ð	˜|ð 	ð 	ð 	ð 	ð 	ð 	ð 	ð 	ð 	ð 	r$   rq   c                   ó   — e Zd ZdS )ÚHeliumForCausalLMNr=   r>   r$   r#   r~   r~   €   r?   r$   r~   c                   ó   — e Zd ZdS )ÚHeliumForSequenceClassificationNr=   r>   r$   r#   r€   r€   „   r?   r$   r€   c                   ó   — e Zd ZdS )ÚHeliumForTokenClassificationNr=   r>   r$   r#   r‚   r‚   ˆ   r?   r$   r‚   )ro   rq   r~   r€   r‚   )r   )'r_   r   Útorch.nnr   Úutilsr   Úutils.genericr   Úgemma.modeling_gemmar   r   r	   Úgranite.modeling_graniter
   Úllama.modeling_llamar   r   r   r   r   Úconfiguration_heliumr   Ú
get_loggerr6   ÚloggerÚModuler   r<   rA   rK   rV   rX   rg   ro   rq   r~   r€   r‚   Ú__all__r>   r$   r#   ú<module>rŽ      sÔ  ðð €€€à €€€Ø Ð Ð Ð Ð Ð à Ð Ð Ð Ð Ð Ø 1Ð 1Ð 1Ð 1Ð 1Ð 1Ø pÐ pÐ pÐ pÐ pÐ pÐ pÐ pÐ pÐ pØ 7Ð 7Ð 7Ð 7Ð 7Ð 7Ø vÐ vÐ vÐ vÐ vÐ vÐ vÐ vÐ vÐ vÐ vÐ vÐ vÐ vØ .Ð .Ð .Ð .Ð .Ð .ð 
ˆÔ	˜HÑ	%Ô	%€ðJð Jð Jð Jð J�B”Iñ Jô Jð Jð"	ð 	ð 	ð 	ð 	Ð0ñ 	ô 	ð 	ð	ð 	ð 	ð 	ð 	�ñ 	ô 	ð 	ð6ð 6ð 6ðð ð ð ð> ð4ð 4ð 4ð 4ð 4Ð&ñ 4ô 4ñ Ôð4ðcð cð cð cð cÐ*ñ cô cð cð	ð 	ð 	ð 	ð 	Ð0ñ 	ô 	ð 	ð
ð 
ð 
ð 
ð 
Ð'¨ñ 
ô 
ð 
ð	ð 	ð 	ð 	ð 	Ð(ñ 	ô 	ð 	ð	ð 	ð 	ð 	ð 	Ð&Dñ 	ô 	ð 	ð	ð 	ð 	ð 	ð 	Ð#>ñ 	ô 	ð 	ðð ð €€€r$   