Ë
    T^(hâî  ã                   ó¬  — d dl Z d dlmZ d dlmZmZ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mZ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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$«       rd dl*m+Z+ d dl,m-Z- d dl.m/Z/ ne0Z- e%jb                  e2«      Z3dZ4dZ5 G d„ dejl                  jn                  «      Z8	 	 dUdeejr                     dee:   fd„Z; G d„ de-«      Z< G d„ de
jz                  «      Z> G d„ de
jz                  «      Z? G d „ d!e
jz                  «      Z@d"„ ZAdVd#„ZB	 dWd$d%d&ejr                  d'ejr                  d(ejr                  d)eej†                     d*ee:e:f   d+e:d,e:d-eeD   d.eeejr                  ejr                  f   eejr                     f   fd/„ZEejŒ                  fd$d%d&ejr                  d0e<dejr                  de:d*ee:e:f   d+e:d,e:d1ejŽ                  d.eejr                     fd2„ZHd$d%d&ejr                  d'ejr                  d(ejr                  d)eej†                     d*ee:e:f   d+e:d,e:d.eejr                     fd3„ZIeHeEeId4œZJ G d5„ d%e
jz                  «      ZK G d6„ d7e
jz                  «      ZLd8ZM e"d9eM«       G d:„ d;e«      «       ZN	 	 dUd<ejr                  d'ejr                  d)eejr                     d=eejr                     d.eejr                  ejr                  ejr                  e:eejr                     eejr                     f   f
d>„ZOd<ejr                  d?ejr                  d@e:dAe:d.ejr                  f
dB„ZPdCZQ e"d9eM«       G dD„ dEeN«      «       ZR G dF„ dGe
jz                  «      ZS e"dHeM«       G dI„ dJeN«      «       ZT e"dKeM«       G dL„ dMeN«      «       ZU e"dNeM«       G dO„ dPeN«      «       ZV e"dQeM«       G dR„ dSeN«      «       ZWg dT¢ZXy)Xé    N)Únullcontext)ÚDictÚOptionalÚTupleÚUnion)Únn)ÚBCEWithLogitsLossÚCrossEntropyLossÚMSELossé   )ÚACT2FN)Ú_prepare_4d_attention_mask)ÚBaseModelOutputÚMaskedLMOutputÚQuestionAnsweringModelOutputÚSequenceClassifierOutputÚTokenClassifierOutput)ÚROPE_INIT_FUNCTIONSÚdynamic_rope_update)ÚPreTrainedModel)Úadd_code_sample_docstringsÚadd_start_docstringsÚ%add_start_docstrings_to_model_forwardÚis_flash_attn_2_availableÚlogging)Úis_triton_availableé   )ÚModernBertConfig)Ú flash_attn_varlen_qkvpacked_func)ÚRotaryEmbedding)Úapply_rotaryzanswerdotai/ModernBERT-baser   c                   ó\   — e Zd Ze	 	 ddeej                     dee   fd„«       Zed„ «       Z	y)ÚApplyRotaryEmbUnpadNÚ
cu_seqlensÚ
max_seqlenc           
      óÚ   — |j                  «       }|j                  \  }}}}	|d d …d d…f   j                  |d|	«      }
t        |
||d||dd¬«       | j	                  |||«       || _        |S )Né   éÿÿÿÿr   FT)Úseqlen_offsetsr$   r%   ÚinterleavedÚinplace)Ú
contiguousÚshapeÚviewr!   Úsave_for_backwardr%   )ÚctxÚqkvÚcosÚsinr$   r%   Ú	total_nnzÚ_threeÚ_nheadsÚheaddimÚqks              úp/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/modernbert/modeling_modernbert.pyÚforwardzApplyRotaryEmbUnpad.forwardC   s‚   € ð �n‰nÓˆØ.1¯i©iÑ+ˆ	�6˜7 Gð ’�B�Q�B�‰Z�_‰_˜Y¨¨GÓ4ˆÜØØØØØ!Ø!ØØõ		
ð 	×Ñ˜c 3¨
Ô3Ø#ˆŒØˆ
ó    c                 óê   — | j                   \  }}}|j                  «       }|j                  \  }}}}|d d …d d…f   j                  |d|«      }	t	        |	||d|| j
                  ddd¬«	       |d d d d d d fS )Nr'   r(   r   FT)r)   r$   r%   r*   r+   Ú	conjugate)Úsaved_tensorsr,   r-   r.   r!   r%   )
r0   Údor2   r3   r$   r4   r5   r6   r7   Údqks
             r9   ÚbackwardzApplyRotaryEmbUnpad.backwardb   s�   € à"×0Ñ0ÑˆˆS�*Ø�]‰]‹_ˆØ.0¯h©hÑ+ˆ	�6˜7 Gð ’�B�Q�B�‰i�n‰n˜Y¨¨GÓ4ˆÜØØØØØ!Ø—~‘~ØØØõ
	
ð �4˜˜t T¨4°Ð5Ð5r;   ©NN)
Ú__name__Ú
__module__Ú__qualname__Ústaticmethodr   ÚtorchÚTensorÚintr:   rA   © r;   r9   r#   r#   B   sQ   „ Øð .2Ø$(ñð
 ˜UŸ\™\Ñ*ðð ˜S‘Mòó ðð< ñ6ó ñ6r;   r#   r$   r%   c                 ó4   — t         j                  | ||||«      S )aÅ  
    Arguments:
        qkv: (total_nnz, 3, nheads, headdim) - input tensor for packed QKV.
        cos, sin: (seqlen_rotary, rotary_dim / 2)
        interleaved: if True, rotate pairs of even and odd dimensions (GPT-J style) instead
            of 1st half and 2nd half (GPT-NeoX style).
        inplace: if True, apply rotary embedding in-place.
        seqlen_offsets: (batch_size,) or int. Each sequence in x is shifted by this amount.
            Most commonly used in inference when we have KV cache.
        cu_seqlens: (batch + 1,) or None
        max_seqlen: int
    Return:
        out: (total_nnz, dim)
    rotary_dim must be <= headdim
    Apply rotary embedding to the first rotary_dim of x.
    )r#   Úapply)r1   r2   r3   r$   r%   s        r9   Úapply_rotary_unpaddedrM   y   s   € ô. ×$Ñ$ S¨#¨s°JÀ
ÓKÐKr;   c                   ó"  ‡ — e Zd ZdZ	 	 	 	 ddededee   deej                     deej                     f
ˆ fd„Z
	 ddej                  d	ej                  dee   d
eej                  eej                  ej                  f   f   fd„Zd
efd„Zˆ xZS )Ú!ModernBertUnpaddedRotaryEmbeddingzP
    The rotary position embeddings applied directly to unpadded sequences.
    ÚdimÚbaser%   ÚdeviceÚdtypec                 óv   •— t         ‰| �  ||d|d¬«       || _        |�|�|�| j                  |||¬«       yyyy)a  
        max_seqlen: if max_seqlen, device, and dtype are provided, we precompute the cos_sin_cache
            up to max_seqlen. If the max_seqlen, device, or dtype during training/inference differ,
            the cos_sin_cache wll be recomputed during the forward pass.
        TF)rP   rQ   Úpos_idx_in_fp32rR   r*   N©rR   rS   )ÚsuperÚ__init__r%   Ú_update_cos_sin_cache)ÚselfrP   rQ   r%   rR   rS   Ú	__class__s         €r9   rX   z*ModernBertUnpaddedRotaryEmbedding.__init__˜   sV   ø€ ô 	‰Ñ˜S t¸TÈ&Ð^cÐÔdØ$ˆŒàÐ! fÐ&8¸UÐ=NØ×&Ñ& z¸&ÈÐ&ÕNð >OÐ&8Ð!r;   r1   r$   Úreturnc                 ó¢   — |�(| j                  ||j                  |j                  ¬«       t        || j                  | j
                  ||¬«      }|S )zØ
        Apply rotary embedding *inplace* to qkv.
        qkv: (total_nnz, 3, nheads, headdim)
        cu_seqlens: (batch + 1,) cumulative sequence lengths
        max_seqlen: int max seq length in the batch
        rV   ©r$   r%   )rY   rR   rS   rM   Ú_cos_cachedÚ_sin_cached)rZ   r1   r$   r%   s       r9   r:   z)ModernBertUnpaddedRotaryEmbedding.forward«   sS   € ð Ð!Ø×&Ñ& z¸#¿*¹*ÈCÏIÉIÐ&ÔVä#ØØ×ÑØ×ÑØ!Ø!ô
ˆð ˆ
r;   c                 óT   — d| j                   › d| j                  › d| j                  › �S )Nzdim=z, base=z, scale_base=)rP   rQ   Ú
scale_base©rZ   s    r9   Ú
extra_reprz,ModernBertUnpaddedRotaryEmbedding.extra_reprÄ   s(   € Ø�d—h‘h�Z˜w t§y¡y k°¸t¿¹Ð>OÐPÐPr;   )g     ˆÃ@NNN©N)rC   rD   rE   Ú__doc__rI   Úfloatr   rG   rR   rS   rX   rH   r   r   r:   Ústrrd   Ú__classcell__©r[   s   @r9   rO   rO   “   sÑ   ø„ ñð Ø$(Ø)-Ø'+ñOàðOð ðOð ˜S‘Mð	Oð
 ˜Ÿ™Ñ&ðOð ˜Ÿ™Ñ$õOð. %)ñ	à�\‰\ðð —L‘Lðð ˜S‘Mð	ð
 
ˆu�|‰|˜U 5§<¡<°·±Ð#=Ñ>Ð>Ñ	?óð2Q˜C÷ Qr;   rO   c                   óì   ‡ — e Zd ZdZdefˆ fd„Z ej                  d¬«      dej                  dej                  fd„«       Z
	 ddeej                     d	eej                     dej                  fd
„Zˆ xZS )ÚModernBertEmbeddingszV
    Same as BertEmbeddings with a tiny tweak for positional embeddings indexing.
    Úconfigc                 ód  •— t         ‰| �  «        || _        t        j                  |j
                  |j                  |j                  ¬«      | _        t        j                  |j                  |j                  |j                  ¬«      | _        t        j                  |j                  «      | _        y )N)Úpadding_idx©ÚepsÚbias)rW   rX   rm   r   Ú	EmbeddingÚ
