§
    ‚Štj€  ã                   ó†  — d dl mZ d dlmZmZ d dlZd dlmc mZ	 d dlmZ ddl
mZ ddlmZ ddlmZmZ dd	lmZ dd
lmZmZmZmZ ddlmZ ddlmZ ddlmZmZ ddl m!Z!m"Z" ddl#m$Z$m%Z% ddl&m'Z' ddl(m)Z)m*Z* ddl+m,Z,m-Z-m.Z. ddl/m0Z0 ddl1m2Z2  G d„ ded¬¦  «        Z3 G d„ dej4        ¦  «        Z5 ed¦  «         G d„ dej4        ¦  «        ¦   «         Z6 G d„ d ej4        ¦  «        Z7e G d!„ d"ej4        ¦  «        ¦   «         Z8 G d#„ d$ej4        ¦  «        Z9d%„ Z: ed&¦  «        dGd'„¦   «         Z;d(ej<        d)e=d*ej<        fd+„Z>	 dHd-ej4        d.ej<        d/ej<        d0ej<        d1ej<        dz  d2e?d3e?d4e'e)         fd5„Z@ ee;¦  «         G d6„ d7ej4        ¦  «        ¦   «         ZA G d8„ d9e¦  «        ZBe* G d:„ d;e%¦  «        ¦   «         ZC G d<„ d=ej4        ¦  «        ZDe* G d>„ d?eC¦  «        ¦   «         ZE	 	 	 dIdAej<        eFej<                 z  dz  dBe=dz  d1ej<        dz  d*ej<        e=z  fdC„ZGe* G dD„ dEeCe¦  «        ¦   «         ZHg dF¢ZIdS )Jé    )ÚCallable)ÚOptionalÚ	TypedDictN)Únné   )Úinitialization)ÚACT2FN)ÚCacheÚDynamicCache)ÚGenerationMixin)Úuse_experts_implementationÚuse_kernel_forward_from_hubÚuse_kernel_func_from_hubÚuse_kernelized_func)Úcreate_causal_mask)ÚGradientCheckpointingLayer)ÚMoeCausalLMOutputWithPastÚMoeModelOutputWithPast)ÚROPE_INIT_FUNCTIONSÚdynamic_rope_update)ÚALL_ATTENTION_FUNCTIONSÚPreTrainedModel)ÚUnpack)ÚTransformersKwargsÚauto_docstring)Úcan_return_tupleÚmaybe_autocastÚmerge_with_config_defaults)Úcapture_outputsé   )ÚGraniteMoeSharedConfigc                   ód   — e Zd ZU dZej        ed<   ej        ed<   eed<   eed<   ej        ed<   dS )ÚGraniteFlashAttentionKwargsaT  
    Keyword arguments for advanced Flash Attention, causal-conv1d, and mamba_ssm kernel usage.
    Use cases include padding-free training and fewer `torch.compile` graph breaks.

    cu_seq_lens_q (`torch.LongTensor`):
        Gets cumulative sequence length for query state.
    cu_seq_lens_k (`torch.LongTensor`):
        Gets cumulative sequence length for key state.
    max_length_q (`int`):
        Maximum sequence length for query state.
    max_length_k (`int`):
        Maximum sequence length for key state.
    seq_idx (`torch.IntTensor):
        Index of each packed sequence.
    Úcu_seq_lens_qÚcu_seq_lens_kÚmax_length_qÚmax_length_kÚseq_idxN)	Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚtorchÚ
LongTensorÚ__annotations__ÚintÚ	IntTensor© ó    ú|/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/granitemoeshared/modeling_granitemoeshared.pyr#   r#   2   sb   € € € € € € ðð ð  Ô#Ð#Ð#Ñ#ØÔ#Ð#Ð#Ñ#ØÐÐÑØÐÐÑØŒ_ÐÐÑÐÐr3   r#   F)Útotalc                   óL   ‡ — e Zd ZdZdefˆ fd„Zdej        dej        fd„Zˆ xZ	S )ÚGraniteMoeSharedMLPz~
    MLP layer for shared experts

    Args:
        config:
            Configuration object with model hyperparameters.
    Úconfigc                 óD  •— t          ¦   «                              ¦   «          |j        | _        |j        | _        t
          |j                 | _        t          j	        | j        | j        dz  d¬¦  «        | _
        t          j	        | j        | j        d¬¦  «        | _        d S )Né   F©Úbias)ÚsuperÚ__init__Úhidden_sizeÚ
input_sizeÚshared_intermediate_sizer	   Ú
hidden_actÚ
activationr   ÚLinearÚinput_linearÚoutput_linear©Úselfr8   Ú	__class__s     €r4   r>   zGraniteMoeSharedMLP.__init__S   s…   ø€ Ý‰Œ×ÒÑÔÐà Ô,ˆŒØ!Ô:ˆÔÝ  Ô!2Ô3ˆŒÝœI d¤o°tÔ7GÈ!Ñ7KÐRWÐXÑXÔXˆÔÝœY tÔ'7¸¼ÈuÐUÑUÔUˆÔÐÐr3   Úhidden_statesÚreturnc                 óÐ   — |                       |¦  «        }|                     dd¬¦  «        }|                      |d         ¦  «        |d         z  }|                      |¦  «        }|S )Nr:   éÿÿÿÿ©Údimr   r    )rE   ÚchunkrC   rF   )rH   rJ   Úchunked_hidden_statess      r4   ÚforwardzGraniteMoeSharedMLP.forward\   sj   € Ø×)Ò)¨-Ñ8Ô8ˆØ -× 3Ò 3°A¸2Ð 3Ñ >Ô >ÐØŸšÐ(=¸aÔ(@ÑAÔAÐDYÐZ[ÔD\Ñ\ˆØ×*Ò*¨=Ñ9Ô9ˆØÐr3   ©
r)   r*   r+   r,   r!   r>   r-   ÚTensorrR   Ú__classcell__©rI   s   @r4   r7   r7   J   s|   ø€ € € € € ðð ðVÐ5ð Vð Vð Vð Vð Vð Vð U¤\ð °e´lð ð ð ð ð ð ð ð r3   r7   ÚRMSNormc                   óT   ‡ — e Zd Zd	deddfˆ fd„Zdej        dej        fd„Zd„ Zˆ xZ	S )
ÚGraniteMoeSharedRMSNormç�íµ ÷Æ°>ÚepsrK   Nc                 ó¬   •— t          ¦   «                              ¦   «          t          j        t	          j        |¦  «        ¦  «        | _        || _        dS )zF
        GraniteMoeSharedRMSNorm is equivalent to T5LayerNorm
        N)r=   r>   r   Ú	Parameterr-   ÚonesÚweightÚvariance_epsilon)rH   r?   r[   rI   s      €r4   r>   z GraniteMoeSharedRMSNorm.__init__f   sD   ø€ õ 	‰Œ×ÒÑÔÐÝ”l¥5¤:¨kÑ#:Ô#:Ñ;Ô;ˆŒØ #ˆÔÐÐr3   rJ   c                 ó  — |j         }|                     t          j        ¦  «        }|                     d¦  «                             dd¬¦  «        }|t          j        || j        z   ¦  «        z  }| j        |                     |¦  «        z  S )Nr:   rM   T)Úkeepdim)	ÚdtypeÚtor-   Úfloat32ÚpowÚmeanÚrsqrtr`   r_   )rH   rJ   Úinput_dtypeÚvariances       r4   rR   zGraniteMoeSharedRMSNorm.forwardn   s|   € Ø#Ô)ˆØ%×(Ò(­¬Ñ7Ô7ˆØ ×$Ò$ QÑ'Ô'×,Ò,¨R¸Ð,Ñ>Ô>ˆØ%­¬°H¸tÔ?TÑ4TÑ(UÔ(UÑUˆØŒ{˜]×-Ò-¨kÑ:Ô:Ñ:Ð:r3   c                 óH   — t          | j        j        ¦  «        › d| j        › �S )Nz, eps=)Útupler_   Úshaper`   )rH   s    r4   Ú
