Ë
    T^(h:š  ã                   ó�  — d Z ddlZddlmZmZmZ ddlZddl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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 ddl m!Z! ddl"m#Z#  e!jH                  e%«      Z&dZ'dZ(d-d„Z) G d„ dejT                  «      Z+ G d„ dejT                  «      Z, G d„ dejT                  «      Z- G d„ de«      Z.dZ/dZ0 ede/«       G d„ de.«      «       Z1 ed e/«       G d!„ d"e.e«      «       Z2 ed#e/«       G d$„ d%e.«      «       Z3 ed&e/«       G d'„ d(e.«      «       Z4 ed)e/«       G d*„ d+e.«      «       Z5g d,¢Z6y).zPyTorch MPT model.é    N)ÚOptionalÚTupleÚUnion)Únn)ÚBCEWithLogitsLossÚCrossEntropyLossÚ	LayerNormÚMSELoss)Ú
functionalé   )Úadd_code_sample_docstringsÚadd_start_docstringsÚ%add_start_docstrings_to_model_forward)ÚGenerationMixin)Ú!_prepare_4d_causal_attention_mask)Ú)BaseModelOutputWithPastAndCrossAttentionsÚ!CausalLMOutputWithCrossAttentionsÚQuestionAnsweringModelOutputÚ SequenceClassifierOutputWithPastÚTokenClassifierOutput)ÚPreTrainedModel)Úloggingé   )Ú	MptConfigzmosaicml/mpt-7br   c                 óR  — t        j                  d|z
  dt         j                  |¬«      j                  ddd|«      }dt	        j
                  t	        j                  | «      «      z  }t        j                  d|dz   t         j                  |¬«      j                  «       }|||z  z  }dt        j                  d|«      z  }|j                  d|dd«      }|| k7  r9t        j                  |dd…ddd…df   |dd…ddd…df   gd¬«      dd…d| …df   }||z  }|j                  d«      S )	a¢  
    Link to paper: https://arxiv.org/abs/2108.12409 - Alibi tensor is not causal as the original paper mentions, it
    relies on a translation invariance of softmax for quick implementation. This implementation has been copied from
    the alibi implementation of MPT source code that led to slightly different results than the Bloom alibi:
    https://huggingface.co/mosaicml/mpt-7b/blob/main/attention.py#L292
    r   )ÚdtypeÚdeviceé   ç      ð?N.©Údimr   )ÚtorchÚarangeÚint32ÚviewÚmathÚceilÚlog2Úint64ÚfloatÚpowÚconcatÚsqueeze)Ú	num_headsÚsequence_lengthÚalibi_bias_maxr   ÚalibiÚnum_heads_power_of_2ÚbaseÚslopess           úb/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/mpt/modeling_mpt.pyÚbuild_mpt_alibi_tensorr6   /   s  € ô �L‰L˜˜_Ñ,¨a´u·{±{È6ÔR×WÑWÐXYÐ[\Ð^_ÐapÓq€EØ¤§	¡	¬$¯)©)°IÓ*>Ó ?Ñ?Ðä�<‰<˜Ð/°!Ñ3¼5¿;¹;ÈvÔV×\Ñ\Ó^€DØ�>Ð$8Ñ8Ñ9€Dà”5—9‘9˜Q Ó%Ñ%€FØ�[‰[˜Ð0°!°QÓ7€Fà˜yÒ(Ü—‘˜v¢a¨¨¨A¨¨s lÑ3°VºA¹sÀ¸sÀC¸KÑ5HÐIÈqÔQÒRSÐU_ÐV_ÐU_ÐadÐRdÑeˆà�F‰N€EØ�=‰=˜ÓÐó    c            
       ó¨   ‡ — e Zd ZdZdefˆ fd„Z	 	 d	dej                  dej                  dee	ej                        deej                     fd„Z
ˆ xZS )
ÚMptAttentionzyMulti-head self attention.
    Using torch or triton attention implemetation enables user to also use additive bias.
    Úconfigc                 ó°  •— t         ‰| �  «        |j                  | _        |j                  | _        |j                  | _        | j                  | j                  z  | _        |j                  j                  | _        | j                  €4dt        j                  | j                  | j                  z  «      z  | _        |j                  j                  | _        |j                  j                  | _        t        j                  | j                  d| j                  z  d¬«      | _        t        j                  | j                  | j                  d¬«      | _        y )Nr   r   F©Úbias)ÚsuperÚ__init__Úhidden_sizeÚn_headsÚmax_seq_lenÚmax_seq_lengthÚhead_dimÚattn_configÚsoftmax_scaler&   ÚsqrtÚ
attn_pdropÚattn_dropout_pÚclip_qkvr   ÚLinearÚWqkvÚout_proj©Úselfr:   Ú	__class__s     €r5   r?   zMptAttention.__init__K   sü   ø€ Ü‰ÑÔØ!×-Ñ-ˆÔØ—~‘~ˆŒØ$×0Ñ0ˆÔØ×(Ñ(¨D¯L©LÑ8ˆŒØ#×/Ñ/×=Ñ=ˆÔØ×ÑÐ%Ø!"¤T§Y¡Y¨t×/?Ñ/?À$Ç,Á,Ñ/NÓ%OÑ!OˆDÔà$×0Ñ0×;Ñ;ˆÔØ×*Ñ*×3Ñ3ˆŒÜ—I‘I˜d×.Ñ.°°D×4DÑ4DÑ0DÈ5ÔQˆŒ	ÜŸ	™	 $×"2Ñ"2°D×4DÑ4DÈ5ÔQˆ�r7   Úhidden_statesÚposition_biasÚpast_key_valueÚattention_maskc                 óÊ  — |j                   d d \  }}| j                  |«      }| j                  r(|j                  | j                   | j                  ¬«      }|j	                  dd¬«      \  }}	}
|j                  ||| j                  | j                  «      j                  dd«      }|	j                  ||| j                  | j                  «      j                  dd«      }	|
j                  ||| j                  | j                  «      j                  dd«      }
|�Kt        |«      dk7  r8t        j                  |d   |	gd¬«      }	t        j                  |d   |
gd¬«      }
|	|
f}n|	|
f}t        j                  ||	j                  dd«      «      | j                  z  }|€|n||d   j                   d   z   }|�—t        |j                   «      dk7  r!t        d	t        |j                   «      › �«      ‚|	j                   d   }t        d|j!                  d«      |z
  «      }t        d|j!                  d«      |z
  «      }|d d …|d …|d …f   }||z   }|�9|j#                  |t        j$                  |j&                  «      j(                  «      }t*        j,                  j/                  |j1                  «       d¬«      j3                  |