vocab_sizeÚhidden_sizeÚpad_token_idÚtok_embeddingsÚ	LayerNormÚnorm_epsÚ	norm_biasÚnormÚDropoutÚembedding_dropoutÚdrop©rZ   rm   r[   s     €r9   rX   zModernBertEmbeddings.__init__Í   sw   ø€ Ü‰ÑÔØˆŒÜ Ÿl™l¨6×+<Ñ+<¸f×>PÑ>PÐ^d×^qÑ^qÔrˆÔÜ—L‘L ×!3Ñ!3¸¿¹Èv×O_ÑO_Ô`ˆŒ	Ü—J‘J˜v×7Ñ7Ó8ˆ�	r;   T©ÚdynamicÚ	input_idsr\   c                 ó`   — | j                  | j                  | j                  |«      «      «      S re   )r~   r{   rw   )rZ   r‚   s     r9   Úcompiled_embeddingsz(ModernBertEmbeddings.compiled_embeddingsÔ   s%   € à�y‰y˜Ÿ™ 4×#6Ñ#6°yÓ#AÓBÓCÐCr;   Úinputs_embedsc                 óú   — |�"| j                  | j                  |«      «      }|S | j                  j                  r| j	                  |«      n.| j                  | j                  | j                  |«      «      «      }|S re   )r~   r{   rm   Úreference_compiler„   rw   )rZ   r‚   r…   Úhidden_statess       r9   r:   zModernBertEmbeddings.forwardØ   su   € ð Ð$Ø ŸI™I d§i¡i°Ó&>Ó?ˆMð Ðð —;‘;×0Ò0ð ×(Ñ(¨Ô3à—Y‘Y˜tŸy™y¨×)<Ñ)<¸YÓ)GÓHÓIð ð
 Ðr;   rB   )rC   rD   rE   rf   r   rX   rG   ÚcompileÚ
LongTensorrH   r„   r   r:   ri   rj   s   @r9   rl   rl   È   s�   ø„ ñð9Ð/õ 9ð €U‡]�]˜4Ô ðD¨U×-=Ñ-=ð DÀ%Ç,Á,ò Dó !ðDð eiñØ! %×"2Ñ"2Ñ3ðØKSÐTY×T`ÑT`ÑKaðà	�‰÷r;   rl   c                   ó`   ‡ — e Zd ZdZdefˆ fd„Zdej                  dej                  fd„Zˆ xZ	S )ÚModernBertMLPa6  Applies the GLU at the end of each ModernBERT layer.

    Compared to the default BERT architecture, this block replaces :class:`~transformers.model.bert.modeling_bert.BertIntermediate`
    and :class:`~transformers.model.bert.modeling_bert.SelfOutput` with a single module that has similar functionality.
    rm   c                 ó¬  •— t         ‰| �  «        || _        t        j                  |j
                  t        |j                  «      dz  |j                  ¬«      | _	        t        |j                     | _        t        j                  |j                  «      | _        t        j                  |j                  |j
                  |j                  ¬«      | _        y )Nr'   ©rr   )rW   rX   rm   r   ÚLinearru   rI   Úintermediate_sizeÚmlp_biasÚWir   Úhidden_activationÚactr|   Úmlp_dropoutr~   ÚWor   s     €r9   rX   zModernBertMLP.__init__í   s“   ø€ Ü‰ÑÔØˆŒÜ—)‘)˜F×.Ñ.´°F×4LÑ4LÓ0MÐPQÑ0QÐX^×XgÑXgÔhˆŒÜ˜&×2Ñ2Ñ3ˆŒÜ—J‘J˜v×1Ñ1Ó2ˆŒ	Ü—)‘)˜F×4Ñ4°f×6HÑ6HÈvÏÉÔ_ˆ�r;   rˆ   r\   c                 ó°   — | j                  |«      j                  dd¬«      \  }}| j                  | j                  | j	                  |«      |z  «      «      S )Nr'   r(   ©rP   )r’   Úchunkr–   r~   r”   )rZ   rˆ   ÚinputÚgates       r9   r:   zModernBertMLP.forwardõ   sI   € Ø—g‘g˜mÓ,×2Ñ2°1¸"Ð2Ó=‰ˆˆtØ�w‰w�t—y‘y §¡¨%£°4Ñ!7Ó8Ó9Ð9r;   )
rC   rD   rE   rf   r   rX   rG   rH   r:   ri   rj   s   @r9   rŒ   rŒ   æ   s2   ø„ ñð`Ð/õ `ð: U§\¡\ð :°e·l±l÷ :r;   rŒ   c            
       ó„   ‡ — e Zd Zddedededeej                     fˆ fd„Z	 ej                  «       ed„ «       «       Zˆ xZS )ÚModernBertRotaryEmbeddingrm   rP   rQ   rR   c                 óÜ  •— t         ‰| �  «        t        |d«      rG|j                  �;|j                  j	                  d|j                  j	                  d«      «      | _        nd| _        |j                  | _        |j                  | _        || _	        t        | j
                     | _        | j                  d |||¬«      \  }| _        | j                  d|d¬«       | j                  | _        y )	NÚrope_scalingÚ	rope_typeÚtypeÚdefault)rP   rQ   Úinv_freqF)Ú
persistent)rW   rX   ÚhasattrrŸ   Úgetr    Úmax_position_embeddingsÚmax_seq_len_cachedÚoriginal_max_seq_lenrm   r   Úrope_init_fnÚattention_scalingÚregister_bufferr£   Úoriginal_inv_freq)rZ   rm   rP   rQ   rR   r£   r[   s         €r9   rX   z"ModernBertRotaryEmbedding.__init__û   sÍ   ø€ Ü‰ÑÔä�6˜>Ô*¨v×/BÑ/BÐ/NØ#×0Ñ0×4Ñ4°[À&×BUÑBU×BYÑBYÐZ`ÓBaÓbˆD�Nà&ˆDŒNØ"(×"@Ñ"@ˆÔØ$*×$BÑ$BˆÔ!àˆŒÜ/°·±Ñ?ˆÔØ+/×+<Ñ+<¸TÀ6ÈsÐY]Ð+<Ó+^Ñ(ˆ�$Ô(Ø×Ñ˜Z¨¸eÐÔDØ!%§¡ˆÕr;   c                 ób  — | j                   d d d …d f   j                  «       j                  |j                  d   dd«      j	                  |j
                  «      }|d d …d d d …f   j                  «       }t        |j
                  j                  t        «      r/|j
                  j                  dk7  r|j
                  j                  nd}t        j                  |d¬«      5  |j                  «       |j                  «       z  j                  dd«      }t        j                  ||fd¬	«      }|j                  «       | j                  z  }|j                  «       | j                  z  }	d d d «       j	                  |j                   ¬
«      	j	                  |j                   ¬
«      fS # 1 sw Y   ŒAxY w)Nr   r(   r   ÚmpsÚcpuF)Údevice_typeÚenabledr'   r˜   )rS   )r£   rg   Úexpandr-   ÚtorR   Ú
isinstancer¡   rh   rG   ÚautocastÚ	transposeÚcatr2   r«   r3   rS   )
rZ   ÚxÚposition_idsÚinv_freq_expandedÚposition_ids_expandedr±   ÚfreqsÚembr2   r3   s
             r9   r:   z!ModernBertRotaryEmbedding.forward  sV  € ð !ŸM™M¨$²°4¨-Ñ8×>Ñ>Ó@×GÑGÈ×HZÑHZÐ[\ÑH]Ð_aÐcdÓe×hÑhÐij×iqÑiqÓrÐØ ,ªQ°²a¨ZÑ 8× >Ñ >Ó @Ðä'1°!·(±(·-±-ÄÔ'EÈ!Ï(É(Ï-É-Ð[`ÒJ`�a—h‘h—m’mÐfkˆÜ�^‰^¨¸UÔCñ 	5Ø&×,Ñ,Ó.Ð1F×1LÑ1LÓ1NÑN×YÑYÐZ[Ð]^Ó_ˆEÜ—)‘)˜U E˜N°Ô3ˆCØ—'‘'“)˜d×4Ñ4Ñ4ˆCØ—'‘'“)˜d×4Ñ4Ñ4ˆC÷		5ð �v‰v˜AŸG™GˆvÓ$ c§f¡f°1·7±7 fÓ&;Ð;Ð;÷	5ð 	5ús   Ã BF%Æ%F.re   )rC   rD   rE   r   rI   rg   r   rG   rR   rX   Úno_gradr   r:   ri   rj   s   @r9   r�   r�   ú   sV   ø„ ñ/Ð/ð /°cð /Àð /ÐPXÐY^×YeÑYeÑPfõ /ð  €U‡]�]ƒ_Øñ<ó ó ô<r;   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..Nr(   r'   r˜   )r-   rG   r¸   )r¹   Úx1Úx2s      r9   Úrotate_halfrÃ     sZ   € à	
ˆ3Ð"�!—'‘'˜"‘+ Ñ"Ð"Ð"Ñ	#€BØ	
ˆ3�—‘˜‘˜qÑ Ñ"Ð"Ñ	#€BÜ�9‰9�r�c˜2�Y BÔ'Ð'r;   c                 óž   — |j                  |«      }|j                  |«      }| |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.
        position_ids (`torch.Tensor`, *optional*):
            Deprecated and unused.
        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Úkr2   r3   rº   Úunsqueeze_dimÚq_embedÚk_embeds           r9   Úapply_rotary_pos_embrË   "  sY   € ð( �-‰-˜Ó
&€CØ
�-‰-˜Ó
&€CØ�3‰wœ; q›>¨CÑ/Ñ0€GØ�3‰wœ; q›>¨CÑ/Ñ0€GØ�GÐÐr;   ÚmoduleÚModernBertAttentionr1   Úattention_maskÚsliding_window_maskrº   Úlocal_attentionÚbsrP   Úoutput_attentionsr\   c	                 óÆ  — | j                  ||¬«      \  }
}|j                  dd«      j                  d¬«      \  }}}t        |||
|«      \  }}| j                  dz  }t        j                  ||j                  dd«      «      |z  }|dk7  r|}||z   }t        j                  j                  |dt
        j                  ¬	«      j                  |j                  «      }t        j                  j                  || j                  | j                  ¬