extra_reprz"GraniteMoeSharedRMSNorm.extra_repru   s&   € Ý˜œÔ)Ñ*Ô*ÐIÐI°$Ô2GÐIÐIÐIr3   )rZ   )
r)   r*   r+   Úfloatr>   r-   rT   rR   rn   rU   rV   s   @r4   rY   rY   d   sŒ   ø€ € € € € ð$ð $¨ð $¸$ð $ð $ð $ð $ð $ð $ð; U¤\ð ;°e´lð ;ð ;ð ;ð ;ðJð Jð Jð Jð Jð Jð Jr3   rY   c                   ór   ‡ — e Zd ZdZdefˆ fd„Zdej        deej        ej        ej        f         fd„Z	ˆ xZ
S )ÚGraniteMoeSharedTopKRoutera¡  Top-k gating that returns the routing decisions without grouping tokens by expert.

    Returns ``(top_k_index, top_k_weights, router_logits)``; the grouping/scattering used to live
    here (via ``expert_size.tolist()``, which broke fullgraph compile) and now happens inside the
    experts forward via ``use_experts_implementation`` so the default ``grouped_mm`` / ``batched_mm``
    paths can compile cleanly.
    r8   c                 óä   •— t          ¦   «                              ¦   «          |j        | _        |j        | _        t          j        t          j	        | j        |j
        ¦  «        ¦  «        | _        d S ©N)r=   r>   Únum_local_expertsÚnum_expertsÚnum_experts_per_tokÚtop_kr   r]   r-   Úemptyr?   r_   rG   s     €r4   r>   z#GraniteMoeSharedTopKRouter.__init__‚   sU   ø€ Ý‰Œ×ÒÑÔÐØ!Ô3ˆÔØÔ/ˆŒ
Ý”l¥5¤;¨tÔ/?ÀÔASÑ#TÔ#TÑUÔUˆŒˆˆr3   rJ   rK   c                 óô   — t          j        || j        ¦  «                             ¦   «         }|                     | j        d¬¦  «        \  }}t          j        |d¬¦  «                             |¦  «        }|||fS )NrM   rN   )	ÚFÚlinearr_   ro   Útopkrw   r-   ÚsoftmaxÚtype_as)rH   rJ   Úrouter_logitsÚtop_k_logitsÚtop_k_indexÚtop_k_weightss         r4   rR   z"GraniteMoeSharedTopKRouter.forwardˆ   so   € Ýœ °´Ñ<Ô<×BÒBÑDÔDˆØ$1×$6Ò$6°t´zÀrÐ$6Ñ$JÔ$JÑ!ˆ�kÝœ l¸Ð;Ñ;Ô;×CÒCÀMÑRÔRˆØ˜M¨=Ð8Ð8r3   )r)   r*   r+   r,   r!   r>   r-   rT   rl   rR   rU   rV   s   @r4   rq   rq   y   sŽ   ø€ € € € € ðð ðVÐ5ð Vð Vð Vð Vð Vð Vð9 U¤\ð 9°e¸E¼LÈ%Ì,ÐX]ÔXdÐ<dÔ6eð 9ð 9ð 9ð 9ð 9ð 9ð 9ð 9r3   rq   c                   óh   ‡ — e Zd ZdZdefˆ fd„Zdej        dej        dej        dej        fd„Zˆ xZ	S )	ÚGraniteMoeSharedExpertsz2Collection of expert weights stored as 3D tensors.r8   c                 ó´  •— t          ¦   «                              ¦   «          |j        | _        |j        | _        |j        | _        t          j	        t          j        | j        d| j        z  | j        ¦  «        ¦  «        | _        t          j	        t          j        | j        | j        | j        ¦  «        ¦  «        | _        t          |j                 | _        d S )Nr:   )r=   r>   rt   ru   r?   Ú
hidden_dimÚintermediate_sizeÚintermediate_dimr   r]   r-   rx   Úgate_up_projÚ	down_projr	   rB   Úact_fnrG   s     €r4   r>   z GraniteMoeSharedExperts.__init__“   s£   ø€ Ý‰Œ×ÒÑÔÐØ!Ô3ˆÔØ Ô,ˆŒØ &Ô 8ˆÔÝœL­¬°TÔ5EÀqÈ4ÔK`ÑG`ÐbfÔbqÑ)rÔ)rÑsÔsˆÔÝœ¥e¤k°$Ô2BÀDÄOÐUYÔUjÑ&kÔ&kÑlÔlˆŒÝ˜VÔ.Ô/ˆŒˆˆr3   rJ   r�   r‚   rK   c                 ó€  — t          j        |¦  «        }t          j        ¦   «         5  t           j        j                             || j        ¬¦  «        }|                     ddd¦  «        }t          j        | 	                    d¬¦  «        d¦  «         
                    ¦   «         }d d d ¦  «         n# 1 swxY w Y   |D ]þ}|d         }|| j        k    rŒt          j        ||         ¦  «        \  }}	||	         }