j&                  «      }t*        j,                  j5                  || j6                  | j8                  ¬
«      }t        j                  ||
«      }|j;                  dddd«      j=                  «       j?                  ||d«      }| jA                  |«      }|||fS )Nr   )ÚminÚmaxr   r    r   r   éÿÿÿÿéþÿÿÿz6Expecting position_bias shape to be 3 dimensions, got ©ÚpÚtraining)!ÚshaperL   rJ   ÚclampÚchunkÚreshaperA   rD   Ú	transposeÚlenr"   ÚcatÚmatmulrF   Ú
ValueErrorrW   ÚsizeÚmasked_fillÚfinfor   rV   r   r   Úsoftmaxr*   ÚtoÚdropoutrI   r\   ÚpermuteÚ
contiguousr%   rM   )rO   rQ   rR   rS   rT   Ú
batch_sizeÚ
seq_lengthÚ	mixed_qkvÚquery_statesÚ
key_statesÚvalue_statesÚattention_scoresÚquery_lengthÚ
key_lengthÚposition_bias_query_indexÚposition_bias_key_indexÚattn_weightsÚcontext_statesÚattn_outputs                      r5   ÚforwardzMptAttention.forwardZ   s,  € ð "/×!4Ñ!4°R°aÐ!8Ñˆ
�Jà—I‘I˜mÓ,ˆ	Ø�=Š=Ø!Ÿ™¨T¯]©]¨NÀÇÁ˜ÓNˆIà1:·±ÀÈ°Ó1JÑ.ˆ�j ,Ø#×+Ñ+¨J¸
ÀDÇLÁLÐRV×R_ÑR_Ó`×jÑjÐklÐnoÓpˆØ×'Ñ'¨
°JÀÇÁÈdÏmÉmÓ\×fÑfÐghÐjkÓlˆ
Ø#×+Ñ+¨J¸
ÀDÇLÁLÐRV×R_ÑR_Ó`×jÑjÐklÐnoÓpˆàÐ%Ü�>Ó" aÒ'Ü"ŸY™Y¨°qÑ(9¸:Ð'FÈAÔN�
Ü$Ÿy™y¨.¸Ñ*;¸\Ð)JÐPQÔR�Ø(¨,Ð7‰Nà(¨,Ð7ˆNä Ÿ<™<¨°j×6JÑ6JÈ2ÈrÓ6RÓSÐVZ×VhÑVhÑhÐà%3Ð%;‘zÀÈnÐ]^ÑN_×NeÑNeÐfgÑNhÑAhˆàÐ$Ü�=×&Ñ&Ó'¨1Ò,Ü Ð#YÔZ]Ð^k×^qÑ^qÓZrÐYsÐ!tÓuÐuØ#×)Ñ)¨"Ñ-ˆJä(+¨A¨}×/AÑ/AÀ!Ó/DÀ|Ñ/SÓ(TÐ%Ü&)¨!¨]×-?Ñ-?ÀÓ-BÀZÑ-OÓ&PÐ#à)ª!Ð-FÑ-GÐI`ÑIaÐ*aÑbˆMà/°-Ñ?ÐàÐ%Ø/×;Ñ;¸NÌEÏKÉKÐXd×XjÑXjÓLk×LoÑLoÓpÐô —}‘}×,Ñ,Ð-=×-CÑ-CÓ-EÈ2Ð,ÓN×QÑQÐR^×RdÑRdÓeˆÜ—}‘}×,Ñ,¨\¸T×=PÑ=PÐ[_×[hÑ[hÐ,ÓiˆäŸ™ l°LÓAˆØ'×/Ñ/°°1°a¸Ó;×FÑFÓH×MÑMÈjÐZdÐfhÓiˆØ—m‘m NÓ3ˆà˜L¨.Ð8Ð8r7   )NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r?   r"   ÚTensorr   r   r|   Ú__classcell__©rP   s   @r5   r9   r9   F   sh   ø„ ñðR˜yõ Rð& 9=Ø15ñ59à—|‘|ð59ð —|‘|ð59ð !  u§|¡|Ñ!4Ñ5ð	59ð
 ! §¡Ñ.÷59r7   r9   c                   ót   ‡ — e Zd Zdefˆ fd„Zdej                  dej                  dej                  fd„Zˆ xZS )ÚMptMLPr:   c                 ó&  •— t         ‰| �  «        |j                  }t        j                  |d|z  d¬«      | _        t        j                  d¬«      | _        t        j                  d|z  |d¬«      | _        |j                  j                  | _        y )Né   Fr<   Únone)Úapproximate)r>   r?   r@   r   rK   Úup_projÚGELUÚactÚ	down_projrE   rH   Úhidden_dropout©rO   r:   r@   rP   s      €r5   r?   zMptMLP.__init__“   sm   ø€ Ü‰ÑÔØ×(Ñ(ˆä—y‘y ¨a°+©oÀEÔJˆŒÜ—7‘7 vÔ.ˆŒÜŸ™ 1 {¡?°KÀeÔLˆŒØ$×0Ñ0×;Ñ;ˆÕr7   rQ   ÚresidualÚreturnc                 óÊ   — | j                  | j                  |«      «      }| j                  |«      }t        j                  || j
                  | j                  ¬«      }||z   }|S )NrZ   )rŒ   rŠ   r�   ÚFrk   rŽ   r\   )rO   rQ   r�   Úintermediate_outputÚoutputs        r5   r|   zMptMLP.forwardœ   sW   € ØŸ™ §¡¨mÓ!<Ó=ˆà"Ÿn™n¨]Ó;Ðä—‘Ð.°$×2EÑ2EÐPT×P]ÑP]Ô^ˆØ˜(Ñ"ˆàˆr7   )	r}   r~   r   r   r?   r"   r�   r|   r‚   rƒ   s   @r5   r…   r…   ’   s5   ø„ ð<˜yõ <ð U§\¡\ð ¸U¿\¹\ð ÈeÏlÉl÷ r7   r…   c                   óÀ   ‡ — e Zd Zdefˆ fd„Z	 	 	 d