«      }t        j                  ||«      }|j                  dd«      j!                  «       }|j#                  |d|«      }|r||fS |fS )N©rº   r   r   r'   r˜   ç      à¿©r(   r(   r(   ©rP   rS   )ÚpÚtraining)Ú
rotary_embr·   ÚunbindrË   Úhead_dimrG   Úmatmulr   Ú
functionalÚsoftmaxÚfloat32r´   rS   ÚdropoutÚattention_dropoutrÙ   r,   r.   )rÌ   r1   rÎ   rÏ   rº   rÐ   rÑ   rP   rÒ   Ú_kwargsr2   r3   ÚqueryÚkeyÚvalueÚscaleÚattn_weightsÚattn_outputs                     r9   Úeager_attention_forwardrê   =  sK  € ð × Ñ  °<Ð Ó@�H€CˆØŸ™ a¨Ó+×2Ñ2°qÐ2Ó9Ñ€Eˆ3�ä% e¨S°#°sÓ;�J€Eˆ3à�O‰O˜TÑ!€EÜ—<‘<  s§}¡}°Q¸Ó':Ó;¸eÑC€Là˜(Ò"Ø,ˆà .Ñ0€Lô —=‘=×(Ñ(¨¸2ÄUÇ]Á]Ð(ÓS×VÑVÐW\×WbÑWbÓc€LÜ—=‘=×(Ñ(¨¸×9QÑ9QÐ\b×\kÑ\kÐ(Ól€LÜ—,‘,˜|¨UÓ3€KØ×'Ñ'¨¨1Ó-×8Ñ8Ó:€KØ×"Ñ" 2 r¨3Ó/€KÙØ˜\Ð*Ð*Øˆ>Ðr;   rÚ   Útarget_dtypec	                 óÄ  —  ||||¬«      }|j                   t        j                  t        j                  fv}
|
rb|j                   }|j	                  |«      }t        |||| j                  r| j                  nd| j                  |¬«      }|j	                  |«      }n3t        |||| j                  r| j                  nd| j                  |¬«      }|j                  ||«      fS )Nr^   ç        )r$   r%   Ú	dropout_pÚdeterministicÚwindow_size)
rS   rG   Úfloat16Úbfloat16r´   r   rÙ   râ   Údeterministic_flash_attnr.   )rÌ   r1   rÚ   r$   r%   rÐ   rÑ   rP   rë   rã   Úconvert_dtypeÚ
orig_dtypeÚattns                r9   Úflash_attention_forwardr÷   b  sÏ   € ñ �S Z¸JÔ
G€Cà—I‘I¤e§m¡m´U·^±^Ð%DÐD€MÙð —Y‘Yˆ
Ø�f‰f�\Ó"ˆä/ØØ!Ø!Ø28·/²/�f×.Ò.ÀsØ ×9Ñ9Ø'ô
ˆð �w‰w�zÓ"‰ä/ØØ!Ø!Ø28·/²/�f×.Ò.ÀsØ ×9Ñ9Ø'ô
ˆð �I‰I�b˜#ÓÐ Ð r;   c                 óv  — | j                  ||¬«      \  }	}
|j                  dd«      j                  d¬«      \  }}}t        |||	|
«      \  }}|dk7  r|}t	        j
                  |||| j                  r| j                  nd|¬«      j                  dd«      j                  «       }|j                  |d	|«      }|fS )
NrÔ   r   r   r'   r˜   rÖ   rí   )rî   Ú	attn_maskr(   )
rÚ   r·   rÛ   rË   ÚFÚscaled_dot_product_attentionrÙ   râ   r,   r.   )rÌ   r1   rÎ   rÏ   rº   rÐ   rÑ   rP   rã   r2   r3   rä   rå   ræ   ré   s                  r9   Úsdpa_attention_forwardrü   �  sÇ   € ð × Ñ  °<Ð Ó@�H€CˆØŸ™ a¨Ó+×2Ñ2°qÐ2Ó9Ñ€Eˆ3�ä% e¨S°#°sÓ;�J€Eˆ3à˜(Ò"Ø,ˆô 	
×&Ñ&ØØØØ28·/²/�f×.Ò.ÀsØ$ô	
÷ 
‰�1�a‹ß	‰‹ð ð ×"Ñ" 2 r¨3Ó/€KØˆ>Ðr;   )Úflash_attention_2ÚeagerÚsdpac                   óz   ‡ — e Zd ZdZd	dedee   fˆ fd„Z	 d
dej                  dee
   dej                  fd„Zˆ xZS )rÍ   a‚  Performs multi-headed self attention on a batch of unpadded sequences.

    If Flash Attention 2 is installed, this module uses Flash Attention to improve throughput.
    If Flash Attention 2 is not installed, the implementation will use PyTorch's SDPA kernel,
    which requires padding and unpadding inputs, adding some overhead.

    See `forward` method for additional details.
    rm   Úlayer_idc                 óâ  •— t         ‰| �  «        || _        || _        |j                  |j
                  z  dk7  r&t        d|j                  › d|j
                  › d�«      ‚|j                  | _        |j                  | _        |j
                  | _	        |j                  |j
                  z  | _
        | j                  | j                  z  | _        t        j                  |j                  d| j                  z  |j                  ¬«      | _        ||j                   z  dk7  r$|j"                  dz  |j"                  dz  f| _        nd| _        |j$                  }|j&                  }| j"                  dk7  r$|j(                  �|j(                  }|j"                  }|j*                  d	k(  rt-        | j                  ||¬
«      | _        nt1        || j                  |¬«      | _        t        j                  |j                  |j                  |j                  ¬«      | _        |j                  dkD  rt        j4                  |j                  «      nt        j6                  «       | _        t;        «       | _        y )Nr   zThe hidden size (z6) is not a multiple of the number of attention heads (ú)r   rŽ   r'   rÖ   rý   )rP   r%   rQ   )rm   rP   rQ   rí   )rW   rX   rm   r  ru   Únum_attention_headsÚ
ValueErrorrâ   ró   Ú	num_headsrÜ   Úall_head_sizer   r�   Úattention_biasÚWqkvÚglobal_attn_every_n_layersrÐ   Úglobal_rope_thetar§   Úlocal_rope_thetaÚ_attn_implementationrO   rÚ   r�   r–   r|   ÚIdentityÚout_dropÚsetÚpruned_heads)rZ   rm   r  Ú
rope_thetar§   r[   s        €r9   rX   zModernBertAttention.__init__Á  s   ø€ Ü‰ÑÔØˆŒØ ˆŒà×Ñ × :Ñ :Ñ:¸aÒ?ÜØ# F×$6Ñ$6Ð#7Ð7mÐnt÷  oIñ  oIð  nJð  JKð  Lóð ð "(×!9Ñ!9ˆÔØ(.×(GÑ(GˆÔ%Ø×3Ñ3ˆŒØ×*Ñ*¨f×.HÑ.HÑHˆŒØ!Ÿ]™]¨T¯^©^Ñ;ˆÔÜ—I‘I˜f×0Ñ0°!°d×6HÑ6HÑ2HÈv×OdÑOdÔeˆŒ	à�f×7Ñ7Ñ7¸1Ò<Ø$*×$:Ñ$:¸aÑ$?À×AWÑAWÐ[\ÑA\Ð#]ˆDÕ à#+ˆDÔ à×-Ñ-ˆ
Ø"(×"@Ñ"@ÐØ×Ñ 8Ò+Ø×&Ñ&Ð2Ø#×4Ñ4�
Ø&,×&<Ñ&<Ð#à×&Ñ&Ð*=Ò=Ü?Ø—M‘MÐ.EÈJôˆD�Oô 8¸vÈ4Ï=É=Ð_iÔjˆDŒOä—)‘)˜F×.Ñ.°×0BÑ0BÈ×I^ÑI^Ô_ˆŒØ@F×@XÑ@XÐ[^Ò@^œŸ
™
 6×#;Ñ#;Ô<Ôdf×doÑdoÓdqˆŒÜ›EˆÕr;   rˆ   rÒ   r\   c           
      ó  — | j                  |«      }|j                  d   }| j                  j                  dk(  r)|j	                  dd| j
                  | j                  «      }n)|j	                  |dd| j
                  | j                  «      }t        | j                  j                     | f|| j                  | j                  || j                  |dœ|¤Ž}|d   }| j                  | j                  |«      «      }|f|dd  z   S )Nr   rý   r(   r   )r1   rÚ   rÐ   rÑ   rP   rÒ   r   )r	  r-   rm   r  r.   r  rÜ   ÚMODERNBERT_ATTENTION_FUNCTIONrÚ   rÐ   r  r  r–   )rZ   rˆ   rÒ   Úkwargsr1   rÑ   Úattn_outputss          r9   r:   zModernBertAttention.forwardé  só   € ð �i‰i˜Ó&ˆà× Ñ  Ñ#ˆØ�;‰;×+Ñ+Ð/BÒBØ—(‘(˜2˜q $§.¡.°$·-±-Ó@‰Cà—(‘(˜2˜r 1 d§n¡n°d·m±mÓDˆCä4°T·[±[×5UÑ5UÑVØð	
àØ—‘Ø ×0Ñ0ØØ×"Ñ"Ø/ñ	
ð ñ	
ˆð % Q™ˆØŸ™ d§g¡g¨mÓ&<Ó=ˆàÐ ,¨q¨rÐ"2Ñ2Ð2r;   re   ©F)rC   rD   rE   rf   r   r   rI   rX   rG   rH   Úboolr:   ri   rj   s   @r9   rÍ   rÍ   ·  sS   ø„ ññ&"Ð/ð &"¸8ÀC¹=õ &"ðV -2ñ3à—|‘|ð3ð $ D™>ð3ð
 
�‰÷3r;   c                   óf  ‡ — e Zd Zddedee   fˆ fd„Z ej                  d¬«      dej                  dej                  fd„«       Z
	 	 	 	 	 	 ddej                  d	eej                     d
eej                     deej                     deej                     dee   dee   dej                  fd„Zˆ xZS )ÚModernBertEncoderLayerrm   r  c                 óž  •— t         ‰| �  «        || _        |dk(  rt        j                  «       | _        n;t        j                  |j                  |j                  |j                  ¬«      | _        t        ||¬«      | _        t        j                  |j                  |j                  |j                  ¬«      | _        t        |«      | _        y )Nr   rp   )rm   r  )rW   rX   rm   r   r  Ú	attn_normrx   ru   ry   rz   rÍ   rö   Úmlp_normrŒ   Úmlp©rZ   rm   r  r[   s      €r9   rX   zModernBertEncoderLayer.__init__  s�   ø€ Ü‰ÑÔØˆŒØ�qŠ=ÜŸ[™[›]ˆD�NäŸ\™\¨&×*<Ñ*<À&Ç/Á/ÐX^×XhÑXhÔiˆDŒNÜ'¨vÀÔIˆŒ	ÜŸ™ V×%7Ñ%7¸V¿_¹_ÐSY×ScÑScÔdˆŒÜ  Ó(ˆ�r;   Tr€   rˆ   r\   c                 óB   — | j                  | j                  |«      «      S re   )r  r  ©rZ   rˆ   s     r9   Úcompiled_mlpz#ModernBertEncoderLayer.compiled_mlp  s   € à�x‰x˜Ÿ™ mÓ4Ó5Ð5r;   rÎ   rÏ   rº   r$   r%   rÒ   c           	      ó
  — | j                  | j                  |«      ||||||¬«      }||d   z   }| j                  j                  r| j	                  |«      n| j                  | j                  |«      «      }	||	z   }|f|dd  z   S )N©rÎ   rÏ   rº   r$   r%   rÒ   r   r   )rö   r  rm   r‡   r"  r  r  )
rZ   rˆ   rÎ   rÏ   rº   r$   r%   rÒ   r  Ú
mlp_outputs
             r9   r:   zModernBertEncoderLayer.forward  sŸ   € ð —y‘yØ�N‰N˜=Ó)Ø)Ø 3Ø%Ø!Ø!Ø/ð !ó 