t          j                             |
| j        |         ¦  «                             dd¬¦  «        \  }}|                      |¦  «        |z  }t          j                             || j        |         ¦  «        }|||	|d f         z  }|                     d|	|                     |j        ¦  «        ¦  «         Œÿ|S )N)Únum_classesr:   r    r   )rM   éþÿÿÿrN   rM   )r-   Ú
zeros_likeÚno_gradr   Ú
functionalÚone_hotru   ÚpermuteÚgreaterÚsumÚnonzeroÚwherer{   r‰   rP   r‹   rŠ   Ú
index_add_rd   rc   )rH   rJ   r�   r‚   Úfinal_hidden_statesÚexpert_maskÚ
expert_hitÚ
expert_idxÚ	top_k_posÚ	token_idxÚcurrent_stateÚgateÚupÚcurrent_hidden_statess                 r4   rR   zGraniteMoeSharedExperts.forwardœ   sø  € õ $Ô.¨}Ñ=Ô=ÐÝŒ]‰_Œ_ð 	Sð 	SÝœ(Ô-×5Ò5°kÈtÔO_Ð5Ñ`Ô`ˆKØ%×-Ò-¨a°°AÑ6Ô6ˆKÝœ {§¢¸8 Ñ'DÔ'DÀaÑHÔH×PÒPÑRÔRˆJð	Sð 	Sð 	Sñ 	Sô 	Sð 	Sð 	Sð 	Sð 	Sð 	Sð 	Søøøð 	Sð 	Sð 	Sð 	Sð
 %ð 
	nð 
	nˆJØ# AœˆJØ˜TÔ-Ò-Ð-ØÝ#(¤;¨{¸:Ô/FÑ#GÔ#GÑ ˆI�yØ)¨)Ô4ˆMÝ”}×+Ò+¨M¸4Ô;LÈZÔ;XÑYÔY×_Ò_Ð`aÐgiÐ_ÑjÔj‰HˆD�"Ø$(§K¢K°Ñ$5Ô$5¸Ñ$:Ð!Ý$&¤M×$8Ò$8Ð9NÐPTÔP^Ð_iÔPjÑ$kÔ$kÐ!Ø$9¸MÈ)ÐU^Ð`dÐJdÔ<eÑ$eÐ!Ø×*Ò*¨1¨iÐ9N×9QÒ9QÐReÔRkÑ9lÔ9lÑmÔmÐmÐmà"Ð"s   ¨A>B2Â2B6Â9B6rS   rV   s   @r4   r„   r„   �   s�   ø€ € € € € à<Ð<ð0Ð5ð 0ð 0ð 0ð 0ð 0ð 0ð#à”|ð#ð ”\ð#ð ”|ð	#ð
 
Œð#ð #ð #ð #ð #ð #ð #ð #r3   r„   c                   óL   ‡ — e Zd ZdZdefˆ fd„Zdej        dej        fd„Zˆ xZ	S )ÚGraniteMoeSharedMoEzISparsely-gated mixture-of-experts block: router decides, experts compute.r8   c                 ó°   •— t          ¦   «                              ¦   «          |j        | _        t	          |¦  «        | _        t          |¦  «        | _        d S rs   )r=   r>   r?   r@   rq   Úrouterr„   ÚexpertsrG   s     €r4   r>   zGraniteMoeSharedMoE.__init__º   sE   ø€ Ý‰Œ×ÒÑÔÐØ Ô,ˆŒÝ0°Ñ8Ô8ˆŒÝ.¨vÑ6Ô6ˆŒˆˆr3   Úlayer_inputrK   c                 óö   — |                      ¦   «         \  }}}|                     d|¦  «        }|                      |¦  «        \  }}}|                      |||¦  «        }	|	                     ||| j        ¦  «        S )NrM   )ÚsizeÚreshaper¦   r§   Úviewr@   )
rH   r¨   ÚbszÚlengthÚemb_sizerJ   r�   r‚   Ú_Úlayer_outputs
             r4   rR   zGraniteMoeSharedMoE.forwardÀ   sv   € Ø +× 0Ò 0Ñ 2Ô 2ÑˆˆV�XØ#×+Ò+¨B°Ñ9Ô9ˆØ(,¯ª°MÑ(BÔ(BÑ%ˆ�] AØ—|’| M°;ÀÑNÔNˆØ× Ò   f¨d¬oÑ>Ô>Ð>r3   rS   rV   s   @r4   r¤   r¤   ·   sq   ø€ € € € € ØSÐSð7Ð5ð 7ð 7ð 7ð 7ð 7ð 7ð? 5¤<ð ?°E´Lð ?ð ?ð ?ð ?ð ?ð ?ð ?ð ?r3   r¤   c                 óœ   — | dd| j         d         dz  …f         }| d| j         d         dz  d…f         }t          j        | |fd¬¦  «        S )z*Rotates half the hidden dims of the input..NrM   r:   rN   )rm   r-   Úcat)ÚxÚx1Úx2s      r4   Úrotate_halfr·   È   s]   € à	
ˆ3Ð"�!”'˜"”+ Ñ"Ð"Ð"Ô	#€BØ	
ˆ3�”˜”˜qÑ Ð"Ð"Ð"Ô	#€BÝŒ9�r�c˜2�Y BÐ'Ñ'Ô'Ð'r3   Úrotary_pos_embc                 ó¾   — |                      |¦  «        }|                      |¦  «        }| |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.
    )Ú	unsqueezer·   )ÚqÚkÚcosÚsinÚunsqueeze_dimÚq_embedÚk_embeds          r4   Úapply_rotary_pos_embrÂ   Ï   sc   € ð& �-Š-˜Ñ
&Ô
&€CØ
�-Š-˜Ñ
&Ô
&€CØ�3‰w�; q™>œ>¨CÑ/Ñ0€GØ�3‰w�; q™>œ>¨CÑ/Ñ0€GØ�GÐÐr3   rJ   Ún_reprK   c                 ó¸   — | j         \  }}}}|dk    r| S | dd…dd…ddd…dd…f                              |||||¦  «        } |                      |||z  ||¦  «        S )zÔ
    This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
    num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
    r    N)rm   Úexpandr«   )rJ   rÃ   ÚbatchÚnum_key_value_headsÚslenÚhead_dims         r4   Ú	repeat_kvrÊ   é   s„   € ð
 2?Ô1DÑ.€EÐ  hØ�‚z€zØÐØ! ! ! ! Q Q Q¨¨a¨a¨a°°°Ð"2Ô3×:Ò:¸5ÐBUÐW\Ð^bÐdlÑmÔm€MØ× Ò  Ð(;¸eÑ(CÀTÈ8ÑTÔTÐTr3   ç        ÚmoduleÚqueryÚkeyÚvalueÚattention_maskÚscalingÚdropoutÚkwargsc                 ó  — t          || j        ¦  «        }t          || j        ¦  «        }	t          j        ||                     dd¦  «        ¦  «        |z  }
|�|
|z   }
t
          j                             |
dt          j        ¬¦  «         	                    |j
        ¦  «        }
t
          j                             |
|| j        ¬¦  «        }
t          j        |
|	¦  «        }|                     dd¦  «                             ¦   «         }||
fS )Nr:   r   rM   )rO   rc   )ÚpÚtrainingr    )rÊ   Únum_key_value_groupsr-   ÚmatmulÚ	transposer   r‘   r}   re   rd   rc   rÒ   rÖ   Ú
contiguous)rÌ   rÍ   rÎ   rÏ   rÐ   rÑ   rÒ   rÓ   Ú
key_statesÚvalue_statesÚattn_weightsÚattn_outputs               r4   Úeager_attention_forwardrß   õ   sé   € õ ˜3 Ô ;Ñ<Ô<€JÝ˜U FÔ$?Ñ@Ô@€Lå”<  z×';Ò';¸A¸qÑ'AÔ'AÑBÔBÀWÑL€LØÐ!Ø# nÑ4ˆå”=×(Ò(¨¸2ÅUÄ]Ð(ÑSÔS×VÒVÐW\ÔWbÑcÔc€LÝ”=×(Ò(¨¸È6Ì?Ð(Ñ[Ô[€LÝ”,˜|¨\Ñ:Ô:€KØ×'Ò'¨¨1Ñ-Ô-×8Ò8Ñ:Ô:€Kà˜Ð$Ð$r3   c                   óÎ   ‡ — e Zd ZdZdedefˆ fd„Z	 	 	 ddej        de	ej        ej        f         dz  dej        dz  d	e
dz  d
ee         de	ej        ej        f         fd„Zˆ xZS )ÚGraniteMoeSharedAttentionz=Multi-headed attention from 'Attention Is All You Need' paperr8   Ú	layer_idxc                 ó¨  •— t          ¦   «                              ¦   «          || _        || _        t	          |d|j        |j        z  ¦  «        | _        |j        |j        z  | _	        |j
        | _        |j        | _        d| _        t          j        |j        |j        | j        z  |j        ¬¦  «        | _        t          j        |j        |j        | j        z  |j        ¬¦  «        | _        t          j        |j        |j        | j        z  |j        ¬¦  «        | _        t          j        |j        | j        z  |j        |j        ¬¦  «        | _        d S )NrÉ   Tr;   )r=   r>   r8   râ   Úgetattrr?   Únum_attention_headsrÉ   rÇ   r×   Úattention_multiplierrÑ   Úattention_dropoutÚ	is_causalr   rD   Úattention_biasÚq_projÚk_projÚv_projÚo_proj©rH   r8   râ   rI   s      €r4   r>   z"GraniteMoeSharedAttention.__init__  s>  ø€ Ý‰Œ×ÒÑÔÐØˆŒØ"ˆŒÝ ¨