dej                  dej                  dej                  deeej                  ej                  f      de	de	fd	„Z
ˆ xZS )ÚMptBlockr:   c                 óÎ  •— t         ‰| �  «        |j                  }t        ||j                  ¬«      | _        d | j
                  _        |j                  | _        t        |«      | _
        t        ||j                  ¬«      | _        d | j                  _        t        |«      | _        |j                  j                  | _        t#        j$                  | j                   «      | _        y )N©Úeps)r>   r?   r@   r	   Úlayer_norm_epsilonÚnorm_1r=   rA   r.   r9   ÚattnÚnorm_2r…   ÚffnrE   rH   Údropout_rater   ÚDropoutÚresid_attn_dropoutr�   s      €r5   r?   zMptBlock.__init__¨   s¦   ø€ Ü‰ÑÔØ×(Ñ(ˆä °×1JÑ1JÔKˆŒàˆ�‰ÔàŸ™ˆŒÜ  Ó(ˆŒ	ä °×1JÑ1JÔKˆŒàˆ�‰Ôä˜&“>ˆŒà"×.Ñ.×9Ñ9ˆÔÜ"$§*¡*¨T×->Ñ->Ó"?ˆÕr7   rQ   rR   rT   Ú
layer_pastÚ	use_cacheÚoutput_attentionsc                 óö   — | j                  |«      }|}| j                  ||||¬«      \  }	}
}| j                  |	«      |z   }| j                  |«      }|}| j	                  ||«      }|f}|r||fz  }|r||
fz  }|S )N)rR   rT   rS   )rœ   r�   r¢   rž   rŸ   )rO   rQ   rR   rT   r£   r¤   r¥   Úlayernorm_outputr�   Úattn_outputsry   rS   r•   Úoutputss                 r5   r|   zMptBlock.forward¼   s«   € ð  Ÿ;™; }Ó5Ðà ˆð 6:·Y±YØØ'Ø)Ø%ð	 6?ó 6
Ñ2ˆ�l Nð ×/Ñ/°Ó=ÀÑHˆàŸ;™; }Ó5Ðð !ˆð —‘Ð*¨HÓ5ˆØ�)ˆáØ˜Ð(Ñ(ˆGáØ˜�Ñ&ˆGàˆr7   )NFF)r}   r~   r   r   r?   r"   r�   r   r   Úboolr|   r‚   rƒ   s   @r5   r—   r—   §   s€   ø„ ð@˜yõ @ð2 CGØØ"'ñ(à—|‘|ð(ð —|‘|ð(ð Ÿ™ð	(ð
 ˜U 5§<¡<°·±Ð#=Ñ>Ñ?ð(ð ð(ð  ÷(r7   r—   c                   óà   ‡ — e Zd ZeZdZdZdgZdgZˆ fd„Z	de