ˆð &¨°Q©Ñ7ˆð �{‰{×,Ò,ð ×Ñ˜mÔ,à—‘˜$Ÿ-™-¨Ó6Ó7ð 	ð
 &¨
Ñ2ˆàÐ ,¨q¨rÐ"2Ñ2Ð2r;   re   )NNNNNF)rC   rD   rE   r   r   rI   rX   rG   r‰   rH   r"  rŠ   r  r:   ri   rj   s   @r9   r  r    sí   ø„ ñ	)Ð/ð 	)¸8ÀC¹=õ 	)ð €U‡]�]˜4Ô ð6¨%¯,©,ð 6¸5¿<¹<ò 6ó !ð6ð 26Ø6:Ø37Ø-1Ø$(Ø,1ñ3à—|‘|ð3ð ! §¡Ñ.ð3ð & e§l¡lÑ3ð	3ð
 ˜u×/Ñ/Ñ0ð3ð ˜UŸ\™\Ñ*ð3ð ˜S‘Mð3ð $ D™>ð3ð 
�‰÷3r;   r  aO  
    This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the
    library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads
    etc.)

    This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.
    Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage
    and behavior.

    Parameters:
        config ([`ModernBertConfig`]):
            Model configuration class with all the parameters of the model. Initializing with a config file does not
            load the weights associated with the model, only the configuration. Check out the
            [`~PreTrainedModel.from_pretrained`] method to load the model weights.
zXThe bare ModernBert Model outputting raw hidden-states without any specific head on top.c                   óÈ   ‡ — e Zd ZeZdZdZddgZdZdZ	dZ
dej                  fd„Ze	 	 	 	 dded	eej$                     d
eeeeeef   f      defˆ fd„«       Zd„ Zˆ fd„Zˆ xZS )ÚModernBertPreTrainedModelÚmodelTrl   r  FrÌ   c                 ó¤  ‡— | j                   j                  Š‰€dŠdt        j                  dt        fˆfd„}| j                   j
                  | j                   j
                  t        j                  d| j                   j                  z  «      z  | j                   j
                  | j                   j                  dz  dœ}t        |t        «      r ||j                  |d   «       y t        |t        «      r- ||j                  |d	   «        ||j                  |d
   «       y t        |t         «      r- ||j"                  |d	   «        ||j                  |d
   «       y t        |t$        «      r ||j&                  |d
   «       y t        |t(        «      r ||j*                  |d
   «       y t        |t,        t.        t0        f«      r ||j2                  |d   «       y t        |t        j4                  «      rW|j6                  j8                  j;                  d«       |j<                  �%|j<                  j8                  j?                  «        y y y )Nr   rÌ   Ústdc                 ó  •— t         j                  j                  | j                  d|‰ |z  ‰|z  ¬«       t	        | t         j
                  «      r7| j                  �*t         j                  j                  | j                  «       y y y )Nrí   )Úmeanr*  ÚaÚb)r   ÚinitÚtrunc_normal_Úweightrµ   r�   rr   Úzeros_)rÌ   r*  Úcutoff_factors     €r9   Úinit_weightz<ModernBertPreTrainedModel._init_weights.<locals>.init_weightX  sq   ø€ Ü�G‰G×!Ñ!Ø—‘ØØØ �. 3Ñ&Ø #Ñ%ð "ô ô ˜&¤"§)¡)Ô,Ø—;‘;Ð*Ü—G‘G—N‘N 6§;¡;Õ/ð +ð -r;   g       @rÕ   )ÚinÚoutÚ	embeddingÚ	final_outr7  r5  r6  r8  g      ð?) rm   Úinitializer_cutoff_factorr   ÚModulerg   Úinitializer_rangeÚmathÚsqrtÚnum_hidden_layersru   rµ   rl   rw   rŒ   r’   r–   rÍ   r	  ÚModernBertPredictionHeadÚdenseÚModernBertForMaskedLMÚdecoderÚ#ModernBertForSequenceClassificationÚ ModernBertForTokenClassificationÚModernBertForQuestionAnsweringÚ
classifierrx   r1  ÚdataÚfill_rr   Úzero_)rZ   rÌ   r4  Ústdsr3  s       @r9   Ú_init_weightsz'ModernBertPreTrainedModel._init_weightsS  sÊ  ø€ ØŸ™×=Ñ=ˆØÐ ØˆMð	0¤§	¡	ð 	0´õ 	0ð —+‘+×/Ñ/Ø—;‘;×0Ñ0´4·9±9¸SÀ4Ç;Á;×C`ÑC`Ñ=`Ó3aÑaØŸ™×6Ñ6ØŸ™×0Ñ0°$Ñ6ñ	
ˆô �fÔ2Ô3Ù˜×-Ñ-¨t°KÑ/@ÕAÜ˜¤Ô.Ù˜Ÿ	™	 4¨¡:Ô.Ù˜Ÿ	™	 4¨¡;Õ/Ü˜Ô 3Ô4Ù˜Ÿ™ T¨$¡ZÔ0Ù˜Ÿ	™	 4¨¡;Õ/Ü˜Ô 8Ô9Ù˜Ÿ™ d¨5¡kÕ2Ü˜Ô 5Ô6Ù˜Ÿ™¨¨U©Õ4ÜØÜ0Ô2RÔTrÐsô
ñ ˜×)Ñ)¨4°Ñ+<Õ=Ü˜¤§¡Ô-Ø�M‰M×Ñ×$Ñ$ SÔ)Ø�{‰{Ð&Ø—‘× Ñ ×&Ñ&Õ(ð 'ð .r;   Úuse_flash_attention_2Útorch_dtypeÚ
device_mapÚcheck_device_mapc                 óð   •— |j                   €,d|_         	 | j                  |t        j                  |d|¬«      S t        ‰| �  ||t        j                  ||¬«      S # t        t
        f$ r
 d |_         Y Œ:w xY w)Nrý   F)rM  rN  Úhard_check_onlyrO  )rL  rM  rN  rO  )Ú_attn_implementation_internalÚ_check_and_enable_flash_attn_2rG   rñ   r  ÚImportErrorrW   Ú_autoset_attn_implementation)Úclsrm   rL  rM  rN  rO  r[   s         €r9   rU  z6ModernBertPreTrainedModel._autoset_attn_implementation‚  s“   ø€ ð ×/Ñ/Ð7Ø3FˆFÔ0ð	<Ø×9Ñ9ØÜ %§¡Ø)Ø$)Ø%5ð :ó ð ô ‰wÑ3ØØ"7ÜŸ™Ø!Ø-ð 4ó 
ð 	
øô ¤Ð,ò <Ø7;�Ö4ð<ús   –#A ÁA5Á4A5c                 óª  — | j                   j                  du ry t        | d«      rTt        | j                  «      dkD  r<| j                   j                  rt
        j                  d«       d| j                   _        | j                  j                  dk(  r<| j                   j                  rt
        j                  d«       d| j                   _        | j                  j                  dk(  r<| j                   j                  rt
        j                  d«       d| j                   _        | j                   j                  €t        «       | j                   _        y y )	NFÚhf_device_mapr   zqIf `accelerate` split the model across devices, `torch.compile` will not work. Falling back to non-compiled mode.r¯   z|Compiling the model with `torch.compile` and using a `torch.mps` device is not supported. Falling back to non-compiled mode.r°   z|Compiling the model with `torch.compile` and using a `torch.cpu` device is not supported. Falling back to non-compiled mode.)