°FÔ4FÈ&ÔJdÑ4dÑeÔeˆŒØ$*Ô$>À&ÔB\Ñ$\ˆÔ!ØÔ2ˆŒØ!'Ô!9ˆÔØˆŒå”iØÔ Ô :¸T¼]Ñ JÐQWÔQfð
ñ 
ô 
ˆŒõ ”iØÔ Ô :¸T¼]Ñ JÐQWÔQfð
ñ 
ô 
ˆŒõ ”iØÔ Ô :¸T¼]Ñ JÐQWÔQfð
ñ 
ô 
ˆŒõ ”iØÔ&¨¬Ñ6¸Ô8JÐQWÔQfð
ñ 
ô 
ˆŒˆˆr3   NrJ   Úposition_embeddingsrÐ   Úpast_key_valuesrÓ   rK   c                 ó"  — |j         d d…         }g |¢d‘| j        ‘R }|                      |¦  «                             |¦  «                             dd¦  «        }|                      |¦  «                             |¦  «                             dd¦  «        }	|                      |¦  «                             |¦  «                             dd¦  «        }
|\  }}t          ||	||¦  «        \  }}	|�|                     |	|
| j	        ¦  «        \  }	}
t          j        | j        j        t          ¦  «        } || ||	|
|f| j        sdn| j        | j        dœ|¤Ž\  }} |j        g |¢d‘R Ž                      ¦   «         }|                      |¦  «        }||fS )NrM   r    r:   rË   )rÒ   rÑ   )rm   rÉ   rê   r¬   rÙ   rë   rì   rÂ   Úupdaterâ   r   Úget_interfacer8   Ú_attn_implementationrß   rÖ   rç   rÑ   r«   rÚ   rí   )rH   rJ   rï   rÐ   rð   rÓ   Úinput_shapeÚhidden_shapeÚquery_statesrÛ   rÜ   r½   r¾   Úattention_interfacerÞ   rÝ   s                   r4   rR   z!GraniteMoeSharedAttention.forward)  sÀ  € ð $Ô)¨#¨2¨#Ô.ˆØ8˜Ð8 bÐ8¨$¬-Ð8Ð8ˆà—{’{ =Ñ1Ô1×6Ò6°|ÑDÔD×NÒNÈqÐRSÑTÔTˆØ—[’[ Ñ/Ô/×4Ò4°\ÑBÔB×LÒLÈQÐPQÑRÔRˆ
Ø—{’{ =Ñ1Ô1×6Ò6°|ÑDÔD×NÒNÈqÐRSÑTÔTˆà&‰ˆˆSÝ#7¸ÀjÐRUÐWZÑ#[Ô#[Ñ ˆ�jàÐ&Ø'6×'=Ò'=¸jÈ,ÐX\ÔXfÑ'gÔ'gÑ$ˆJ˜å(?Ô(MØŒKÔ,Õ.Eñ)
ô )
Ðð %8Ð$7ØØØØØð	%
ð  $œ}ÐH�C�C°$Ô2HØ”Lð	%
ð 	%
ð ð	%
ð 	%
Ñ!ˆ�\ð *�kÔ)Ð;¨;Ð;¸Ð;Ð;Ð;×FÒFÑHÔHˆØ—k’k +Ñ.Ô.ˆØ˜LÐ(Ð(r3   ©NNN)r)   r*   r+   r,   r!   r0   r>   r-   rT   rl   r
   r   r   rR   rU   rV   s   @r4   rá   rá     sæ   ø€ € € € € àGÐGð
Ð5ð 
À#ð 
ð 
ð 
ð 
ð 
ð 
ð4 IMØ.2Ø(,ð&)ð &)à”|ð&)ð # 5¤<°´Ð#=Ô>ÀÑEð&)ð œ tÑ+ð	&)ð
  ™ð&)ð Ð+Ô,ð&)ð 
ˆuŒ|˜Uœ\Ð)Ô	*ð&)ð &)ð &)ð &)ð &)ð &)ð &)ð &)r3   rá   c                   ó  ‡ — e Zd Zdedefˆ fd„Z	 	 	 	 	 	 ddej        dej        dz  dej        dz  d	e	dz  d
e
dz  de
dz  deej        ej        f         dz  dee         deej        eej        ej        f         dz  f         fd„Zˆ xZS )ÚGraniteMoeSharedDecoderLayerr8   râ   c                 óŽ  •— t          ¦   «                              ¦   «          |j        | _        t          ||¬¦  «        | _        t          |j        |j        ¬¦  «        | _        t          |j        |j        ¬¦  «        | _        t          |¦  «        | _
        |j        | _        |j        dk    rd nt          |¦  «        | _        d S )N)r8   râ   ©r[   r   )r=   r>   r?   rá   Ú	self_attnrY   Úrms_norm_epsÚinput_layernormÚpost_attention_layernormr¤   Úblock_sparse_moeÚresidual_multiplierrA   r7   Ú