j                  fd„Zedeeej                   ej                   f      d	eeej                   ej                   f      fd
„«       Zˆ xZS )ÚMptPreTrainedModelÚtransformerTr—   z
lm_head.*.c                 ó$   •— t        ‰| �  |i |¤Ž y ©N)r>   r?   )rO   ÚinputsÚkwargsrP   s      €r5   r?   zMptPreTrainedModel.__init__î   s   ø€ Ü‰Ñ˜&Ð+ FÓ+r7   Úmodulec                 ó  — t        |t        j                  «      rm|j                  j                  j                  d| j                  j                  ¬«       |j                  �%|j                  j                  j                  «        yyt        |t        j                  «      rz|j                  j                  j                  d| j                  j                  ¬«       |j                  �2|j                  j                  |j                     j                  «        yyt        |t        «      rV|j                  �$|j                  j                  j                  «        |j                  j                  j                  d«       yy)zInitialize the weights.g        )ÚmeanÚstdNr   )Ú
isinstancer   rK   ÚweightÚdataÚnormal_r:   Úinitializer_ranger=   Úzero_Ú	EmbeddingÚpadding_idxr	   Úfill_)rO   r²   s     r5   Ú_init_weightsz MptPreTrainedModel._init_weightsñ   s  € ä�fœbŸi™iÔ(ð �M‰M×Ñ×&Ñ&¨C°T·[±[×5RÑ5RÐ&ÔSØ�{‰{Ð&Ø—‘× Ñ ×&Ñ&Õ(ð 'ä˜¤§¡Ô-Ø�M‰M×Ñ×&Ñ&¨C°T·[±[×5RÑ5RÐ&ÔSØ×!Ñ!Ð-Ø—‘×"Ñ" 6×#5Ñ#5Ñ6×<Ñ<Õ>ð .ä˜¤	Ô*Ø�{‰{Ð&Ø—‘× Ñ ×&Ñ&Ô(Ø�M‰M×Ñ×$Ñ$ SÕ)ð +r7   rS   r‘   c                 ól   ‡‡‡— | d   d   j                   \  }}ŠŠ||z  Št        ˆˆˆfd„| D «       «      S )zw
        Converts the cache to the format expected by Mpt, i.e. to tuple(tuple([batch_size * num_heads, ...]))
        r   c              3   óv   •K  — | ]0  }|d    j                  ‰‰‰«      |d   j                  ‰‰‰«      f–— Œ2 y­w©r   r   N)r`   )Ú.0r£   Úbatch_size_times_num_headsrD   ro   s     €€€r5   ú	<genexpr>z;MptPreTrainedModel._convert_to_mpt_cache.<locals>.<genexpr>  sK   øè ø€ ò 
ð
 ð ˜1‘×%Ñ%Ð&@À(ÈJÓWØ˜1‘×%Ñ%Ð&@À*ÈhÓWôñ
ùs   ƒ69)r]   Útuple)rS   rn   r.   rÄ   rD   ro   s      @@@r5   Ú_convert_to_mpt_cachez(MptPreTrainedModel._convert_to_mpt_cache  sM   ú€ ð 7EÀQÑ6GÈÑ6J×6PÑ6PÑ3ˆ
�I˜x¨Ø%/°)Ñ%;Ð"ô õ 
ð
 -ô
ó 
ð 	
r7   )r}   r~   r   r   Úconfig_classÚbase_model_prefixÚsupports_gradient_checkpointingÚ_no_split_modulesÚ_keys_to_ignore_on_load_missingr?   r   ÚModuler¿   Ústaticmethodr   r"   r�   rÇ   r‚   rƒ   s   @r5   r¬   r¬   ç   s‹   ø„ Ø€LØ%ÐØ&*Ð#Ø#˜ÐØ'4 oÐ#ô,ð* B§I¡Ió *ð" ð
Ø˜e E§L¡L°%·,±,Ð$>Ñ?Ñ@ð
à	ˆu�U—\‘\ 5§<¡<Ð/Ñ0Ñ	1ò
ó ô
r7   r¬   a*  

    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 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 ([`MptConfig`]): 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.
a®  
    Args:
        input_ids (`torch.LongTensor` of shape `(batch_size, input_ids_length)`):
            `input_ids_length` = `sequence_length` if `past_key_values` is `None` else `past_key_values[0][0].shape[2]`
            (`sequence_length` of input past key value states). Indices of input sequence tokens in the vocabulary.

            If `past_key_values` is used, only `input_ids` that do not have their past calculated should be passed as
            `input_ids`.

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

            [What are input IDs?](../glossary#input-ids)
        past_key_values (`Tuple[Tuple[torch.Tensor]]` of length `config.n_layers`):
            Contains precomputed hidden-states (key and values in the attention blocks) as computed by the model (see
            `past_key_values` output below). Can be used to speed up sequential decoding. The `input_ids` which have
            their past given to this model should not be passed as `input_ids` as they have already been computed.

            Each element of `past_key_values` is a tuple (past_key, past_value):
            - past_key: [batch_size * num_heads, head_dim, kv_length]
            - past_value: [batch_size * num_heads, kv_length, head_dim]
        attention_mask (`torch.FloatTensor` 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)

        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.

            If `past_key_values` is used, optionally only the last `inputs_embeds` have to be input (see
            `past_key_values`).
        use_cache (`bool`, *optional*):
            If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see
            `past_key_values`).
        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 [`~file_utils.ModelOutput`] instead of a plain tuple.
z]The bare Mpt Model transformer outputting raw hidden-states without any specific head on top.c                   ó’  ‡ — e Zd Zdefˆ fd„Zd„ Zdd„Zdej                  fd„Z	 e
e«       eeee¬«      	 	 	 	 	 	 	 	 ddeej"                     d	eeeej                  ej                  f   d
f      deej                     deej"                     dee   dee   dee   dee   deeej                  d
f   ef   fd„«       «       Zˆ xZS )ÚMptModelr:   c                 óô  •— t         ‰| �  |«       |j                  | _        |j                  | _        t        j                  |j                  | j                  «      | _        t        j                  t        |j                  «      D �cg c]  }t        |«      ‘Œ c}«      | _        t        | j                  |j                  ¬«      | _        d | j                   _        d| _        | j'                  «        y c c}w )Nr™   F)r>   r?   r@   rA   r.   r   r¼   Ú
vocab_sizeÚwteÚ
ModuleListÚrangeÚn_layersr—   Úblocksr	   r›   Únorm_fr=   Úgradient_checkpointingÚ	post_init)rO   r:   Ú_rP   s      €r5   r?   zMptModel.__init__\  s¶   ø€ Ü‰Ñ˜Ô à!×-Ñ-ˆÔØŸ™ˆŒô —<‘< × 1Ñ 1°4×3CÑ3CÓDˆŒô —m‘m¼uÀVÇ_Á_Ó?UÖ$V¸!¤X¨fÕ%5Ò$VÓWˆŒô   × 0Ñ 0°f×6OÑ6OÔPˆŒàˆ�‰Ôà&+ˆÔ#ð 	�‰Õùò %Ws   ÂC5c                 ó   — | j                   S r¯   ©rÓ   ©rO   s    r5   Úget_input_embeddingszMptModel.get_input_embeddingsr  s   € Ø�x‰xˆr7   c                 ó   — t        ||||«      S r¯   )r6   )rO   r.   r/   r0   r   s        r5   r6   zMptModel.build_mpt_alibi_tensoru  s   € Ü% i°À.ÐRXÓYÐYr7   Únew_embeddingsc                 ó   — || _         y r¯   rÝ   ©rO   rá   s     r5   Úset_input_embeddingszMptModel.set_input_embeddingsx  s	   € Ø!ˆ�r7   ©Ú
checkpointÚoutput_typerÈ   Ú	input_idsÚpast_key_values.rT   Úinputs_embedsr¤   r¥   Úoutput_hidden_statesÚreturn_dictr‘   c	           
      óh  — |�|n| j                   j                  }|�|n| j                   j                  }|�|n| j                   j                  }|�|n| j                   j                  }|�|�t        d«      ‚|�|j                  \  }
}n|�|j                  \  }
}}nt        d«      ‚|€"t        d gt        | j                  «      z  «      }|€| j                  |«      }|}|rdnd }|rdnd }|rdnd }| j                  r%| j                  r|rt        j                  d«       d}|}d}|d   �|d   d   j                  d   }||z   }|€$t        j                   |
|f|j"                  ¬«      }n|j%                  |j"                  «      }| j'                  | j(                  | j                   j*                  |j"                  ¬«      }t-        ||
|f||«      }|j/                  «       }t1        | j                  |«      D ]w  \  }}|r||fz   }| j                  r.| j                  r"| j3                  |j4                  ||||||«      }n |||||||¬	«      }|d   }|d
u r	||d   fz   }|sŒk|||rdnd   fz   }Œy | j7                  |«      }|r||fz   }|st        d„ ||||fD «       «      S t9        ||||¬«      S )NzDYou cannot specify both input_ids and inputs_embeds at the same timez5You have to specify either input_ids or inputs_embeds© zZ`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`...Fr   r   ©r   )r£   rT   r¤   r¥   rR   Tr   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wr¯   rî   )rÃ   Úvs     r5   rÅ   z#MptModel.forward.<locals>.<genexpr>è  s   è ø€ Òw˜qÐijÑivœÑwùs   ‚Š)Úlast_hidden_stateré   rQ   Ú
attentions)r:   r¥   rë   r¤   Úuse_return_dictre   r]   rÆ   rb   r×   rÓ   rÙ   r\   ÚloggerÚwarning_oncer"   Úonesr   rj   r6   r.   rB   r   rª   ÚzipÚ_gradient_checkpointing_funcÚ__call__rØ   r   )rO   rè   ré   rT   rê   r¤   r¥   rë   rì   r±   rn   ro   rÛ   rQ   ÚpresentsÚall_self_attentionsÚall_hidden_statesÚseq_length_with_pastÚpast_key_values_lengthr1   Úcausal_maskÚblockr£   r©   s                           r5   r|   zMptModel.forward{  s%  € ð$ 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð "+Ð!6‘I¸D¿K¹K×<QÑ<Qˆ	Ø%0Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàÐ  ]Ð%>ÜÐcÓdÐdØÐ"Ø%.§_¡_Ñ"ˆJ™
ØÐ&Ø(5×(;Ñ(;Ñ%ˆJ˜
¡AäÐTÓUÐUàÐ"Ü# T F¬S°·±Ó-=Ñ$=Ó>ˆOàÐ Ø ŸH™H YÓ/ˆMà%ˆá"‘2¨ˆÙ$5™b¸4ÐÙ"6™B¸DÐà×&Ò&¨4¯=ª=ÙÜ×#Ñ#Øpôð "�	ð  *ÐØ!"ÐØ˜1ÑÐ)Ø%4°QÑ%7¸Ñ%:×%@Ñ%@ÀÑ%CÐ"Ø#7Ð:PÑ#PÐ ØÐ!Ü"ŸZ™Z¨Ð5IÐ(JÐS`×SgÑSgÔh‰Nà+×.Ñ.¨}×/CÑ/CÓDˆNà×+Ñ+¨D¯N©N¸D¿K¹K×<SÑ<SÐ\i×\pÑ\pÐ+Óqˆä7Ø˜Z¨Ð4°mÐE[ó
ˆð "×&Ñ&Ó(ˆä!$ T§[¡[°/Ó!Bò 	^ÑˆE�:Ù#Ø$5¸Ð8HÑ$HÐ!à×*Ò*¨t¯}ª}Ø×;Ñ;Ø—N‘NØ!ØØØØØ%ó‘ñ  Ø!Ø)Ø#.Ø'Ø&7Ø"'ô�ð $ A™JˆMØ˜DÑ Ø# w¨q¡z mÑ3�â Ø&9¸WÉ)ÁQÐYZÑ=[Ð<]Ñ&]Ñ#ð;	^ð@ Ÿ™ MÓ2ˆáØ 1°]Ð4DÑ DÐáÜÑw ]°HÐ>OÐQdÐ$eÔwÓwÐwä8Ø+Ø$Ø+Ø*ô	
ð 	
r7   ©é   N©NNNNNNNN)r}   r~   r   r   r?   rß   r6   r"   r�   rä   r   ÚMPT_INPUTS_DOCSTRINGr   Ú_CHECKPOINT_FOR_DOCr   Ú_CONFIG_FOR_DOCr   Ú
LongTensorr   rª   r   r|   r‚   rƒ   s   @r5   rÐ   rÐ   W  sA  ø„ ð
˜yõ ò,óZð"°5·<±<ó "ñ +Ð+?Ó@ÙØ&Ø=Ø$ôð 15ØSWØ15Ø48Ø$(Ø,0Ø/3Ø&*ñn
à˜E×,Ñ,Ñ-ðn
ð " %¨¨e¯l©l¸E¿L¹LÐ.HÑ(IÈ3Ð(NÑ"OÑPðn
ð ! §¡Ñ.ð	n
ð
   × 0Ñ 0Ñ1ðn
ð ˜D‘>ðn
ð $ D™>ðn
ð ' t™nðn
ð ˜d‘^ðn
ð 
ˆu�U—\‘\ 3Ð&Ñ'Ð)RÐRÑ	Sòn
óó Aôn
r7   rÐ   z†
    The MPT Model transformer with a language modeling head on top (linear layer with weights tied to the input
    embeddings).
    c                   óL  ‡ — e Zd ZdgZdefˆ fd„Zd„ Zdej                  fd„Z	 e
e«       eeee¬«      	 	 	 	 	 	 	 	 	 ddeej"                     d	eeeej                  ej                  f   d
f      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ej                     ef   fd„«       «       Zdeeej                  ej                  f   d
f   dej"                  deeej                  ej                  f   d
f   fd„Zˆ xZS )ÚMptForCausalLMzlm_head.weightr:   c                 óÆ   •— t         ‰| �  |«       t        |«      | _        t	        j
                  |j                  |j                  d¬«      | _        | j                  «        y ©NFr<   )