rm   r‡   r¥   ÚlenrX  ÚloggerÚwarning_oncerR   r¡   r   rc   s    r9   Ú_maybe_set_compilez,ModernBertPreTrainedModel._maybe_set_compile£  s  € Ø�;‰;×(Ñ(¨EÑ1Øä�4˜Ô)¬c°$×2DÑ2DÓ.EÈÒ.IØ�{‰{×,Ò,Ü×#Ñ#ð9ôð -2ˆD�K‰KÔ)à�;‰;×Ñ˜uÒ$Ø�{‰{×,Ò,Ü×#Ñ#ð9ôð -2ˆD�K‰KÔ)à�;‰;×Ñ˜uÒ$Ø�{‰{×,Ò,Ü×#Ñ#ð9ôð -2ˆD�K‰KÔ)à�;‰;×(Ñ(Ð0Ü,?Ó,AˆD�K‰KÕ)ð 1r;   c                 óÎ   •— t        ‰| �  |i |¤Ž}| j                  j                  dv r<| j                  j                  rt        j                  d«       d| j                  _        |S )N>   NTzcResizing token embeddings with `torch.compile` is not supported. Falling back to non-compiled mode.F)rW   Úresize_token_embeddingsrm   r‡   rZ  r[  )rZ   Úargsr  Úmodel_embedsr[   s       €r9   r^  z1ModernBertPreTrainedModel.resize_token_embeddingsÂ  s[   ø€ Ü‘wÑ6¸ÐGÀÑGˆà�;‰;×(Ñ(¨LÑ8Ø�{‰{×,Ò,Ü×#Ñ#Øyôð -2ˆD�K‰KÔ)àÐr;   )FNNT)rC   rD   rE   r   Úconfig_classÚbase_model_prefixÚsupports_gradient_checkpointingÚ_no_split_modulesÚ_supports_flash_attn_2Ú_supports_sdpaÚ_supports_flex_attnr   r:  rK  Úclassmethodr  r   rG   rS   r   rh   r   rI   rU  r\  r^  ri   rj   s   @r9   r'  r'  F  s¿   ø„ ð
 $€LØÐØ&*Ð#Ø/Ð1IÐJÐØ!ÐØ€NØÐð-) B§I¡Ió -)ð^ ð ',Ø-1Ø;?Ø!%ñ
ð  $ð
ð ˜eŸk™kÑ*ð	
ð
 ˜U 3¨¨S°#¨X©Ð#6Ñ7Ñ8ð
ð ô
ó ð
ò@B÷>
ð 
r;   r'  ÚinputsÚlabelsc                 ó¢  — |j                  dt        j                  ¬«      }t        j                  |j	                  «       d¬«      j	                  «       }t        |j                  «       j                  «       «      }t        j                  j                  j                  t        j                  |dt        j                  ¬«      d«      }| j                  «       dk(  r| j	                  «       |   }n*| j                  ^}	}
}|	|
z  } | j                  |g|¢­Ž |   }|�|j	                  «       |   nd}|�|j	                  «       |   nd}||||||fS )	aˆ  
    Remove padding from input sequences.

    Args:
        inputs: (batch, seqlen, ...) or (batch, seqlen)
        attention_mask: (batch, seqlen), bool / int, 1 means valid and 0 means not valid.
        position_ids: (batch, seqlen), int, position ids
        labels: (batch, seqlen), int, labels

    Returns:
        unpadded_inputs: (total_nnz, ...), where total_nnz = number of tokens selected in attention_mask.
        indices: (total_nnz)
        cu_seqlens: (batch + 1), the cumulative sequence lengths
        max_seqlen_in_batch: int
        unpadded_position_ids: (total_nnz) or None
        unpadded_labels: (total_nnz) or None
    r(   r×   F)Úas_tupler   )r   r   r'   N)ÚsumrG   Úint32ÚnonzeroÚflattenrI   ÚmaxÚitemr   rÞ   ÚpadÚcumsumrP   r-   r.   )ri  rÎ   rº   rj  Úseqlens_in_batchÚindicesÚmax_seqlen_in_batchr$   Úunpadded_inputsÚbatchÚseqlenÚrestr-   Úunpadded_position_idsÚunpadded_labelss                  r9   Ú_unpad_modernbert_inputr~  Ï  s,  € ð. &×)Ñ)¨b¼¿¹Ð)ÓDÐÜ�m‰m˜N×2Ñ2Ó4¸uÔE×MÑMÓO€GÜÐ.×2Ñ2Ó4×9Ñ9Ó;Ó<ÐÜ—‘×$Ñ$×(Ñ(¬¯©Ð6FÈAÔUZ×U`ÑU`Ô)aÐciÓj€Jà‡z�zƒ|�qÒØ Ÿ.™.Ó*¨7Ñ3‰à%Ÿ|™|Ðˆˆv˜Ø˜‘ˆØ%˜&Ÿ+™+ eÐ3¨dÒ3°GÑ<ˆà?KÐ?W˜L×0Ñ0Ó2°7Ò;Ð]aÐØ39Ð3E�f—n‘nÓ& wÒ/È4€Oà˜G ZÐ1DÐF[Ð]lÐlÐlr;   rv  ry  rz  c                 ól  — | j                  «       dk(  rHt        j                  ||z  | j                  | j                  ¬«      }| ||<   |j                  ||«      }|S | j                  ^}}t        j                  ||z  g|¢­| j                  | j                  dœŽ}| ||<    |j
                  ||g|¢­Ž }|S )aQ  
    Add padding to sequences.

    Args:
        inputs: (total_nnz, ...) or (total_nnz,), where total_nnz = number of tokens selected in attention_mask.
        indices: (total_nnz)
        batch: int, batch size
        seqlen: int, max sequence length

    Returns:
        padded_inputs: (batch, seqlen, ...) or (batch, seqlen)
    r   )rS   rR   )rP   rG   ÚzerosrS   rR   r.   r-   )ri  rv  ry  rz  ÚoutputÚpadded_inputsÚ_r{  s           r9   Ú_pad_modernbert_outputr„  ø  s¬   € ð$ ‡z�zƒ|�qÒÜ—‘˜U V™^°6·<±<ÈÏÉÔVˆØ ˆˆw‰ØŸ™ E¨6Ó2ˆð Ðð —<‘<ˆˆˆDÜ—‘˜U V™^Ð]¨dÑ]¸&¿,¹,ÈvÏ}É}Ò]ˆØ ˆˆw‰Ø#˜Ÿ™ E¨6Ð9°DÒ9ˆàÐr;   aø  
    Args:
        input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
            Indices of input sequence tokens in the vocabulary. With Flash Attention 2.0, padding will be ignored
            by default should you provide it.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.

            [What are input IDs?](../glossary#input-ids)
        attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
            Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:

            - 1 for tokens that are **not masked**,
            - 0 for tokens that are **masked**.

            [What are attention masks?](../glossary#attention-mask)

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.

            If you want to change padding behavior, you should read [`modeling_opt._prepare_decoder_attention_mask`]
            and modify to your needs. See diagram 1 in [the paper](https://arxiv.org/abs/1910.13461) for more
            information on the default strategy.

            - 1 indicates the head is **not masked**,
            - 0 indicates the head is **masked**.
        sliding_window_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
            Mask to avoid performing attention on padding or far-away tokens. In ModernBert, only every few layers
            perform global attention, while the rest perform local attention. This mask is used to avoid attending to
            far-away tokens in the local attention layers when not using Flash Attention.
        position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
            config.n_positions - 1]`.

            [What are position IDs?](../glossary#position-ids)
        inputs_embeds (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
            Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This
            is useful if you want more control over how to convert `input_ids` indices into associated vectors than the
            model's internal embedding lookup matrix.
        indices (`torch.Tensor` of shape `(total_unpadded_tokens,)`, *optional*):
            Indices of the non-padding tokens in the input sequence. Used for unpadding the output.
        cu_seqlens (`torch.Tensor` of shape `(batch + 1,)`, *optional*):
            Cumulative sequence lengths of the input sequences. Used to index the unpadded tensors.
        max_seqlen (`int`, *optional*):
            Maximum sequence length in the batch excluding padding tokens. Used to unpad input_ids and pad output tensors.
        batch_size (`int`, *optional*):
            Batch size of the input sequences. Used to pad the output tensors.
        seq_len (`int`, *optional*):
            Sequence length of the input sequences including padding tokens. Used to pad the output tensors.
        output_attentions (`bool`, *optional*):
            Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
            tensors for more detail.
        output_hidden_states (`bool`, *optional*):
            Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
            more detail.
        return_dict (`bool`, *optional*):
            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
c            "       óü  ‡ — e Zd Zdefˆ fd„Zd„ Zd„ Z ee«       e	e
ee¬«      	 	 	 	 	 	 	 	 	 	 	 	 	 ddeej                     deej                      deej                      d	eej                     d
eej                      deej                      deej                      dee   dee   dee   dee   dee   dee   deeej                   df   ef   fd„«       «       Zdej                   dedej                   fd„Zˆ xZS )ÚModernBertModelrm   c           	      óŠ  •— t         ‰| �  |«       || _        t        |«      | _        t        j                  t        |j                  «      D �cg c]  }t        ||«      ‘Œ c}«      | _
        t        j                  |j                  |j                  |j                  ¬«      | _        d| _        | j#                  «        y c c}w )Nrp   F)rW   rX   rm   rl   Ú
embeddingsr   Ú
ModuleListÚranger>  r  Úlayersrx   ru   ry   rz   Ú
final_normÚgradient_checkpointingÚ	post_initr  s      €r9   rX   zModernBertModel.__init__Y  s“   ø€ Ü‰Ñ˜Ô ØˆŒÜ.¨vÓ6ˆŒÜ—m‘mÜFKÈF×LdÑLdÓFeÖf¸(Ô# F¨HÕ5Òfó
ˆŒô Ÿ,™, v×'9Ñ'9¸v¿¹ÐU[×UeÑUeÔfˆŒØ&+ˆÔ#Ø�‰Õùò	 gs   ÁC c                 ó.   — | j                   j                  S re   ©rˆ  rw   rc   s    r9   Úget_input_embeddingsz$ModernBertModel.get_input_embeddingsd  s   € Ø�‰×-Ñ-Ð-r;   c                 ó&   — || j                   _        y re   r�  )rZ   ræ   s     r9   Úset_input_embeddingsz$ModernBertModel.set_input_embeddingsg  s   € Ø).ˆ�‰Õ&r;   ©Ú
checkpointÚoutput_typera  r‚   rÎ   rÏ   rº   r…   rv  r$   r%   Ú
batch_sizeÚseq_lenrÒ   Úoutput_hidden_statesÚreturn_dictr\   .c                 ój  ‡‡	‡
— |�|n| j                   j                  }|�|n| j                   j                  }|�|n| j                   j                  }|d u |d uz  rt	        d«      ‚|rdnd }|rdnd }| j                  «        |�| j                  ||«       ‰	€)‰
€'|�|j                  d d \  Š	Š
n|j                  d d \  Š	Š
|�|j                  n|j                  }|€(t        j                  ‰	‰
f|t        j                  ¬«      }d}| j                   j                  dk(  rM‰€‰|€‡|€…d}|€0t        j                  «       5  t        ||¬«      ^}Š}}}d d d «       nQt        ||¬«      ^}Š}}}n>|€&t        j                  ‰
|¬	«      j!                  d
«      }| j#                  ||¬«      \  }}| j%                  ||¬«      }| j&                  D ]t  }|r||fz   }| j(                  r/| j*                  r#| j-                  |j.                  |||||||«      }n ||||||||¬«      }|d
   }|sŒ]t1        |«      dkD  sŒl||d   fz   }Œv |r||fz   }| j3                  |«      }|r't5        |‰‰	‰
¬«      }|�t7        ˆ	ˆˆ
fd„|D «       «      }|st7        d„ |||fD «       «      S t9        |||¬«      S # 1 sw Y   �ŒxY w)Nz:You must specify exactly one of input_ids or inputs_embedsrJ   r'   rV   Frý   T)ri  rÎ   )rR   r   )rÒ   )r‚   r…   r$  r   ©ri  rv  ry  rz  c              3   ó<   •K  — | ]  }t        |‰‰‰¬ «      –— Œ y­w)rœ  N)r„  )Ú.0Úhsr—  rv  r˜  s     €€€r9   ú	<genexpr>z*ModernBertModel.forward.<locals>.<genexpr>Ù  s(   øè ø€ ò *àô +°"¸gÈZÐ`g×hÐhñ*ùs   ƒc              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wre   rJ   )rž  Úvs     r9   r   z*ModernBertModel.forward.<locals>.<genexpr>ß  s   è ø€ Òm˜qÐ_`Ñ_lœÑmùs   ‚Š)Úlast_hidden_staterˆ   Ú
attentions)rm   rÒ   r™  Úuse_return_dictr  r\  Ú%warn_if_padding_and_no_attention_maskr-   rR   rG   Úonesr  r  r¿   r~  ÚarangerÅ   Ú_update_attention_maskrˆ  r‹  r�  rÙ   Ú_gradient_checkpointing_funcÚ__call__rY  rŒ  r„  Útupler   )rZ   r‚   rÎ   rÏ   rº   r…   rv  r$   r%   r—  r˜  rÒ   r™  rš  Úall_hidden_statesÚall_self_attentionsrR   Úrepadrƒ  rˆ   Úencoder_layerÚlayer_outputss         `  ``           r9   r:   zModernBertModel.forwardj  sG  ú€ ð, 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà˜Ð -°tÐ";Ò<ÜÐYÓZÐZá"6™B¸DÐÙ$5™b¸4Ðà×ÑÔ!àÐ Ø×6Ñ6°yÀ.ÔQàÐ ' /ØÐ(Ø&3×&9Ñ&9¸"¸1Ð&=Ñ#�
™Gà&/§o¡o°b°qÐ&9Ñ#�
˜GØ%.Ð%:�×!Ò!À×@TÑ@TˆàÐ!Ü"ŸZ™Z¨°WÐ(=ÀfÔTY×T^ÑT^Ô_ˆNàˆØ�;‰;×+Ñ+Ð/BÒBØˆ :Ð#5¸*Ð:LØ�Ø Ð(ÜŸ™›ñ ÜI`Ø#,¸^ôJÐF˜	 7¨J¸
ÀQ÷ð ô
 JaØ,¸^ôJÐF�M 7¨J¸
ÁQð Ð#Ü$Ÿ|™|¨G¸FÔC×MÑMÈaÓP�à26×2MÑ2MØÐ2Cð 3Nó 3Ñ/ˆNÐ/ð Ÿ™°)È=˜ÓYˆà!Ÿ[™[ò 	PˆMÙ#Ø$5¸Ð8HÑ$HÐ!à×*Ò*¨t¯}ª}Ø $× AÑ AØ!×*Ñ*Ø!Ø"Ø'Ø ØØØ%ó	!‘ñ !.Ø!Ø#1Ø(;Ø!-Ø)Ø)Ø&7ô!�ð *¨!Ñ,ˆMÚ ¤S¨Ó%7¸!Ó%;Ø&9¸]È1Ñ=MÐ<OÑ&OÑ#ð7	Pñ:  Ø 1°]Ð4DÑ DÐàŸ™¨Ó6ˆáÜ2Ø$¨g¸ZÐPWôˆMð !Ð,Ü$)õ *à/ô*ó %Ð!ñ
 ÜÑm ]Ð4EÐGZÐ$[ÔmÓmÐmÜØ+Ø+Ø*ô
ð 	
÷Añ ús   Ä>J(Ê(J2c                 ó   — |r†| j                   j                  dk(  r't        j                  d«       d| j                   _        nF| j                   j                  dk7  r-t        j                  d| j                   j                  › d�«       t	        || j
                  «      }t        j                  |j                  d   «      j                  d«      }t        j                  ||j                  z
  «      }|| j                   j                  dz  k  j                  d«      j                  d«      j                  |j                  «      }|j                  |j!                  «       t        j"                  | j
                  «      j$                  «      }||fS )Nrÿ   z’Outputting attentions is only supported with the 'eager' attention implementation, not with "sdpa". Falling back to `attn_implementation="eager"`.rþ   zZOutputting attentions is only supported with the eager attention implementation, not with zT. Consider setting `attn_implementation="eager"`. Setting `output_attentions=False`.r'   r   )rm   r  rZ  r[  r   rS   rG   r¨  r-   rÅ   ÚabsÚTrÐ   r´   rR   Úmasked_fillÚlogical_notÚfinfoÚmin)rZ   rÎ   rÒ   Úglobal_attention_maskÚrowsÚdistanceÚwindow_maskrÏ   s           r9   r©  z&ModernBertModel._update_attention_maskæ  sS  € ÙØ�{‰{×/Ñ/°6Ò9Ü×#Ñ#ðVôð 4;�—‘Õ0Ø—‘×1Ñ1°WÒ<Ü×#Ñ#ð Ø $§¡× @Ñ @ÐAð B:ð:ôô !;¸>È4Ï:É:Ó VÐô �|‰|Ð1×7Ñ7¸Ñ:Ó;×EÑEÀaÓHˆä—9‘9˜T D§F¡F™]Ó+ˆð ˜Ÿ™×4Ñ4¸Ñ9Ñ9×DÑDÀQÓG×QÑQÐRSÓT×WÑWÐXf×XmÑXmÓnð 	ð 4×?Ñ?À×@WÑ@WÓ@YÔ[`×[fÑ[fÐgk×gqÑgqÓ[r×[vÑ[vÓwÐà$Ð&9Ð9Ð9r;   ©NNNNNNNNNNNNN)rC   rD   rE   r   rX   r‘  r“  r   ÚMODERNBERT_INPUTS_DOCSTRINGr   Ú_CHECKPOINT_FOR_DOCr   Ú_CONFIG_FOR_DOCr   rG   rŠ   rH   rI   r  r   r   r:   r©  ri   rj   s   @r9   r†  r†  T  s¨  ø„ ð
	Ð/õ 	ò.ò/ñ +Ð+FÓGÙØ&Ø#Ø$ôð 15Ø15Ø6:Ø37Ø04Ø*.Ø-1Ø$(Ø$(Ø!%Ø,0Ø/3Ø&*ñt