shared_mlprî   s      €r4   r>   z%GraniteMoeSharedDecoderLayer.__init__S  s°   ø€ Ý‰Œ×ÒÑÔÐØ!Ô-ˆÔÝ2¸&ÈIÐVÑVÔVˆŒÝ6°vÔ7IÈvÔObÐcÑcÔcˆÔÝ(?ÀÔ@RÐX^ÔXkÐ(lÑ(lÔ(lˆÔ%Ý 3°FÑ ;Ô ;ˆÔØ#)Ô#=ˆÔ Ø"(Ô"AÀQÒ"FÐ"F˜$˜$ÕL_Ð`fÑLgÔLgˆŒˆˆr3   NFrJ   rÐ   Úposition_idsrð   Úoutput_attentionsÚ	use_cacherï   rÓ   rK   c                 ó4  — |}	|                       |¦  «        } | j        d|||||||dœ|¤Ž\  }}
|	|| j        z  z   }|}	|                      |¦  «        }|                      |¦  «        }| j        €|}n||                      |¦  «        z   }|	|| j        z  z   }|S )N)rJ   rÐ   r  rð   r  r  rï   r2   )r   rþ   r  r  r  r  )rH   rJ   rÐ   r  rð   r  r  rï   rÓ   Úresidualr°   Úmoe_hidden_statess               r4   rR   z$GraniteMoeSharedDecoderLayer.forward]  sÜ   € ð !ˆØ×,Ò,¨]Ñ;Ô;ˆð *˜4œ>ð 	
Ø'Ø)Ø%Ø+Ø/ØØ 3ð	
ð 	
ð ð	
ð 	
Ñˆ�qð ! =°4Ô3KÑ#KÑKˆà ˆØ×5Ò5°mÑDÔDˆØ ×1Ò1°-Ñ@Ô@ÐàŒ?Ð"Ø-ˆMˆMà-°·²ÀÑ0NÔ0NÑNˆMØ  =°4Ô3KÑ#KÑKˆØÐr3   )NNNFFN)r)   r*   r+   r!   r0   r>   r-   rT   r.   r
   Úboolrl   r   r#   ÚFloatTensorrR   rU   rV   s   @r4   rû   rû   R  s2  ø€ € € € € ðhÐ5ð hÀ#ð hð hð hð hð hð hð /3Ø04Ø(,Ø).Ø!&ØHLð%ð %à”|ð%ð œ tÑ+ð%ð Ô&¨Ñ-ð	%ð
  ™ð%ð   $™;ð%ð ˜$‘;ð%ð # 5¤<°´Ð#=Ô>ÀÑEð%ð Ð4Ô5ð%ð 
ˆuÔ  %¨Ô(9¸5Ô;LÐ(LÔ"MÐPTÑ"TÐTÔ	Uð%ð %ð %ð %ð %ð %ð %ð %r3   rû   c                   ó†   ‡ — e Zd ZU eed<   dZdZdgZdgZdZ	dZ
dZdZdZeedœZ ej        ¦   «         ˆ fd„¦   «         Zˆ xZS )ÚGraniteMoeSharedPreTrainedModelr8   ÚmodelTrû   rð   )rJ   Ú
attentionsc                 óŠ  •— t          ¦   «                              |¦  «         t          |t          ¦  «        rNt	          j        |j        d| j        j        ¬¦  «         t	          j        |j	        d| j        j        ¬¦  «         d S t          |t          ¦  «        r(t	          j        |j        d| j        j        ¬¦  «         d S d S )NrË   )rg   Ústd)r=   Ú_init_weightsÚ
isinstancer„   ÚinitÚnormal_r‰   r8   Úinitializer_rangerŠ   rq   r_   )rH   rÌ   rI   s     €r4   r  z-GraniteMoeSharedPreTrainedModel._init_weights–  s·   ø€ å‰Œ×Ò˜fÑ%Ô%Ð%Ý�fÕ5Ñ6Ô6ð 	UÝŒL˜Ô,°3¸D¼KÔ<YÐZÑZÔZÐZÝŒL˜Ô)°¸¼Ô9VÐWÑWÔWÐWÐWÐWÝ˜Õ :Ñ;Ô;ð 	UÝŒL˜œ¨S°d´kÔ6SÐTÑTÔTÐTÐTÐTð	Uð 	Ur3   )r)   r*   r+   r!   r/   Úbase_model_prefixÚsupports_gradient_checkpointingÚ_no_split_modulesÚ_skip_keys_device_placementÚ_supports_flash_attnÚ_supports_sdpaÚ_supports_flex_attnÚ_can_compile_fullgraphÚ_supports_attention_backendrû   rá   Ú_can_record_outputsr-   r�   r  rU   rV   s   @r4   r  r  …  s±   ø€ € € € € € à"Ð"Ð"Ñ"ØÐØ&*Ð#Ø7Ð8ÐØ#4Ð"5ÐØÐØ€NØÐØ!ÐØ"&Ðà5Ø/ðð Ðð
 €U„]�_„_ðUð Uð Uð Uñ „_ðUð Uð Uð Uð Ur3   r  c                   óÔ   ‡ — e Zd ZU ej        ed<   ddefˆ fd„Ze	 	 	 ddedz  de	d         de
dz  ded	ef         fd
„¦   «         Z ej        ¦   «         ed„ ¦   «         ¦   «         Zˆ xZS )ÚGraniteMoeSharedRotaryEmbeddingÚinv_freqNr8   c                 ó²  •— t          ¦   «                              ¦   «          |j        | _        |j        | _        || _        | j        j        d         | _        | j        }| j        dk    rt          | j                 } || j        |¦  «        \  }| _
        |                      d|d¬¦  «         |                      d|                     ¦   «         d¬¦  «         d S )NÚ	rope_typeÚdefaultr$  F)Ú
persistentÚoriginal_inv_freq)r=   r>   Úmax_position_embeddingsÚmax_seq_len_cachedÚoriginal_max_seq_lenr8   Úrope_parametersr&  Úcompute_default_rope_parametersr   Úattention_scalingÚregister_bufferÚclone)rH   r8   ÚdeviceÚrope_init_fnr$  rI   s        €r4   r>   z(GraniteMoeSharedRotaryEmbedding.__init__£  sÊ   ø€ Ý‰Œ×ÒÑÔÐØ"(Ô"@ˆÔØ$*Ô$BˆÔ!àˆŒàœÔ4°[ÔAˆŒØ!%Ô!EˆØŒ>˜YÒ&Ð&Ý.¨t¬~Ô>ˆLØ+7¨<¸¼ÀVÑ+LÔ+LÑ(ˆ�$Ô(à×Ò˜Z¨¸eÐÑDÔDÐDØ×ÒÐ0°(·.².Ñ2BÔ2BÈuÐÑUÔUÐUÐUÐUr3   r2  ztorch.deviceÚseq_lenrK   ztorch.Tensorc                 óü   — | j         d         }t          | dd¦  «        p| j        | j        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_thetarÉ   Ng      ð?r   r:   ©rc   )r2  rc   )	r-  rä   r?   rå   r-   ÚarangeÚint64rd   ro   )r8   r2  r4  ÚbaserO   Úattention_factorr$  s          r4   r.  z?GraniteMoeSharedRotaryEmbedding.compute_default_rope_parameters³  sŒ   € ð& Ô% lÔ3ˆÝ�f˜j¨$Ñ/Ô/Ðc°6Ô3EÈÔIcÑ3cˆàÐð Ø•U”\ ! S¨!µ5´;Ð?Ñ?Ô?×BÒBÈ&ÕX]ÔXcÐBÑdÔdÐgjÑjÑkñ
ˆð Ð)Ð)Ð)r3   c                 óN  — | j         d d d …d f                              ¦   «                              |j        d         dd¦  «                             |j        ¦  «        }|d d …d d d …f                              ¦   «         }t          |j        j        t          ¦  «        r|j        j        dk    r|j        j        nd}t          |d¬¦  «        5  |                     ¦   «         |                     ¦   «         z   
                    dd¦  «        }t          j        ||fd¬	¦  «        }|                     ¦   «         | j        z  }|                     ¦   «         | j        z  }	d d d ¦  «         n# 1 swxY w Y   |                     |j        ¬