r>   r?   rÐ   r­   r   rK   r@   rÒ   Úlm_headrÚ   rN   s     €r5   r?   zMptForCausalLM.__init__ü  sI   ø€ Ü‰Ñ˜Ô Ü# FÓ+ˆÔÜ—y‘y ×!3Ñ!3°V×5FÑ5FÈUÔSˆŒð 	�‰Õr7   c                 ó   — | j                   S r¯   ©r  rÞ   s    r5   Úget_output_embeddingsz$MptForCausalLM.get_output_embeddings  s   € Ø�|‰|Ðr7   rá   c                 ó   — || _         y r¯   r  rã   s     r5   Úset_output_embeddingsz$MptForCausalLM.set_output_embeddings  s	   € Ø%ˆ�r7   rå   rè   ré   .rT   rê   Úlabelsr¤   r¥   rë   rì   r‘   c
           
      ó¬  — |	�|	n| j                   j                  }	| j                  ||||||||	¬«      }|d   }| j                  |«      }d}|�E|j	                  |j
                  «      } | j                  ||fd| j                   j                  i|
¤Ž}|	s|f|dd z   }|�|f|z   S |S t        |||j                  |j                  |j                  ¬«      S )a³  
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for language modeling. Note that the labels **are shifted** inside the model, i.e. you can set
            `labels = input_ids` Indices are selected in `[-100, 0, ..., config.vocab_size]` All labels set to `-100`
            are ignored (masked), the loss is only computed for labels in `[0, ..., config.vocab_size]`
        N©ré   rT   rê   r¤   r¥   rë   rì   r   rÒ   r   ©ÚlossÚlogitsré   rQ   ró   )r:   rô   r­   r  rj   r   Úloss_functionrÒ   r   ré   rQ   ró   )rO   rè   ré   rT   rê   r  r¤   r¥   rë   rì   r±   Útransformer_outputsrQ   Ú	lm_logitsr  r•   s                   r5   r|   zMptForCausalLM.forward
  s  € ð2 &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà"×.Ñ.ØØ+Ø)Ø'ØØ/Ø!5Ø#ð /ó 	