à˜E×,Ñ,Ñ-ðt
ð ! §¡Ñ.ðt
ð & e§l¡lÑ3ð	t
ð
 ˜u×/Ñ/Ñ0ðt
ð   §¡Ñ-ðt
ð ˜%Ÿ,™,Ñ'ðt
ð ˜UŸ\™\Ñ*ðt
ð ˜S‘Mðt
ð ˜S‘Mðt
ð ˜#‘ðt
ð $ D™>ðt
ð ' t™nðt
ð ˜d‘^ðt
ð 
ˆu�U—\‘\ 3Ð&Ñ'¨Ð8Ñ	9òt
óó Hðt
ðl:°U·\±\ð :ÐVZð :Ð_d×_kÑ_k÷ :r;   r†  c                   ó\   ‡ — e Zd Zdefˆ fd„Zdej                  dej                  fd„Zˆ xZS )r?  rm   c                 óJ  •— t         ‰| �  «        || _        t        j                  |j
                  |j
                  |j                  «      | _        t        |j                     | _
        t        j                  |j
                  |j                  |j                  ¬«      | _        y )Nrp   )rW   rX   rm   r   r�   ru   Úclassifier_biasr@  r   Úclassifier_activationr”   rx   ry   rz   r{   r   s     €r9   rX   z!ModernBertPredictionHead.__init__  sq   ø€ Ü‰ÑÔØˆŒÜ—Y‘Y˜v×1Ñ1°6×3EÑ3EÀv×G]ÑG]Ó^ˆŒ
Ü˜&×6Ñ6Ñ7ˆŒÜ—L‘L ×!3Ñ!3¸¿¹Èv×O_ÑO_Ô`ˆ�	r;   rˆ   r\   c                 ó`   — | j                  | j                  | j                  |«      «      «      S re   )r{   r”   r@  r!  s     r9   r:   z ModernBertPredictionHead.forward  s#   € Ø�y‰y˜Ÿ™ $§*¡*¨]Ó";Ó<Ó=Ð=r;   )	rC   rD   rE   r   rX   rG   rH   r:   ri   rj   s   @r9   r?  r?    s-   ø„ ðaÐ/õ að> U§\¡\ð >°e·l±l÷ >r;   r?  zZThe ModernBert Model with a decoder head on top that is used for masked language modeling.c            #       ó`  ‡ — e Zd ZdgZdefˆ fd„Zd„ Zdej                  fd„Z	 e
j                  d¬«      d	e
j                  d
e
j                  fd„«       Z ee«       eeee¬«      	 	 	 	 	 	 	 	 	 	 	 	 	 	 ddee
j*                     dee
j                     dee
j                     dee
j                     dee
j                     dee
j                     dee
j                     dee
j                     dee   dee   dee   dee   dee   dee   d
eee
j                     ef   fd„«       «       Zˆ xZS )rA  zdecoder.weightrm   c                 ót  •— t         ‰| �  |«       || _        t        |«      | _        t        |«      | _        t        j                  |j                  |j                  |j                  ¬«      | _        | j                  j                  | _        | j                  j                  | _        | j                  «        y )NrŽ   )rW   rX   rm   r†  r(  r?  Úheadr   r�   ru   rt   Údecoder_biasrB  Úsparse_predictionÚsparse_pred_ignore_indexrŽ  r   s     €r9   rX   zModernBertForMaskedLM.__init__  s…   ø€ Ü‰Ñ˜Ô ØˆŒÜ$ VÓ,ˆŒ
Ü,¨VÓ4ˆŒ	Ü—y‘y ×!3Ñ!3°V×5FÑ5FÈV×M`ÑM`ÔaˆŒà!%§¡×!>Ñ!>ˆÔØ(,¯©×(LÑ(LˆÔ%ð 	�‰Õr;   c                 ó   — | j                   S re   ©rB  rc   s    r9   Úget_output_embeddingsz+ModernBertForMaskedLM.get_output_embeddings&  s   € Ø�|‰|Ðr;   Únew_embeddingsc                 ó   — || _         y re   rÍ  )rZ   rÏ  s     r9   Úset_output_embeddingsz+ModernBertForMaskedLM.set_output_embeddings)  s	   € Ø%ˆ�r;   Tr€   r�  r\   c                 óB   — | j                  | j                  |«      «      S re   )rB  rÈ  )rZ   r�  s     r9   Úcompiled_headz#ModernBertForMaskedLM.compiled_head,  s   € à�|‰|˜DŸI™I fÓ-Ó.Ð.r;   r”  r‚   rÎ   rÏ   rº   r…   rj  rv  r$   r%   r—  r˜  rÒ   r™  rš  c                 óH  — |�|n| j                   j                  }| j                  «        | j                   j                  dk(  rÁ|€¿|€½|	€»|
€)|€'|�|j                  d d \  }
}n|j                  d d \  }
}|�|j
                  n|j
                  }|€(t        j                  |
|f|t        j                  ¬«      }|€4t        j                  «       5  t        ||||¬«      \  }}}}	}}d d d «       nt        ||||¬«      \  }}}}	}}| j                  ||||||||	|
||||¬«      }|d   }| j                  rK|�I|j                  d«      }|j                  |j                  d   d«      }|| j                  k7  }||   }||   }| j                   j                  r| j!                  |«      n| j#                  | j%                  |«      «      }d }|�(| j'                  ||| j                   j(                  ¬«      }| j                   j                  dk(  rN| j                   j*                  s|€
t-        «       nt        j                  «       5  t/        |||
|¬	«      }d d d «       |s|f}|�|f|z   S |S t1        |||j2                  |j4                  ¬
«      S # 1 sw Y   �Œ�xY w# 1 sw Y   ŒHxY w)Nrý   r'   rV   )ri  rÎ   rº   rj  ©r‚   rÎ   rÏ   rº   r…   rv  r$   r%   r—  r˜  rÒ   r™  rš  r   r(   )rt   rœ  ©ÚlossÚlogitsrˆ   r¤  )rm   r¥  r\  r  r-   rR   rG   r§  r  r¿   r~  r(  rÊ  r.   rË  r‡   rÓ  rB  rÈ  Úloss_functionrt   Úrepad_logits_with_gradr   r„  r   rˆ   r¤  )rZ   r‚   rÎ   rÏ   rº   r…   rj  rv  r$   r%   r—  r˜  rÒ   r™  rš  r  rR   Úoutputsr£  Úmask_tokensrØ  r×  r�  s                          r9   r:   zModernBertForMaskedLM.forward0  sì  € ð0 &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆØ×ÑÔ!à�;‰;×+Ñ+Ð/BÒBØˆ :Ð#5¸*Ð:LØÐ%¨'¨/Ø$Ð0Ø.;×.AÑ.AÀ"À1Ð.EÑ+˜
¡Gà.7¯o©o¸b¸qÐ.AÑ+˜
 GØ-6Ð-B˜×)Ò)È×H\ÑH\�à!Ð)Ü%*§Z¡Z°¸WÐ0EÈfÔ\a×\fÑ\fÔ%g�Nà Ð(ÜŸ™›ñ Ü[rØ#,¸^ÐZfÐouô\ÑX˜	 7¨J¸
ÀLÐRX÷ð ô
 \sØ,¸^ÐZfÐouô\ÑX�M 7¨J¸
ÀLÐRXð —*‘*ØØ)Ø 3Ø%Ø'ØØ!Ø!Ø!ØØ/Ø!5Ø#ð ó 
ˆð $ A™JÐà×!Ò! fÐ&8à—[‘[ “_ˆFØ 1× 6Ñ 6°v·|±|ÀA±ÈÓ KÐð ! D×$AÑ$AÑAˆKØ 1°+Ñ >ÐØ˜KÑ(ˆFð �{‰{×,Ò,ð ×ÑÐ0Ô1à—‘˜dŸi™iÐ(9Ó:Ó;ð 	ð ˆØÐØ×%Ñ% f¨fÀÇÁ×AWÑAWÐ%ÓXˆDà�;‰;×+Ñ+Ð/BÒBØ"&§+¡+×"DÒ"DÈÈ””Ô\a×\iÑ\iÓ\kñ rÜ/°vÀwÐV`ÐipÔq�÷rñ Ø�YˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEäØØØ!×/Ñ/Ø×)Ñ)ô	
ð 	
÷mñ ú÷^rð rús   ÃJÉJÊJÊJ!©NNNNNNNNNNNNNN)rC   rD   rE   Ú_tied_weights_keysr   rX   rÎ  r   r�   rÑ  rG   r‰   rH   rÓ  r   r¾  r   r¿  r   rÀ  r   rŠ   rI   r  r   r   r:   ri   rj   s   @r9   rA  rA    sÜ  ø„ ð
 +Ð+ÐðÐ/õ òð&°B·I±Ió &ð €U‡]�]˜4Ô ð/ E§L¡Lð /°U·\±\ò /ó !ð/ñ +Ð+FÓGÙØ&Ø"Ø$ôð 15Ø15Ø6:Ø/3Ø04Ø)-Ø*.Ø-1Ø$(Ø$(Ø!%Ø,0Ø/3Ø&*ñ]
à˜E×,Ñ,Ñ-ð]
ð ! §¡Ñ.ð]
ð & e§l¡lÑ3ð	]
ð
 ˜uŸ|™|Ñ,ð]