¦  «        |	                     |j        ¬
¦  «        fS )Nr   rM   r    ÚmpsÚcpuF)Údevice_typeÚenabledr:   rN   r7  )r$  ro   rÅ   rm   rd   r2  r  ÚtypeÚstrr   rÙ   r-   r³   r½   r/  r¾   rc   )
rH   r´   r  Úinv_freq_expandedÚposition_ids_expandedr?  ÚfreqsÚembr½   r¾   s
             r4   rR   z'GraniteMoeSharedRotaryEmbedding.forwardÑ  s·  € ð !œM¨$°°°°4¨-Ô8×>Ò>Ñ@Ô@×GÒGÈÔHZÐ[\ÔH]Ð_aÐcdÑeÔe×hÒhÐijÔiqÑrÔrÐØ ,¨Q¨Q¨Q°°a°a°a¨ZÔ 8× >Ò >Ñ @Ô @Ðå'1°!´(´-ÅÑ'EÔ'EÐkÈ!Ì(Ì-Ð[`ÒJ`ÐJ`�a”h”m�mÐfkˆÝ¨¸UÐCÑCÔCð 	5ð 	5Ø&×,Ò,Ñ.Ô.Ð1F×1LÒ1LÑ1NÔ1NÑN×YÒYÐZ[Ð]^Ñ_Ô_ˆEÝ”)˜U E˜N°Ð3Ñ3Ô3ˆCØ—'’'‘)”)˜dÔ4Ñ4ˆCØ—'’'‘)”)˜dÔ4Ñ4ˆCð		5ð 	5ð 	5ñ 	5ô 	5ð 	5ð 	5ð 	5ð 	5ð 	5ð 	5øøøð 	5ð 	5ð 	5ð 	5ð �vŠv˜AœGˆvÑ$Ô$ c§f¢f°1´7 fÑ&;Ô&;Ð;Ð;s   ÃBE&Å&E*Å-E*rs   rù   )r)   r*   r+   r-   rT   r/   r!   r>   Ústaticmethodr   r0   rl   ro   r.  r�   r   rR   rU   rV   s   @r4   r#  r#     sú   ø€ € € € € € ØŒlÐÐÑðVð VÐ5ð Vð Vð Vð Vð Vð Vð  à04Ø+/Ø"ð*ð *Ø&¨Ñ-ð*à˜Ô(ð*ð �t‘ð*ð 
ˆ~˜uÐ$Ô	%ð	*ð *ð *ñ „\ð*ð: €U„]�_„_Øð<ð <ñ Ôñ „_ð<ð <ð <ð <ð <r3   r#  c                   óâ   ‡ — e Zd Zdefˆ fd„Zeee	 	 	 	 	 	 ddej	        dz  dej
        dz  dej	        dz  dedz  dej        dz  d	edz  d
ee         defd„¦   «         ¦   «         ¦   «         Zˆ xZS )ÚGraniteMoeSharedModelr8   c                 óö  •‡— t          ¦   «                              ‰¦  «         ‰j        | _        ‰j        | _        t          j        ‰j        ‰j        | j        ¦  «        | _        t          j	        ˆfd„t          ‰j        ¦  «        D ¦   «         ¦  «        | _        t          ‰j        ‰j        ¬¦  «        | _        t!          ‰¬¦  «        | _        d| _        ‰j        | _        |                      ¦   «          d S )Nc                 ó0   •— g | ]}t          ‰|¦  «        ‘ŒS r2   )rû   )Ú.0râ   r8   s     €r4   ú
<listcomp>z2GraniteMoeSharedModel.__init__.<locals>.<listcomp>ê  s$   ø€ ÐnÐnÐnÀÕ)¨&°)Ñ<Ô<ÐnÐnÐnr3   rý   ©r8   F)r=   r>   Úpad_token_idÚpadding_idxÚ
vocab_sizer   Ú	Embeddingr?   Úembed_tokensÚ
ModuleListÚrangeÚnum_hidden_layersÚlayersrY   rÿ   Únormr#  Ú
rotary_embÚgradient_checkpointingÚembedding_multiplierÚ	post_initrG   s    `€r4   r>   zGraniteMoeSharedModel.__init__ã  sà   øø€ Ý‰Œ×Ò˜Ñ Ô Ð Ø!Ô.ˆÔØ Ô+ˆŒåœL¨Ô):¸FÔ<NÐPTÔP`ÑaÔaˆÔÝ”mØnÐnÐnÐnÍeÐTZÔTlÑNmÔNmÐnÑnÔnñ
ô 
ˆŒõ ,¨FÔ,>ÀFÔDWÐXÑXÔXˆŒ	Ý9ÀÐHÑHÔHˆŒØ&+ˆÔ#Ø$*Ô$?ˆÔ!ð 	�ŠÑÔÐÐÐr3   NÚ	input_idsrÐ   r  rð   Úinputs_embedsr  rÓ   rK   c           
      óZ  — |d u |d uz  rt          d¦  «        ‚|r|€t          | j        ¬¦  «        }|€|                      |¦  «        }|€V|�|                     ¦   «         nd}t          j        |j        d         |j        ¬¦  «        |z   }| 	                    d¦  «        }t          | j        ||||¬¦  «        }	|| j        z  }|}
|                      |
|¦  «        }| j        d | j        j        …         D ]} ||
f||	|||dœ|¤Ž}
Œ|                      |
¦  «        }
t!          |
|¬¦  «        S )	Nz:You must specify exactly one of input_ids or inputs_embedsrN  r   r    )r2  )r8   r^  rÐ   rð   r  )rï   rÐ   r  rð   r  )Úlast_hidden_staterð   )Ú
ValueErrorr   r8   rS  Úget_seq_lengthr-   r8  rm   r2  rº   r   r[  rY  rW  rV  rX  r   )rH   r]  rÐ   r  rð   r^  r  rÓ   Úpast_seen_tokensÚcausal_maskrJ   rï   Údecoder_layers                r4   rR   zGraniteMoeSharedModel.forwardô  s—  € ð ˜Ð -°tÐ";Ñ<ð 	[ÝÐYÑZÔZÐZàð 	?˜Ð0Ý*°$´+Ð>Ñ>Ô>ˆOàÐ Ø ×-Ò-¨iÑ8Ô8ˆMàÐØCRÐC^˜×=Ò=Ñ?Ô?Ð?ÐdeÐÝ œ<¨Ô(;¸AÔ(>À}ÔG[Ð\Ñ\Ô\Ð_oÑoˆLØ'×1Ò1°!Ñ4Ô4ˆLå(Ø”;Ø'Ø)Ø+Ø%ð
ñ 
ô 
ˆð &¨Ô(AÑAˆØ%ˆð #Ÿošo¨m¸\ÑJÔJÐà!œ[Ð)H¨4¬;Ô+HÐ)HÔIð 		ð 		ˆMØ)˜MØðà$7Ø*Ø)Ø /Ø#ðð ð ðð ˆMˆMð Ÿ	š	 -Ñ0Ô0ˆå%Ø+Ø+ð
ñ 
ô 
ð 	
r3   )NNNNNN)r)   r*   r+   r!   r>   r   r   r   r-   r.   rT   r
   r  r  r   r   r   rR   rU   rV   s   @r4   rI  rI  á  s  ø€ € € € € ðÐ5ð ð ð ð ð ð ð"  ØØð .2Ø.2Ø04Ø(,Ø26Ø!%ð5
ð 5
àÔ# dÑ*ð5
ð œ tÑ+ð5
ð Ô&¨Ñ-ð	5
ð
  ™ð5
ð Ô(¨4Ñ/ð5
ð ˜$‘;ð5
ð Ð+Ô,ð5
ð 
 ð5
ð 5
ð 5
ñ „^ñ „_ñ  Ôð5
ð 5
ð 5
ð 5
ð 5
r3   rI  r:   Úgate_logitsru   c                 óÆ  ‡— | �t          | t          ¦  «        sdS t          | t          ¦  «        r/| d         j        Št          j        ˆfd„| D ¦   «         d¬¦  «        }t          j        j                             |d¬¦  «        }t          j        ||d¬¦  «        \  }}t          j        j         	                    ||¦  «        }|€@t          j
        |                     ¦   «         d¬¦  «        }	t          j
        |d¬¦  «        }
�n.|j        \  }}|j        d         ||z  z  }|ddd…dd…ddf                              |||||f¦  «                             d||¦  «                             ‰¦  «        }t          j        |                     ¦   «         |z  d¬¦  «        t          j        |d¬¦  «        z  }	|ddd…dd…df                              ||||f¦  «                             d|¦  «                             ‰¦  «        }t          j        ||z  d¬¦  «        t          j        |d¬¦  «        z  }
t          j        |	|
                     d¦  «        z  ¦  «        }||z  S )aÄ  
    Computes auxiliary load balancing loss as in Switch Transformer - implemented in Pytorch.

    See Switch Transformer (https://huggingface.co/papers/2101.03961) for more details. This function implements the loss
    function presented in equations (4) - (6) of the paper. It aims at penalizing cases where the routing between
    experts is too unbalanced.

    Args:
        gate_logits:
            Logits from the `gate`, should be a tuple of model.config.num_hidden_layers tensors of
            shape [batch_size X sequence_length, num_experts].
        num_experts:
            Number of experts
        top_k:
            The number of experts to route per-token, can be also interpreted as the `top-k` routing
            parameter.
        attention_mask (`torch.Tensor`, *optional*):
            The attention_mask used in forward function
            shape [batch_size X sequence_length] if not None.

    Returns:
        The auxiliary loss.
    Nr   c                 ó:   •— g | ]}|                      ‰¦  «        ‘ŒS r2   )rd   )rL  Ú
layer_gateÚcompute_devices     €r4   rM  z,load_balancing_loss_func.<locals>.<listcomp>Q  s&   ø€ Ð-jÐ-jÐ-jÐPZ¨j¯mªm¸NÑ.KÔ.KÐ-jÐ-jÐ-jr3   rN   rM   )r  rl   r2  r-   r³   r   r‘   r}   r|   r’   rg   ro   rm   rÅ   r«   rd   r•   rº   )rf  ru   rw   rÐ   Úconcatenated_gate_logitsÚrouting_weightsr°   Úselected_expertsrš   Útokens_per_expertÚrouter_prob_per_expertÚ
batch_sizeÚsequence_lengthrV  Úexpert_attention_maskÚ router_per_expert_attention_maskÚoverall_lossrj  s                    @r4   Úload_balancing_loss_funcru  /  s�  ø€ ð: Ð¥*¨[½%Ñ"@Ô"@ÐØˆqå�+�uÑ%Ô%ð sØ$ QœÔ.ˆÝ#(¤9Ð-jÐ-jÐ-jÐ-jÐ^iÐ-jÑ-jÔ-jÐpqÐ#rÑ#rÔ#rÐ å”hÔ)×1Ò1Ð2JÐPRÐ1ÑSÔS€Oåœ* _°eÀÐDÑDÔDÑ€AÐå”(Ô%×-Ò-Ð.>ÀÑLÔL€KàÐå!œJ {×'8Ò'8Ñ':Ô':ÀÐBÑBÔBÐõ "'¤¨OÀÐ!CÑ!CÔ!CÐÑà&4Ô&:Ñ#ˆ
�OØ4Ô:¸1Ô=À*ÈÑB^Ñ_Ðð ˜4    A A A t¨TÐ1Ô2ßŠVÐ&¨
°OÀUÈKÐXÑYÔYßŠW�R˜ Ñ,Ô,ßŠR�ÑÔð	 	õ "œI k×&7Ò&7Ñ&9Ô&9Ð<QÑ&QÐWXÐYÑYÔYÕ\aÔ\eØ! qð]
ñ ]
ô ]
ñ 
Ðð ˜4    A A A tÐ+Ô,ßŠVÐ&¨
°OÀ[ÐQÑRÔRßŠW�R˜Ñ%Ô%ßŠR�ÑÔð	 	)õ "'¤¨?Ð=]Ñ+]ÐcdÐ!eÑ!eÔ!eÕhmÔhqØ,°!ði
ñ i
ô i
ñ "
Ðõ ”9Ð.Ð1G×1QÒ1QÐRSÑ1TÔ1TÑTÑUÔU€LØ˜+Ñ%Ð%r3   c                   ó  ‡ — e Zd ZddiZddiZddgdgfiZdefˆ fd„Zee		 	 	 	 	 	 	 	 dde
j        d	z  de
j        d	z  de
j        d	z  ded	z  de
j        d	z  de
j        d	z  ded	z  dee
j        z  deez  fd„¦   «         ¦   «         Zˆ xZS )ÚGraniteMoeSharedForCausalLMzlm_head.weightzmodel.embed_tokens.weightÚlm_headÚcolwise_gather_outputrJ   Úlogitsr8   c                 ó^  •— t          ¦   «                              |¦  «         t          |¦  «        | _        |j        | _        t          j        |j        |j        d¬¦  «        | _        |j	        | _	        |j
        | _        |j        | _        |j        | _        |                      ¦   «          d S )NFr;   )r=   r>   rI  r  rQ  r   rD   r?   rx  Úrouter_aux_loss_coefrt   ru   rv   Úlogits_scalingr\  rG   s     €r4   r>   z$GraniteMoeSharedForCausalLM.__init__‡  s–   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ý*¨6Ñ2Ô2ˆŒ
Ø Ô+ˆŒÝ”y Ô!3°VÔ5FÈUÐSÑSÔSˆŒØ$*Ô$?ˆÔ!Ø!Ô3ˆÔØ#)Ô#=ˆÔ Ø$Ô3ˆÔð 	�ŠÑÔÐÐÐr3   Nr   r]  rÐ   r  rð   r^  ÚlabelsÚoutput_router_logitsÚlogits_to_keeprK   c	           	      ó2  — |�|n| j         j        } | j        d|||||dœ|	¤Ž}
|
j        }t	          |t
          ¦  «        rt          | d¦  «        n|}|                      |dd…|dd…f         ¦  «        }|| j         j        z  }d}|� | j	        ||fd| j         j
        i|	¤Ž}d}|rHt          |
j        | j        | j        |¦  «        }|�%|| j        |                     |j        ¦  «        z  z  }t%          ||||
j        |
j        |
j        |
j        ¬¦  «        S )ax  
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
            config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
            (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.

        Example:

        ```python
        >>> from transformers import AutoTokenizer, GraniteMoeSharedForCausalLM

        >>> model = GraniteMoeSharedForCausalLM.from_pretrained("ibm/PowerMoE-3b")
        >>> tokenizer = AutoTokenizer.from_pretrained("ibm/PowerMoE-3b")

        >>> prompt = "Hey, are you conscious? Can you talk to me?"
        >>> inputs = tokenizer(prompt, return_tensors="pt")

        >>> # Generate
        >>> generate_ids = model.generate(inputs.input_ids, max_length=30)
        >>> tokenizer.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
        "Hey, are you conscious? Can you talk to me?\nI'm not conscious, but I can talk to you."
        ```N)r]  rÐ   r  rð   r^  rQ  )ÚlossÚaux_lossrz  rð   rJ   r  r   r2   )r8   r  r  r`  r  r0   Úslicerx  r}  Úloss_functionrQ  ru  r   ru   rv   r|  rd   r2  r   rð   rJ   r  )rH   r]  rÐ   r  rð   r^  r~  r  r€  rÓ   ÚoutputsrJ   Úslice_indicesrz  r‚  rƒ  s                   r4   rR   z#GraniteMoeSharedForCausalLM.forward”  s‘  € ðJ %9Ð$DÐ Ð È$Ì+ÔJjð 	ð �$”*ð 