Ðð ,¨AÑ.ˆà—L‘L Ó/ˆ	àˆØÐà—Y‘Y˜y×/Ñ/Ó0ˆFà%�4×%Ñ%ØØñð  Ÿ;™;×1Ñ1ðð ñ	ˆDñ Ø�\Ð$7¸¸Ð$;Ñ;ˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä0ØØØ/×?Ñ?Ø-×;Ñ;Ø*×5Ñ5ô
ð 	
r7   ÚpastÚbeam_idxc           	      ó¶   ‡— |D ��ci c]/  }|D ](  }|j                   |j                  |j                   «      “Œ* Œ1 c}}Št        ˆfd„|D «       «      }|S c c}}w )aL  
        This function is used to re-order the `past_key_values` cache if [`~PreTrainedModel.beam_search`] or
        [`~PreTrainedModel.beam_sample`] is called. This is required to match `past_key_values` with the correct
        beam_idx at every generation step.

        Output shares the same memory storage as `past`.
        c              3   ó²   •K  — | ]N  }|d    j                  d ‰|d    j                     «      |d   j                  d ‰|d    j                     «      f–— ŒP y­wrÂ   )Úindex_selectr   )rÃ   r£   Údevice_to_beam_idxs     €r5   rÅ   z0MptForCausalLM._reorder_cache.<locals>.<genexpr>Y  se   øè ø€ ò 
ð
 ð ˜1‘×*Ñ*¨1Ð.@ÀÈAÁ×AUÑAUÑ.VÓWØ˜1‘×*Ñ*¨1Ð.@ÀÈAÁ×AUÑAUÑ.VÓWôñ
ùs   ƒAA)r   rj   rÆ   )rO   r  r  r£   Ú
past_stateÚreordered_pastr!  s         @r5   Ú_reorder_cachezMptForCausalLM._reorder_cacheK  sr   ø€ ð QU÷
ØBLÐgqò
ØYcˆJ×Ñ˜xŸ{™{¨:×+<Ñ+<Ó=Ñ=ð
Øó
Ðô ó 
ð
 #ô
ó 
ˆð Ðùó
s   ‡4A©	NNNNNNNNN)r}   r~   r   Ú_tied_weights_keysr   r?   r  r"   r�   r  r   r  r   r  r   r  r   r  r   rª   r   r|   r$  r‚   rƒ   s   @r5   r
  r
  ò  s¦  ø„ ð +Ð+Ðð˜yõ òð&°E·L±Ló &ñ +Ð+?Ó@ÙØ&Ø5Ø$ôð 15ØSWØ15Ø04Ø)-Ø$(Ø,0Ø/3Ø&*ñ9
à˜E×,Ñ,Ñ-ð9
ð " %¨¨e¯l©l¸E¿L¹LÐ.HÑ(IÈ3Ð(NÑ"OÑPð9
ð ! §¡Ñ.ð	9
ð
   §¡Ñ-ð9
ð ˜Ÿ™Ñ&ð9
ð ˜D‘>ð9
ð $ D™>ð9
ð ' t™nð9
ð ˜d‘^ð9
ð 
ˆu�U—\‘\Ñ"Ð$EÐEÑ	Fò9
óó Að9
ðvØ˜% §¡¨e¯l©lÐ :Ñ;¸SÐ@ÑAðØMR×M]ÑM]ðà	ˆu�U—\‘\ 5§<¡<Ð/Ñ0°#Ð5Ñ	6÷r7   r
  aÒ  
    The MPT Model transformer with a sequence classification head on top (linear layer).

    [`MptForSequenceClassification`] uses the last token in order to do the classification, as other causal models
    (e.g. GPT-1) do.

    Since it does classification on the last token, it requires to know the position of the last token. If a
    `pad_token_id` is defined in the configuration, it finds the last token that is not a padding token in each row. If
    no `pad_token_id` is defined, it simply takes the last value in each row of the batch. Since it cannot guess the
    padding tokens when `inputs_embeds` are passed instead of `input_ids`, it does the same (take the last value in
    each row of the batch).
    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eej                  ej                  f   df      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ej                     e	f   fd„«       «       Zˆ xZS )ÚMptForSequenceClassificationr:   c                 óè   •— t         ‰| �  |«       |j                  | _        t        |«      | _        t        j                  |j                  |j                  d¬«      | _        | j                  «        y r  )
r>   r?   Ú
num_labelsrÐ   r­   r   rK   r@   ÚscorerÚ   rN   s     €r5   r?   z%MptForSequenceClassification.__init__s  sV   ø€ Ü‰Ñ˜Ô Ø ×+Ñ+ˆŒÜ# FÓ+ˆÔÜ—Y‘Y˜v×1Ñ1°6×3DÑ3DÈ5ÔQˆŒ
ð 	�‰Õr7   rå   rè   ré   .rT   rê   r  r¤   r¥   rë   rì   r‘   c
           
      ór  — |	�|	n| j                   j                  }	| j                  ||||||||	¬«      }
|
d   }| j                  |«      }|�|j                  d   }n|j                  d   }| j                   j
                  €|dk7  rt        d«      ‚| j                   j
                  €d}nÃ|�“|| j                   j
                  k7  j                  |j                  t        j                  «      }t        j                  |j                  d   |j                  t        j                  ¬«      }||z  j                  d«      }n.d}t        j                  | j                  j                   › d�«       |t        j                  ||j                  ¬	«      |f   }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/                  «       «      }nc |||«      }nY| j                   j"                  dk(  rt1        «       } |||«      }n,| j                   j"                  dk(  rt3        «       } |||«      }|	s|f|
dd z   }|�|f|z   S |S t5        |||
j6                  |
j8                  |
j:                  ¬«      S )á�  
        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   r   z=Cannot handle batch sizes > 1 if no padding token is defined.rX   )r   r   zŠ will not detect padding tokens in `inputs_embeds`. Results may be unexpected if using padding tokens in conjunction with `inputs_embeds.`rï   Ú
regressionÚsingle_label_classificationÚmulti_label_classificationr  )r:   rô   r­   r+  r]   Úpad_token_idre   rj   r   r"   r$   r#   Úargmaxrõ   rö   rP   r}   Úproblem_typer*  r   ÚlongÚintr
   r-   r   r   r   ré   rQ   ró   )rO   rè   ré   rT   rê   r  r¤   r¥   rë   rì   r  rQ   r  rn   Úlast_non_pad_tokenÚnon_pad_maskÚtoken_indicesÚpooled_logitsr  Úloss_fctr•   s                        r5   r|   z$MptForSequenceClassification.forward|  sí  € ð0 &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà"×.Ñ.ØØ+Ø)Ø'ØØ/Ø!5Ø#ð /ó 	