ð   §¡Ñ-ð]
ð ˜Ÿ™Ñ&ð]
ð ˜%Ÿ,™,Ñ'ð]
ð ˜UŸ\™\Ñ*ð]
ð ˜S‘Mð]
ð ˜S‘Mð]
ð ˜#‘ð]
ð $ D™>ð]
ð ' t™nð]
ð ˜d‘^ð]
ð" 
ˆu�U—\‘\Ñ" NÐ2Ñ	3ò#]
óó Hô]
r;   rA  zVThe ModernBert Model with a sequence classification head on top that performs pooling.c            #       óÐ  ‡ — e Zd Zdefˆ fd„Z ee«       eee	e
¬«      	 	 	 	 	 	 	 	 	 	 	 	 	 	 ddeej                     deej                     deej                     deej                     deej                     d	eej                     d
eej                     deej                     dee   dee   dee   dee   dee   dee   deeej                     e	f   fd„«       «       Zˆ xZS )rC  rm   c                 ón  •— t         ‰| �  |«       |j                  | _        || _        t	        |«      | _        t        |«      | _        t        j                  j                  |j                  «      | _        t        j                  |j                  |j                  «      | _        | j!                  «        y re   )rW   rX   Ú
num_labelsrm   r†  r(  r?  rÈ  rG   r   r|   Úclassifier_dropoutr~   r�   ru   rF  rŽ  r   s     €r9   rX   z,ModernBertForSequenceClassification.__init__›  s‚   ø€ Ü‰Ñ˜Ô Ø ×+Ñ+ˆŒØˆŒä$ VÓ,ˆŒ
Ü,¨VÓ4ˆŒ	Ü—H‘H×$Ñ$ V×%>Ñ%>Ó?ˆŒ	ÜŸ)™) F×$6Ñ$6¸×8IÑ8IÓJˆŒð 	�‰Õr;   r”  r‚   rÎ   rÏ   rº   r…   rj  rv  r$   r%   r—  r˜  rÒ   r™  rš  r\   c                 óf  — |�|n| j                   j                  }| j                  «        | j                  ||||||||	|
||||¬«      }|d   }| j                   j                  dk(  r
|dd…df   }nQ| j                   j                  dk(  r8||j                  d«      z  j                  d¬«      |j                  dd	¬
«      z  }| j                  |«      }| j                  |«      }| j                  |«      }d}|��‡| j                   j                  €�| j                  dk(  rd| j                   _
        nl| j                  dkD  rL|j                  t        j                  k(  s|j                  t        j                  k(  rd| j                   _
        nd| j                   _
        | j                   j                  dk(  rIt!        «       }| j                  dk(  r& ||j#                  «       |j#                  «       «      }nŒ |||«      }n‚| j                   j                  dk(  r=t%        «       } ||j'                  d| j                  «      |j'                  d«      «      }n,| j                   j                  dk(  rt)        «       } |||«      }|s|f}|�|f|z   S |S t+        |||j,                  |j.                  ¬«      S )a�  
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the sequence classification/regression loss. Indices should be in `[0, ...,
            config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If
            `config.num_labels > 1` a classification loss is computed (Cross-Entropy).
        NrÕ  r   rV  r,  r(   r   r˜   T)rP   ÚkeepdimÚ
regressionÚsingle_label_classificationÚmulti_label_classificationrÖ  )rm   r¥  r\  r(  Úclassifier_poolingrÅ   rm  rÈ  r~   rF  Úproblem_typerá  rS   rG   ÚlongrI   r   Úsqueezer
   r.   r	   r   rˆ   r¤  )rZ   r‚   rÎ   rÏ   rº   r…   rj  rv  r$   r%   r—  r˜  rÒ   r™  rš  r  rÛ  r£  Úpooled_outputrØ  r×  Úloss_fctr�  s                          r9   r:   z+ModernBertForSequenceClassification.forward¨  s  € ð< &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆØ×ÑÔ!à—*‘*ØØ)Ø 3Ø%Ø'ØØ!Ø!Ø!ØØ/Ø!5Ø#ð ó 
ˆð $ A™JÐà�;‰;×)Ñ)¨UÒ2Ø 1²!°Q°$Ñ 7ÑØ�[‰[×+Ñ+¨vÒ5Ø!2°^×5MÑ5MÈbÓ5QÑ!Q× VÑ VÐ[\Ð VÓ ]Ð`n×`rÑ`rØ˜tð asó añ !Ðð Ÿ	™	Ð"3Ó4ˆØŸ	™	 -Ó0ˆØ—‘ Ó/ˆàˆØÑØ�{‰{×'Ñ'Ð/Ø—?‘? aÒ'Ø/;�D—K‘KÕ,Ø—_‘_ qÒ(¨f¯l©l¼e¿j¹jÒ.HÈFÏLÉLÔ\a×\eÑ\eÒLeØ/L�D—K‘KÕ,à/K�D—K‘KÔ,à�{‰{×'Ñ'¨<Ò7Ü"›9�Ø—?‘? aÒ'Ù# F§N¡NÓ$4°f·n±nÓ6FÓG‘Dá# F¨FÓ3‘DØ—‘×)Ñ)Ð-JÒJÜ+Ó-�Ù §¡¨B°·±Ó @À&Ç+Á+ÈbÃ/ÓR‘Ø—‘×)Ñ)Ð-IÒIÜ,Ó.�Ù ¨Ó/�áØ�YˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä'ØØØ!×/Ñ/Ø×)Ñ)ô	
ð 	
r;   rÝ  )rC   rD   rE   r   rX   r   r¾  r   r¿  r   rÀ  r   rG   rŠ   rH   rI   r  r   r   r:   ri   rj   s   @r9   rC  rC  –  sˆ  ø„ ð
Ð/õ ñ +Ð+FÓGÙØ&Ø,Ø$ôð 15Ø15Ø6:Ø/3Ø04Ø)-Ø*.Ø-1Ø$(Ø$(Ø!%Ø,0Ø/3Ø&*ñW
à˜E×,Ñ,Ñ-ðW
ð ! §¡Ñ.ðW
ð & e§l¡lÑ3ð	W
ð
 ˜uŸ|™|Ñ,ðW
ð   §¡Ñ-ðW
ð ˜Ÿ™Ñ&ðW
ð ˜%Ÿ,™,Ñ'ðW
ð ˜UŸ\™\Ñ*ðW
ð ˜S‘MðW
ð ˜S‘MðW
ð ˜#‘ðW
ð $ D™>ðW
ð ' t™nðW
ð ˜d‘^ðW
ð" 
ˆu�U—\‘\Ñ"Ð$<Ð<Ñ	=ò#W
óó HôW
r;   rC  zlThe ModernBert Model with a token classification head on top, e.g. for Named Entity Recognition (NER) tasks.c            #       óÐ  ‡ — e Zd Zdefˆ fd„Z ee«       eee	e
¬«      	 	 	 	 	 	 	 	 	 	 	 	 	 	 ddeej                     deej                     deej                     deej                     deej                     d	eej                     d
eej                     deej                     dee   dee   dee   dee   dee   dee   deeej                     e	f   fd„«       «       Zˆ xZS )rD  rm   c                 ó`  •— t         ‰| �  |«       |j                  | _        t        |«      | _        t        |«      | _        t        j                  j                  |j                  «      | _        t        j                  |j                  |j                  «      | _        | j                  «        y re   ©rW   rX   rá  r†  r(  r?  rÈ  rG   r   r|   râ  r~   r�   ru   rF  rŽ  r   s     €r9   rX   z)ModernBertForTokenClassification.__init__  s{   ø€ Ü‰Ñ˜Ô Ø ×+Ñ+ˆŒä$ VÓ,ˆŒ
Ü,¨VÓ4ˆŒ	Ü—H‘H×$Ñ$ V×%>Ñ%>Ó?ˆŒ	ÜŸ)™) F×$6Ñ$6¸×8IÑ8IÓJˆŒð 	�‰Õr;   r”  r‚   rÎ   rÏ   rº   r…   rj  rv  r$   r%   r—  r˜  rÒ   r™  rš  r\   c                 óò  — |�|n| j                   j                  }| j                  «        | j                  ||||||||	|
||||¬«      }|d   }| j	                  |«      }| j                  |«      }| j                  |«      }d}|�<t        «       } ||j                  d| j                  «      |j                  d«      «      }|s|f|dd z   }|�|f|z   S |S t        |||j                  |j                  ¬«      S )zÛ
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the token classification loss. Indices should be in `[0, ..., config.num_labels - 1]`.
        NrÕ  r   r(   r   rÖ  )rm   r¥  r\  r(  rÈ  r~   rF  r
   r.   rá  r   rˆ   r¤  )rZ   r‚   rÎ   rÏ   rº   r…   rj  rv  r$   r%   r—  r˜  rÒ   r™  rš  rÛ  r£  rØ  r×  rí  r�  s                        r9   r:   z(ModernBertForTokenClassification.forward  s"  € ð6 &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆØ×ÑÔ!à—*‘*ØØ)Ø 3Ø%Ø'ØØ!Ø!Ø!ØØ/Ø!5Ø#ð ó 
ˆð $ A™JÐà ŸI™IÐ&7Ó8ÐØ ŸI™IÐ&7Ó8ÐØ—‘Ð!2Ó3ˆàˆØÐÜ'Ó)ˆHÙ˜FŸK™K¨¨D¯O©OÓ<¸f¿k¹kÈ"»oÓNˆDáØ�Y ¨¨ Ñ,ˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä$ØØØ!×/Ñ/Ø×)Ñ)ô	
ð 	
r;   rÝ  )rC   rD   rE   r   rX   r   r¾  r   r¿  r   rÀ  r   rG   rŠ   rH   rI   r  r   r   r:   ri   rj   s   @r9   rD  rD    sw  ø„ ð

Ð/õ 
ñ +Ð+FÓGÙØ&Ø)Ø$ôð 15Ø15Ø6:Ø/3Ø04Ø)-Ø*.Ø-1Ø$(Ø$(Ø!%Ø,0Ø/3Ø&*ñ;
à˜E×,Ñ,Ñ-ð;
ð ! §¡Ñ.ð;
ð & e§l¡lÑ3ð	;
ð
 ˜uŸ|™|Ñ,ð;
ð   §¡Ñ-ð;
ð ˜Ÿ™Ñ&ð;
ð ˜%Ÿ,™,Ñ'ð;
ð ˜UŸ\™\Ñ*ð;
ð ˜S‘Mð;
ð ˜S‘Mð;
ð ˜#‘ð;
ð $ D™>ð;
ð ' t™nð;
ð ˜d‘^ð;
ð  
ˆu�U—\‘\Ñ"Ð$9Ð9Ñ	:ò!;
óó Hô;
r;   rD  zæ
    The ModernBert Model with a span classification head on top for extractive question-answering tasks like SQuAD
    (a linear layer on top of the hidden-states output to compute `span start logits` and `span end logits`).
    c            #       óÎ  ‡ — e Zd Zdefˆ fd„Z ee«       eee	e
¬«      	 	 	 	 	 	 	 	 	 	 	 	 	 ddeej                     deej                     deej                     deej                     deej                     d	eej                     d
eej                     deej                     dee   dee   dee   dee   dee   dee   deeej                     e	f   fd„«       «       Zˆ xZS )rE  rm   c                 ó`  •— t         ‰| �  |«       |j                  | _        t        |«      | _        t        |«      | _        t        j                  j                  |j                  «      | _        t        j                  |j                  |j                  «      | _        | j                  «        y re   rð  r   s     €r9   rX   z'ModernBertForQuestionAnswering.__init__e  sy   ø€ Ü‰Ñ˜Ô Ø ×+Ñ+ˆŒä$ VÓ,ˆŒ
Ü,¨VÓ4ˆŒ	Ü—H‘H×$Ñ$ V×%>Ñ%>Ó?ˆŒ	ÜŸ)™) F×$6Ñ$6¸×8IÑ8IÓJˆŒà�‰Õr;   r”  r‚   rÎ   rÏ   rº   Ústart_positionsÚend_positionsrv  r$   r%   r—  r˜  rÒ   r™  rš  r\   c                 óT  — |�|n| j                   j                  }| j                  «        | j                  |||||||	|
||||¬«      }|d   }| j	                  |«      }| j                  |«      }| j                  |«      }|j                  dd¬«      \  }}|j                  d«      j                  «       }|j                  d«      j                  «       }d }|�|� | j                  ||||fi |¤Ž}|s||f|dd  z   }|�|f|z   S |S t        ||||j                  |j                  ¬«      S )N)rÎ   rÏ   rº   rv  r$   r%   r—  r˜  rÒ   r™  rš  r   r   r(   r˜   )r×  Ústart_logitsÚ
end_logitsrˆ   r¤  )rm   r¥  r\  r(  rÈ  r~   rF  Úsplitrë  r,   rÙ  r   rˆ   r¤  )rZ   r‚   rÎ   rÏ   rº   rô  rõ  rv  r$   r%   r—  r˜  rÒ   r™  rš  r  rÛ  r£  rØ  r÷  rø  r×  r�  s                          r9   r:   z&ModernBertForQuestionAnswering.forwardp  sg  € ð0 &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆØ×ÑÔ!à—*‘*ØØ)Ø 3Ø%ØØ!Ø!Ø!ØØ/Ø!5Ø#ð ó 
ˆð $ A™JÐà ŸI™IÐ&7Ó8ÐØ ŸI™IÐ&7Ó8ÐØ—‘Ð!2Ó3ˆà#)§<¡<°°r <Ó#:Ñ ˆ�jØ#×+Ñ+¨BÓ/×:Ñ:Ó<ˆØ×'Ñ'¨Ó+×6Ñ6Ó8ˆ
àˆØÐ&¨=Ð+DØ%�4×%Ñ% l°JÀÐQ^ÑiÐbhÑiˆDáØ" JÐ/°'¸!¸"°+Ñ=ˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä+ØØ%Ø!Ø!×/Ñ/Ø×)Ñ)ô
ð 	
r;   r½  )rC   rD   rE   r   rX   r   r¾  r   r¿  r   rÀ  r   rG   rH   rI   r  r   r   r:   ri   rj   s   @r9   rE  rE  ]  sr  ø„ ð	Ð/õ 	ñ +Ð+FÓGÙØ&Ø0Ø$ôð 26Ø6:Ø/3Ø26Ø04Ø*.Ø-1Ø$(Ø$(Ø!%Ø,0Ø/3Ø&*ñ;
à˜EŸL™LÑ)ð;
ð ! §¡Ñ.ð;
ð & e§l¡lÑ3ð	;
ð
 ˜uŸ|™|Ñ,ð;
ð " %§,¡,Ñ/ð;
ð   §¡Ñ-ð;
ð ˜%Ÿ,™,Ñ'ð;
ð ˜UŸ\™\Ñ*ð;
ð ˜S‘Mð;
ð ˜S‘Mð;
ð ˜#‘ð;
ð $ D™>ð;
ð ' t™nð;
ð ˜d‘^ð;
ð" 
ˆu�U—\‘\Ñ"Ð$@Ð@Ñ	Aò#;
óó Hô;
r;   rE  )r†  r'  rA  rC  rD  rE  rB   )Nr   r  )Yr<  Ú
contextlibr   Útypingr   r   r   r   rG   Útorch.nn.functionalr   rÞ   rú   Útorch.nnr	   r
   r   Úactivationsr   Úmodeling_attn_mask_utilsr   Úmodeling_outputsr   r   r   r   r   Úmodeling_rope_utilsr   r   Úmodeling_utilsr   Úutilsr   r   r   r   r   Úutils.import_utilsr   Úconfiguration_modernbertr   Úflash_attn.flash_attn_interfacer   Úflash_attn.layers.rotaryr    Úflash_attn.ops.triton.rotaryr!   ÚobjectÚ