ØØ)Ø%Ø+Ø'ð
ð 
ð ð
ð 
ˆð  Ô1ˆÝ8BÀ>ÕSVÑ8WÔ8WÐk�˜~˜o¨tÑ4Ô4Ð4Ð]kˆØ—’˜m¨A¨A¨A¨}¸a¸a¸aÐ,?Ô@ÑAÔAˆØ˜$œ+Ô4Ñ4ˆàˆØÐà%�4Ô%ØØðð ð  œ;Ô1ðð ð	ð ˆDð ˆØð 	MÝ/ØÔ%ØÔ ØÔ(Øñ	ô ˆHð Ð!Ø˜Ô1°H·K²KÀÄÑ4LÔ4LÑLÑL�Ý(ØØØØ#Ô3Ø!Ô/ØÔ)Ø!Ô/ð
ñ 
ô 
ð 	
r3   )NNNNNNNr   )r)   r*   r+   Ú_tied_weights_keysÚ_tp_planÚ_pp_planr!   r>   r   r   r-   r.   rT   r
   r  r  r0   rl   r   rR   rU   rV   s   @r4   rw  rw  �  s`  ø€ € € € € à*Ð,GÐHÐØÐ2Ð3€HØ˜_Ð-°¨zÐ:Ð;€HðÐ5ð ð ð ð ð ð ð Øð .2Ø.2Ø04Ø(,Ø26Ø*.Ø,0Ø-.ðQ
ð Q
àÔ# dÑ*ðQ
ð œ tÑ+ðQ
ð Ô&¨Ñ-ð	Q
ð
  ™ðQ
ð Ô(¨4Ñ/ðQ
ð Ô  4Ñ'ðQ
ð # T™kðQ
ð ˜eœlÑ*ðQ
ð 
Ð*Ñ	*ðQ
ð Q
ð Q
ñ Ôñ „^ðQ
ð Q
ð Q
ð Q
ð Q
r3   rw  )rw  rI  r  )r    )rË   )Nr:   N)JÚcollections.abcr   Útypingr   r   r-   Útorch.nn.functionalr   r‘   rz   Ú r   r  Úactivationsr	   Úcache_utilsr
   r   Ú
generationr   Úintegrationsr   r   r   r   Úmasking_utilsr   Úmodeling_layersr   Úmodeling_outputsr   r   Úmodeling_rope_utilsr   r   Úmodeling_utilsr   r   Úprocessing_utilsr   Úutilsr   r   Úutils.genericr   r   r   Úutils.output_capturingr   Úconfiguration_granitemoesharedr!   r#   ÚModuler7   rY   rq   r„   r¤   r·   rÂ   rT   r0   rÊ   ro   rß   rá   rû   r  r#  rI  rl   ru  rw  Ú__all__r2   r3   r4   ú<module>rŸ     så  ðð* %Ð $Ð $Ð $Ð $Ð $Ø &Ð &Ð &Ð &Ð &Ð &Ð &Ð &à €€€Ø Ð Ð Ð Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð à &Ð &Ð &Ð &Ð &Ð &Ø !Ð !Ð !Ð !Ð !Ð !Ø .Ð .Ð .Ð .Ð .Ð .Ð .Ð .Ø )Ð )Ð )Ð )Ð )Ð )ðð ð ð ð ð ð ð ð ð ð ð ð 0Ð /Ð /Ð /Ð /Ð /Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø QÐ QÐ QÐ QÐ QÐ QÐ QÐ QØ KÐ KÐ KÐ KÐ KÐ KÐ KÐ KØ FÐ FÐ FÐ FÐ FÐ FÐ FÐ FØ &Ð &Ð &Ð &Ð &Ð &Ø 7Ð 7Ð 7Ð 7Ð 7Ð 7Ð 7Ð 7Ø YÐ YÐ YÐ YÐ YÐ YÐ YÐ YÐ YÐ YØ 5Ð 5Ð 5Ð 5Ð 5Ð 5Ø BÐ BÐ BÐ BÐ BÐ Bðð ð ð ð  )°5ð ñ ô ð ð0ð ð ð ð ˜"œ)ñ ô ð ð4 Ð˜YÑ'Ô'ðJð Jð Jð Jð J˜bœiñ Jô Jñ (Ô'ðJð(9ð 9ð 9ð 9ð 9 ¤ñ 9ô 9ð 9ð, ð$#ð $#ð $#ð $#ð $#˜bœiñ $#ô $#ñ Ôð$#ðN?ð ?ð ?ð ?ð ?˜"œ)ñ ?ô ?ð ?ð"(ð (ð (ð ÐÐ*Ñ+Ô+ðð ð ñ ,Ô+ðð2	U˜Uœ\ð 	U°#ð 	U¸%¼,ð 	Uð 	Uð 	Uð 	Uð& ð%ð %ØŒIð%àŒ<ð%ð 
Œð%ð Œ<ð	%ð
 ”L 4Ñ'ð%ð ð%ð ð%ð Ð'Ô(ð%ð %ð %ð %ð2 ÐÐ)Ñ*Ô*ð@)ð @)ð @)ð @)ð @) ¤	ñ @)ô @)ñ +Ô*ð@)ðF0ð 0ð 0ð 0ð 0Ð#=ñ 0ô 0ð 0ðf ðUð Uð Uð Uð U oñ Uô Uñ „ðUð4><ð ><ð ><ð ><ð >< b¤iñ ><ô ><ð ><ðB ðJ
ð J
ð J
ð J
ð J
Ð;ñ J
ô J
ñ „ðJ
ð^ #Ø
Ø*.ð	O&ð O&Ø”  e¤lÔ 3Ñ3°dÑ:ðO&à�t‘ðO&ð ”L 4Ñ'ð	O&ð
 „\�CÑðO&ð O&ð O&ð O&ðd ðe
ð e
ð e
ð e
ð e
Ð"AÀ?ñ e
ô e
ñ „ðe
ðP fÐ
eÐ
e€€€r3   