Ðð ,¨AÑ.ˆØ—‘˜MÓ*ˆàÐ Ø"Ÿ™¨Ñ+‰Jà&×,Ñ,¨QÑ/ˆJà�;‰;×#Ñ#Ð+°
¸a²ÜÐ\Ó]Ð]Ø�;‰;×#Ñ#Ð+Ø!#ÑØÐ"à%¨¯©×)AÑ)AÑA×EÑEÀfÇmÁmÔUZ×U`ÑU`ÓaˆLÜ!ŸL™L¨¯©¸Ñ)<ÀVÇ]Á]ÔZ_×ZeÑZeÔfˆMØ"/°,Ñ">×!FÑ!FÀrÓ!JÑà!#ÐÜ×ÑØ—>‘>×*Ñ*Ð+ð ,Zð Zôð
 œuŸ|™|¨J¸v¿}¹}ÔMÐOaÐaÑbˆàˆØÑØ�{‰{×'Ñ'Ð/Ø—?‘? 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Ò'Ù# M×$9Ñ$9Ó$;¸V¿^¹^Ó=MÓN‘Dá# M°6Ó:‘DØ—‘×)Ñ)Ð-JÒJÜ+Ó-�Ù ¨vÓ6‘Ø—‘×)Ñ)Ð-IÒIÜ,Ó.�Ù ¨vÓ6�ÙØ#Ð%Ð(;¸A¸BÐ(?Ñ?ˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä/ØØ Ø/×?Ñ?Ø-×;Ñ;Ø*×5Ñ5ô
ð 	
r7   r%  )r}   r~   r   r   r?   r   r  r   r  r   r  r   r"   r  r   r�   rª   r   r|   r‚   rƒ   s   @r5   r(  r(  c  s6  ø„ ð ˜yõ ñ +Ð+?Ó@ÙØ&Ø4Ø$ôð 15ØSWØ15Ø04Ø)-Ø$(Ø,0Ø/3Ø&*ñY
à˜E×,Ñ,Ñ-ðY
ð " %¨¨e¯l©l¸E¿L¹LÐ.HÑ(IÈ3Ð(NÑ"OÑPðY
ð ! §¡Ñ.ð	Y
ð
   §¡Ñ-ðY
ð ˜Ÿ™Ñ&ðY
ð ˜D‘>ðY
ð $ D™>ðY
ð ' t™nðY
ð ˜d‘^ðY
ð 
ˆu�U—\‘\Ñ"Ð$DÐDÑ	EòY
óó AôY
r7   r(  z¢
    MPT Model with a token classification head on top (a linear layer on top of the hidden-states output) 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eej                  ej                  f   df      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ej                     e	f   fd„«       «       Zˆ xZS )ÚMptForTokenClassificationr:   c                 ó°  •— t         ‰| �  |«       |j                  | _        t        |«      | _        t        |d«      r|j                  �|j                  }n't        |d«      r|j                  �|j                  }nd}t        j                  |«      | _
        t        j                  |j                  |j                  «      | _        | j                  «        y )NÚclassifier_dropoutrŽ   gš™™™™™¹?)r>   r?   r*  rÐ   r­   Úhasattrr>  rŽ   r   r¡   rk   rK   r@   Ú
classifierrÚ   )rO   r:   r>  rP   s      €r5   r?   z"MptForTokenClassification.__init__æ  s¯   ø€ Ü‰Ñ˜Ô Ø ×+Ñ+ˆŒä# FÓ+ˆÔÜ�6Ð/Ô0°V×5NÑ5NÐ5ZØ!'×!:Ñ!:ÑÜ�VÐ-Ô.°6×3HÑ3HÐ3TØ!'×!6Ñ!6Ñà!$ÐÜ—z‘zÐ"4Ó5ˆŒÜŸ)™) F×$6Ñ$6¸×8IÑ8IÓJˆŒð 	�‰Õr7   rå   rè   ré   .rT   rê   r  r¤   r¥   rë   rì   r‘   c
           
      ó  — |	�|	n| j                   j                  }	| j                  ||||||||	¬«      }|d   }| j                  |«      }| j	                  |«      }d}|�l|j                  |j                  «      }|j                  \  }}t        «       } ||j                  ||z  | j                  «      |j                  ||z  «      «      }|	s|f|dd z   }|�|f|z   S |S t        |||j                  |j                  ¬«      S )r-  Nr  r   r   )r  r  rQ   ró   )r:   rô   r­   rk   r@  rj   r   r]   r   r%   r*  r   rQ   ró   )rO   rè   ré   rT   rê   r  r¤   r¥   rë   rì   Údeprecated_argumentsr  rQ   r  r  rn   ro   r:  r•   s                      r5   r|   z!MptForTokenClassification.forward÷  s+  € ð2 &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà"×.Ñ.ØØ+Ø)Ø'ØØ/Ø!5Ø#ð /ó 	
Ðð ,¨AÑ.ˆØŸ™ ]Ó3ˆØ—‘ Ó/ˆàˆØÐà—Y‘Y˜vŸ}™}Ó-ˆFØ%+§\¡\Ñ"ˆJ˜
Ü'Ó)ˆHÙØ—‘˜J¨Ñ3°T·_±_ÓEÀvÇ{Á{ÐS]Ð`jÑSjÓGkóˆDñ Ø�YÐ!4°Q°RÐ!8Ñ8ˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä$ØØØ-×;Ñ;Ø*×5Ñ5ô	
ð 	
r7   r%  )r}   r~   r   r   r?   r   r  r   r  r   r  r   r"   r  r   r�   rª   r   r|   r‚   rƒ   s   @r5   r<  r<  Þ  s*  ø„ ð˜yõ ñ" +Ð+?Ó@ÙØ&Ø)Ø$ôð 15ØSWØ15Ø04Ø)-Ø$(Ø,0Ø/3Ø&*ñ7
à˜E×,Ñ,Ñ-ð7
ð " %¨¨e¯l©l¸E¿L¹LÐ.HÑ(IÈ3Ð(NÑ"OÑPð7
ð ! §¡Ñ.ð	7
ð
   §¡Ñ-ð7
ð ˜Ÿ™Ñ&ð7
ð ˜D‘>ð7
ð $ D™>ð7
ð ' t™nð7
ð ˜d‘^ð7
ð 
ˆu�U—\‘\Ñ"Ð$9Ð9Ñ	:ò7
óó Aô7
r7   r<  zì
    The MPT Model transformer with a span classification head on top for extractive question-answering tasks like SQuAD
    (a linear layers on top of the hidden-states output to compute `span start logits` and `span end logits`).
    c                   ó.  ‡ — e Zd Zˆ fd„Z eej                  d«      «      	 	 	 	 	 	 	 	 d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ef   fd„«       Zˆ xZS )ÚMptForQuestionAnsweringc                 ó®   •— t         ‰| �  |«       t        |«      | _        t	        j
                  |j                  d«      | _        | j                  «        y )Nr   )	r>   r?   rÐ   r­   r   rK   r@   Ú