get_loggerrC   rZ  r¿  rÀ  ÚautogradÚFunctionr#   rH   rI   rM   rO   r:  rl   rŒ   r�   rÃ   rË   rŠ   r  rê   rò   rS   r÷   rü   r  rÍ   r  ÚMODERNBERT_START_DOCSTRINGr'  r~  r„  r¾  r†  r?  rA  rC  rD  rE  Ú__all__rJ   r;   r9   ú<module>r     s  ðó, Ý "ß /Ó /ã ß Ð Ý ß AÑ Aå !Ý B÷õ ÷ LÝ -÷õ õ 6Ý 6ñ ÔÝPÝ8Þ9à€Oà	ˆ×	Ñ	˜HÓ	%€à3Ð Ø$€ô46˜%Ÿ.™.×1Ñ1ô 46ðv *.Ø $ñLð ˜Ÿ™Ñ&ð	Lð
 ˜‘óLô42Q¨ô 2Qôj˜2Ÿ9™9ô ô<:�B—I‘Iô :ô(< §	¡	ô <òB(óðH ).ñ"Ø!ð"à	�‰ð"ð —L‘Lð"ð Ÿ™ð	"ð
 ˜5×+Ñ+Ñ,ð"ð ˜3 ˜8‘_ð"ð 	ð"ð 
ð"ð   ‘~ð"ð ˆ5�—‘˜uŸ|™|Ð+Ñ,¨e°E·L±LÑ.AÐAÑBó"ð\ !&§¡ñ(!Ø!ð(!à	�‰ð(!ð 2ð(!ð —‘ð	(!ð
 ð(!ð ˜3 ˜8‘_ð(!ð 	ð(!ð 
ð(!ð —+‘+ð(!ð ˆ5�<‰<Ñó(!ðV Ø!ð à	�‰ð ð —L‘Lð ð Ÿ™ð	 ð
 ˜5×+Ñ+Ñ,ð ð ˜3 ˜8‘_ð ð 	ð ð 
ð ð ˆ5�<‰<Ñó ðH 1Ø$Ø"ñ!Ð ôM3˜"Ÿ)™)ô M3ô`+3˜RŸY™Yô +3ð\Ð ñ" Ø^ØóôB ó Bó	ðBðP ,0Ø%)ñ	&mØ�L‰Lð&mà—L‘Lð&mð ˜5Ÿ<™<Ñ(ð&mð �U—\‘\Ñ"ð	&mð
 ˆ5�<‰<˜Ÿ™ u§|¡|°S¸(À5Ç<Á<Ñ:PÐRZÐ[`×[gÑ[gÑRhÐhÑió&mðRØ�L‰Lðà�\‰\ðð ðð ð	ð
 ‡\�\óð>:Ð ñz Ø^Øóôk:Ð/ó k:ó	ðk:ô\	>˜rŸy™yô 	>ñ Ø`Øóô}
Ð5ó }
ó	ð}
ñ@ Ø\Øóôk
Ð*Có k
ó	ðk
ñ\ ØrØóôN
Ð'@ó N
ó	ðN
ñb ðð óôM
Ð%>ó M
óðM
ò`�r;   