qa_outputsrÚ   rN   s     €r5   r?   z MptForQuestionAnswering.__init__?  sA   ø€ Ü‰Ñ˜Ô Ü# FÓ+ˆÔÜŸ)™) F×$6Ñ$6¸Ó:ˆŒð 	�‰Õr7   zbatch_size, sequence_lengthrè   rT   rê   Ústart_positionsÚend_positionsr¥   rë   rì   r‘   c	                 ó"  — |�|n| j                   j                  }| j                  ||||||¬«      }	|	d   }
| j                  |
«      }|j	                  dd¬«      \  }}|j                  d«      j                  «       }|j                  d«      j                  «       }d}|�·|�µt        |j                  «       «      dkD  r|j                  d«      }t        |j                  «       «      dkD  r|j                  d«      }|j                  d«      }|j                  d|«      }|j                  d|«      }t        |¬«      } |||«      } |||«      }||z   dz  }|s||f|	dd z   }|�|f|z   S |S t        ||||	j                  |	j                  ¬	«      S )
a  
        start_positions (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for position (index) of the start of the labelled span for computing the token classification loss.
            Positions are clamped to the length of the sequence (`sequence_length`). Position outside of the sequence
            are not taken into account for computing the loss.
        end_positions (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for position (index) of the end of the labelled span for computing the token classification loss.
            Positions are clamped to the length of the sequence (`sequence_length`). Position outside of the sequence
            are not taken into account for computing the loss.
        N)rT   rê   r¥   rë   rì   r   r   rX   r    )Úignore_indexr   )r  Ústart_logitsÚ
end_logitsrQ   ró   )r:   rô   r­   rF  Úsplitr-   rm   rb   rf   r^   r   r   rQ   ró   )rO   rè   rT   rê   rG  rH  r¥   rë   rì   r©   Úsequence_outputr  rK  rL  Ú
total_lossÚignored_indexr:  Ú
start_lossÚend_lossr•   s                       r5   r|   zMptForQuestionAnswering.forwardG  s»  € ð, &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà×"Ñ"ØØ)Ø'Ø/Ø!5Ø#ð #ó 
ˆð " !™*ˆà—‘ Ó1ˆØ#)§<¡<°°r <Ó#:Ñ ˆ�jØ#×+Ñ+¨BÓ/×:Ñ:Ó<ˆØ×'Ñ'¨Ó+×6Ñ6Ó8ˆ
àˆ
ØÐ&¨=Ð+Dä�?×'Ñ'Ó)Ó*¨QÒ.Ø"1×"9Ñ"9¸"Ó"=�Ü�=×%Ñ%Ó'Ó(¨1Ò,Ø -× 5Ñ 5°bÓ 9�à(×-Ñ-¨aÓ0ˆMØ-×3Ñ3°A°}ÓEˆOØ)×/Ñ/°°=ÓAˆMä'°]ÔCˆHÙ! ,°Ó@ˆJÙ 
¨MÓ:ˆHØ$ xÑ/°1Ñ4ˆJáØ" JÐ/°'¸!¸"°+Ñ=ˆFØ/9Ð/E�Z�M FÑ*ÐQÈ6ÐQä+ØØ%Ø!Ø!×/Ñ/Ø×)Ñ)ô
ð 	
r7   r  )r}   r~   r   r?   r   r  Úformatr   r"   r  ÚFloatTensorrª   r   r   r   r|   r‚   rƒ   s   @r5   rD  rD  7  sú   ø„ ôñ +Ð+?×+FÑ+FÐGdÓ+eÓfð 15Ø6:Ø59Ø6:Ø48Ø,0Ø/3Ø&*ñB
à˜E×,Ñ,Ñ-ðB
ð ! ×!2Ñ!2Ñ3ðB
ð   × 1Ñ 1Ñ2ð	B
ð
 " %×"2Ñ"2Ñ3ðB
ð   × 0Ñ 0Ñ1ðB
ð $ D™>ðB
ð ' t™nðB
ð ˜d‘^ðB
ð 
ˆuÐ2Ð2Ñ	3òB
ó gôB
r7   rD  )r
  rÐ   r¬   r(  r<  rD  r  )7r€   r&   Útypingr   r   r   r"   Útorch.utils.checkpointr   Útorch.nnr   r   r	   r
   r   r“   Ú
file_utilsr   r   r   Ú
generationr   Úmodeling_attn_mask_utilsr   Úmodeling_outputsr   r   r   r   r   Úmodeling_utilsr   Úutilsr   Úconfiguration_mptr   Ú
get_loggerr}   rõ   r  r  r6   rÍ   r9   r…   r—   r¬   ÚMPT_START_DOCSTRINGr  rÐ   r
  r(  r<  rD  Ú__all__rî   r7   r5   ú<module>rb     s¬  ðñ ã ß )Ñ )ã Û Ý ß LÓ LÝ $ç qÑ qÝ )Ý I÷õ õ .Ý Ý (ð 
ˆ×	Ñ	˜HÓ	%€à'Ð Ø€óô.I9�2—9‘9ô I9ôXˆR�Y‰Yô ô*=ˆr�y‰yô =ô@,
˜ô ,
ð^Ð ð/Ð ñd ØcØóôT
Ð!ó T
ó	ðT
ñn ðð óôgÐ'¨ó góðgñT ðð óôi
Ð#5ó i
óði
ñX ðð óôO
Ð 2ó O
óðO
ñd ðð óôL
Ð0ó L
óðL
ò^�r7   