§
    ‚Štjùp ã                   ó8  — d Z ddlZddlmZ ddlZddlmZ ddlmZmZm	Z	 ddl
mZ ddlmZmZ dd	lmZ dd
lmZ ddlmZmZ ddlmZ ddlmZ ddlmZmZmZ ddlm Z   ej!        e"¦  «        Z# ed¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z$ ed¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z% ed¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z& ed¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z' ed¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z( ed ¬¦  «        e G d!„ d"e¦  «        ¦   «         ¦   «         Z) ed#¬¦  «        e G d$„ d%e¦  «        ¦   «         ¦   «         Z* ed&¬¦  «        e G d'„ d(e¦  «        ¦   «         ¦   «         Z+ ed)¬¦  «        e G d*„ d+e¦  «        ¦   «         ¦   «         Z, ed,¬¦  «        e G d-„ d.e¦  «        ¦   «         ¦   «         Z- G d/„ d0ej.        ¦  «        Z/ G d1„ d2ej.        ¦  «        Z0 G d3„ d4ej.        ¦  «        Z1 G d5„ d6ej.        ¦  «        Z2 G d7„ d8ej.        ¦  «        Z3 G d9„ d:ej.        ¦  «        Z4 G d;„ d<ej.        ¦  «        Z5 G d=„ d>e¦  «        Z6 G d?„ d@ej.        ¦  «        Z7 G dA„ dBej.        ¦  «        Z8 G dC„ dDej.        ¦  «        Z9 G dE„ dFej.        ¦  «        Z:e G dG„ dHe¦  «        ¦   «         Z; edI¬¦  «         G dJ„ dKe;¦  «        ¦   «         Z<dL„ Z= G dM„ dNej.        ¦  «        Z> edO¬¦  «         G dP„ dQe;¦  «        ¦   «         Z? edR¬¦  «         G dS„ dTe;¦  «        ¦   «         Z@ edU¬¦  «         G dV„ dWe;¦  «        ¦   «         ZA edX¬¦  «         G dY„ dZe;¦  «        ¦   «         ZB ed[¬¦  «         G d\„ d]e;¦  «        ¦   «         ZC ed^¬¦  «         G d_„ d`e;¦  «        ¦   «         ZDe G da„ dbe;¦  «        ¦   «         ZEe G dc„ dde;¦  «        ¦   «         ZFg de¢ZGdS )fzPyTorch LUKE model.é    N)Ú	dataclass)Únn)ÚBCEWithLogitsLossÚCrossEntropyLossÚMSELossé   )Úinitialization)ÚACT2FNÚgelu)Úcreate_bidirectional_mask)ÚGradientCheckpointingLayer)ÚBaseModelOutputÚBaseModelOutputWithPooling)ÚPreTrainedModel)Úapply_chunking_to_forward)ÚModelOutputÚauto_docstringÚloggingé   )Ú
LukeConfigz3
    Base class for outputs of the LUKE model.
    )Úcustom_introc                   ó`   — e Zd ZU dZdZej        dz  ed<   dZe	ej        df         dz  ed<   dS )ÚBaseLukeModelOutputWithPoolingax  
    pooler_output (`torch.FloatTensor` of shape `(batch_size, hidden_size)`):
        Last layer hidden-state of the first token of the sequence (classification token) further processed by a
        Linear layer and a Tanh activation function.
    entity_last_hidden_state (`torch.FloatTensor` of shape `(batch_size, entity_length, hidden_size)`):
        Sequence of entity hidden-states at the output of the last layer of the model.
    entity_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
        Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
        shape `(batch_size, entity_length, hidden_size)`. Entity hidden-states of the model at the output of each
        layer plus the initial entity embedding outputs.
    NÚentity_last_hidden_state.Úentity_hidden_states©
Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚtorchÚFloatTensorÚ__annotations__r   Útuple© ó    úd/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/luke/modeling_luke.pyr   r   %   sZ   € € € € € € ð
ð 
ð :>Ð˜eÔ/°$Ñ6Ð=Ð=Ñ=ØAEÐ˜% Ô 1°3Ð 6Ô7¸$Ñ>ÐEÐEÑEÐEÐEr&   r   zV
    Base class for model's outputs, with potential hidden states and attentions.
    c                   ó`   — e Zd ZU dZdZej        dz  ed<   dZe	ej        df         dz  ed<   dS )ÚBaseLukeModelOutputa„  
    entity_last_hidden_state (`torch.FloatTensor` of shape `(batch_size, entity_length, hidden_size)`):
        Sequence of entity hidden-states at the output of the last layer of the model.
    entity_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
        Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
        shape `(batch_size, entity_length, hidden_size)`. Entity hidden-states of the model at the output of each
        layer plus the initial entity embedding outputs.
    Nr   .r   r   r%   r&   r'   r)   r)   <   sZ   € € € € € € ðð ð :>Ð˜eÔ/°$Ñ6Ð=Ð=Ñ=ØAEÐ˜% Ô 1°3Ð 6Ô7¸$Ñ>ÐEÐEÑEÐEÐEr&   r)   c                   ó0  — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	ej        dz  ed<   dZ
ej        dz  ed<   dZej        dz  ed<   dZeej                 dz  ed<   dZeej        d	f         dz  ed
<   dZeej        d	f         dz  ed<   dS )ÚLukeMaskedLMOutputa:  
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        The sum of masked language modeling (MLM) loss and entity prediction loss.
    mlm_loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Masked language modeling (MLM) loss.
    mep_loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Masked entity prediction (MEP) loss.
    logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):
        Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax).
    entity_logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):
        Prediction scores of the entity prediction head (scores for each entity vocabulary token before SoftMax).
    entity_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
        Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
        shape `(batch_size, entity_length, hidden_size)`. Entity hidden-states of the model at the output of each
        layer plus the initial entity embedding outputs.
    NÚlossÚmlm_lossÚmep_lossÚlogitsÚentity_logitsÚhidden_states.r   Ú
attentions)r   r   r   r    r,   r!   r"   r#   r-   r.   r/   r0   r1   r$   r   r2   r%   r&   r'   r+   r+   P   sø   € € € € € € ðð ð" &*€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø)-€HˆeÔ $Ñ&Ð-Ð-Ñ-Ø)-€HˆeÔ $Ñ&Ð-Ð-Ñ-Ø'+€FˆEÔ Ñ$Ð+Ð+Ñ+Ø.2€M�5Ô$ tÑ+Ð2Ð2Ñ2Ø59€M�5˜Ô*Ô+¨dÑ2Ð9Ð9Ñ9ØAEÐ˜% Ô 1°3Ð 6Ô7¸$Ñ>ÐEÐEÑEØ7;€J��eÔ'¨Ð,Ô-°Ñ4Ð;Ð;Ñ;Ð;Ð;r&   r+   z2
    Outputs of entity classification models.
    c                   óÚ   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	e
ej        df         dz  ed<   dZe
ej        df         dz  ed<   dZe
ej        df         dz  ed<   dS )	ÚEntityClassificationOutputá¿  
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Classification loss.
    logits (`torch.FloatTensor` of shape `(batch_size, config.num_labels)`):
        Classification scores (before SoftMax).
    entity_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
        Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
        shape `(batch_size, entity_length, hidden_size)`. Entity hidden-states of the model at the output of each
        layer plus the initial entity embedding outputs.
    Nr,   r/   .r1   r   r2   ©r   r   r   r    r,   r!   r"   r#   r/   r1   r$   r   r2   r%   r&   r'   r4   r4   r   óµ   € € € € € € ð	ð 	ð &*€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø'+€FˆEÔ Ñ$Ð+Ð+Ñ+Ø:>€M�5˜Ô*¨CÐ/Ô0°4Ñ7Ð>Ð>Ñ>ØAEÐ˜% Ô 1°3Ð 6Ô7¸$Ñ>ÐEÐEÑEØ7;€J��eÔ'¨Ð,Ô-°Ñ4Ð;Ð;Ñ;Ð;Ð;r&   r4   z7
    Outputs of entity pair classification models.
    c                   óÚ   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	e
ej        df         dz  ed<   dZe
ej        df         dz  ed<   dZe
ej        df         dz  ed<   dS )	ÚEntityPairClassificationOutputr5   Nr,   r/   .r1   r   r2   r6   r%   r&   r'   r9   r9   ‹   r7   r&   r9   z7
    Outputs of entity span classification models.
    c                   óÚ   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	e
ej        df         dz  ed<   dZe
ej        df         dz  ed<   dZe
ej        df         dz  ed<   dS )	ÚEntitySpanClassificationOutputaÎ  
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Classification loss.
    logits (`torch.FloatTensor` of shape `(batch_size, entity_length, config.num_labels)`):
        Classification scores (before SoftMax).
    entity_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
        Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
        shape `(batch_size, entity_length, hidden_size)`. Entity hidden-states of the model at the output of each
        layer plus the initial entity embedding outputs.
    Nr,   r/   .r1   r   r2   r6   r%   r&   r'   r;   r;   ¤   r7   r&   r;   z4
    Outputs of sentence classification models.
    c                   óÚ   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	e
ej        df         dz  ed<   dZe
ej        df         dz  ed<   dZe
ej        df         dz  ed<   dS )	ÚLukeSequenceClassifierOutputa  
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Classification (or regression if config.num_labels==1) loss.
    logits (`torch.FloatTensor` of shape `(batch_size, config.num_labels)`):
        Classification (or regression if config.num_labels==1) scores (before SoftMax).
    entity_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
        Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
        shape `(batch_size, entity_length, hidden_size)`. Entity hidden-states of the model at the output of each
        layer plus the initial entity embedding outputs.
    Nr,   r/   .r1   r   r2   r6   r%   r&   r'   r=   r=   ½   r7   r&   r=   z@
    Base class for outputs of token classification models.
    c                   óÚ   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	e
ej        df         dz  ed<   dZe
ej        df         dz  ed<   dZe
ej        df         dz  ed<   dS )	ÚLukeTokenClassifierOutputaÐ  
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Classification loss.
    logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.num_labels)`):
        Classification scores (before SoftMax).
    entity_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
        Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
        shape `(batch_size, entity_length, hidden_size)`. Entity hidden-states of the model at the output of each
        layer plus the initial entity embedding outputs.
    Nr,   r/   .r1   r   r2   r6   r%   r&   r'   r?   r?   Ö   r7   r&   r?   z/
    Outputs of question answering models.
    c                   óø   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	ej        dz  ed<   dZ
eej        df         dz  ed<   dZeej        df         dz  ed<   dZeej        df         dz  ed	<   dS )
Ú LukeQuestionAnsweringModelOutputa‡  
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Total span extraction loss is the sum of a Cross-Entropy for the start and end positions.
    entity_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
        Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
        shape `(batch_size, entity_length, hidden_size)`. Entity hidden-states of the model at the output of each
        layer plus the initial entity embedding outputs.
    Nr,   Ústart_logitsÚ
end_logits.r1   r   r2   )r   r   r   r    r,   r!   r"   r#   rB   rC   r1   r$   r   r2   r%   r&   r'   rA   rA   ï   sÍ   € € € € € € ðð ð &*€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø-1€L�%Ô# dÑ*Ð1Ð1Ñ1Ø+/€J�Ô! DÑ(Ð/Ð/Ñ/Ø:>€M�5˜Ô*¨CÐ/Ô0°4Ñ7Ð>Ð>Ñ>ØAEÐ˜% Ô 1°3Ð 6Ô7¸$Ñ>ÐEÐEÑEØ7;€J��eÔ'¨Ð,Ô-°Ñ4Ð;Ð;Ñ;Ð;Ð;r&   rA   z,
    Outputs of multiple choice models.
    c                   óÚ   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	e
ej        df         dz  ed<   dZe
ej        df         dz  ed<   dZe
ej        df         dz  ed<   dS )	ÚLukeMultipleChoiceModelOutputa  
    loss (`torch.FloatTensor` of shape *(1,)*, *optional*, returned when `labels` is provided):
        Classification loss.
    logits (`torch.FloatTensor` of shape `(batch_size, num_choices)`):
        *num_choices* is the second dimension of the input tensors. (see *input_ids* above).

        Classification scores (before SoftMax).
    entity_hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
        Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
        shape `(batch_size, entity_length, hidden_size)`. Entity hidden-states of the model at the output of each
        layer plus the initial entity embedding outputs.
    Nr,   r/   .r1   r   r2   r6   r%   r&   r'   rE   rE     sµ   € € € € € € ðð ð &*€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø'+€FˆEÔ Ñ$Ð+Ð+Ñ+Ø:>€M�5˜Ô*¨CÐ/Ô0°4Ñ7Ð>Ð>Ñ>ØAEÐ˜% Ô 1°3Ð 6Ô7¸$Ñ>ÐEÐEÑEØ7;€J��eÔ'¨Ð,Ô-°Ñ4Ð;Ð;Ñ;Ð;Ð;r&   rE   c                   ó8   ‡ — e Zd ZdZˆ fd„Z	 	 	 	 dd„Zd„ Zˆ xZS )ÚLukeEmbeddingszV
    Same as BertEmbeddings with a tiny tweak for positional embeddings indexing.
    c                 ó"  •— t          ¦   «                              ¦   «          t          j        |j        |j        |j        ¬¦  «        | _        t          j        |j        |j        ¦  «        | _	        t          j        |j
        |j        ¦  «        | _        t          j        |j        |j        ¬¦  «        | _        t          j        |j        ¦  «        | _        |j        | _        t          j        |j        |j        | j        ¬¦  «        | _	        d S )N©Úpadding_idx©Úeps)ÚsuperÚ__init__r   Ú	EmbeddingÚ
vocab_sizeÚhidden_sizeÚpad_token_idÚword_embeddingsÚmax_position_embeddingsÚposition_embeddingsÚtype_vocab_sizeÚtoken_type_embeddingsÚ	LayerNormÚlayer_norm_epsÚDropoutÚhidden_dropout_probÚdropoutrJ   ©ÚselfÚconfigÚ	__class__s     €r'   rN   zLukeEmbeddings.__init__'  sÝ   ø€ Ý‰Œ×ÒÑÔÐÝ!œ|¨FÔ,=¸vÔ?QÐ_eÔ_rÐsÑsÔsˆÔÝ#%¤<°Ô0NÐPVÔPbÑ#cÔ#cˆÔ Ý%'¤\°&Ô2HÈ&ÔJ\Ñ%]Ô%]ˆÔ"åœ fÔ&8¸fÔ>SÐTÑTÔTˆŒÝ”z &Ô"<Ñ=Ô=ˆŒð "Ô.ˆÔÝ#%¤<ØÔ*¨FÔ,>ÈDÔL\ð$
ñ $
ô $
ˆÔ Ð Ð r&   Nc                 ó:  — |€E|�.t          || j        ¦  «                             |j        ¦  «        }n|                      |¦  «        }|�|                     ¦   «         }n|                     ¦   «         d d…         }|€+t          j        |t          j        | j	        j        ¬¦  «        }|€|  
                    |¦  «        }|                      |¦  «        }|                      |¦  «        }||z   |z   }|                      |¦  «        }|                      |¦  «        }|S )Néÿÿÿÿ©ÚdtypeÚdevice)Ú"create_position_ids_from_input_idsrJ   Útore   Ú&create_position_ids_from_inputs_embedsÚsizer!   ÚzerosÚlongÚposition_idsrS   rU   rW   rX   r\   )	r^   Ú	input_idsÚtoken_type_idsrl   Úinputs_embedsÚinput_shaperU   rW   Ú
embeddingss	            r'   ÚforwardzLukeEmbeddings.forward6  s  € ð ÐØÐ$åAÀ)ÈTÔM]Ñ^Ô^×aÒaÐbkÔbrÑsÔs��à#×JÒJÈ=ÑYÔY�àÐ Ø#Ÿ.š.Ñ*Ô*ˆKˆKà'×,Ò,Ñ.Ô.¨s°¨sÔ3ˆKàÐ!Ý"œ[¨½E¼JÈtÔO`ÔOgÐhÑhÔhˆNàÐ Ø ×0Ò0°Ñ;Ô;ˆMà"×6Ò6°|ÑDÔDÐØ $× :Ò :¸>Ñ JÔ JÐà"Ð%8Ñ8Ð;PÑPˆ
Ø—^’^ JÑ/Ô/ˆ
Ø—\’\ *Ñ-Ô-ˆ
ØÐr&   c                 ó  — |                      ¦   «         dd…         }|d         }t          j        | j        dz   || j        z   dz   t          j        |j        ¬¦  «        }|                     d¦  «                             |¦  «        S )z×
        We are provided embeddings directly. We cannot infer which are padded so just generate sequential position ids.

        Args:
            inputs_embeds: torch.Tensor

        Returns: torch.Tensor
        Nrb   r   rc   r   )ri   r!   ÚarangerJ   rk   re   Ú	unsqueezeÚexpand)r^   ro   rp   Úsequence_lengthrl   s        r'   rh   z5LukeEmbeddings.create_position_ids_from_inputs_embedsW  s‡   € ð $×(Ò(Ñ*Ô*¨3¨B¨3Ô/ˆØ% aœ.ˆå”|ØÔ˜qÑ  /°DÔ4DÑ"DÀqÑ"HÕPUÔPZÐcpÔcwð
ñ 
ô 
ˆð ×%Ò% aÑ(Ô(×/Ò/°Ñ<Ô<Ð<r&   )NNNN)r   r   r   r    rN   rr   rh   Ú__classcell__©r`   s   @r'   rG   rG   "  st   ø€ € € € € ðð ð
ð 
ð 
ð 
ð 
ð" ØØØðð ð ð ðB=ð =ð =ð =ð =ð =ð =r&   rG   c                   ó`   ‡ — e Zd Zdefˆ fd„Z	 ddej        dej        dej        dz  fd„Zˆ xZS )	ÚLukeEntityEmbeddingsr_   c                 ó$  •— t          ¦   «                              ¦   «          || _        t          j        |j        |j        d¬¦  «        | _        |j        |j        k    r&t          j	        |j        |j        d¬¦  «        | _
        t          j        |j        |j        ¦  «        | _        t          j        |j        |j        ¦  «        | _        t          j        |j        |j        ¬¦  «        | _        t          j        |j        ¦  «        | _        d S )Nr   rI   F©ÚbiasrK   )rM   rN   r_   r   rO   Úentity_vocab_sizeÚentity_emb_sizeÚentity_embeddingsrQ   ÚLinearÚentity_embedding_denserT   rU   rV   rW   rX   rY   rZ   r[   r\   r]   s     €r'   rN   zLukeEntityEmbeddings.__init__j  sÚ   ø€ Ý‰Œ×ÒÑÔÐØˆŒå!#¤¨fÔ.FÈÔH^ÐlmÐ!nÑ!nÔ!nˆÔØÔ! VÔ%7Ò7Ð7Ý*,¬)°FÔ4JÈFÔL^ÐejÐ*kÑ*kÔ*kˆDÔ'å#%¤<°Ô0NÐPVÔPbÑ#cÔ#cˆÔ Ý%'¤\°&Ô2HÈ&ÔJ\Ñ%]Ô%]ˆÔ"åœ fÔ&8¸fÔ>SÐTÑTÔTˆŒÝ”z &Ô"<Ñ=Ô=ˆŒˆˆr&   NÚ
entity_idsrl   rn   c                 ó‚  — |€t          j        |¦  «        }|                      |¦  «        }| j        j        | j        j        k    r|                      |¦  «        }|                      |                     d¬¦  «        ¦  «        }|dk     	                    |¦  «         
                    d¦  «        }||z  }t          j        |d¬¦  «        }||                     d¬¦  «                             d¬¦  «        z  }|                      |¦  «        }||z   |z   }|                      |¦  «        }|                      |¦  «        }|S )Nr   )Úminrb   éþÿÿÿ©ÚdimgH¯¼šò×z>)r!   Ú
zeros_liker�   r_   r€   rQ   rƒ   rU   ÚclampÚtype_asru   ÚsumrW   rX   r\   )	r^   r„   rl   rn   r�   rU   Úposition_embedding_maskrW   rq   s	            r'   rr   zLukeEntityEmbeddings.forwardx  sE  € ð Ð!Ý"Ô-¨jÑ9Ô9ˆNà ×2Ò2°:Ñ>Ô>ÐØŒ;Ô&¨$¬+Ô*AÒAÐAØ $× ;Ò ;Ð<MÑ NÔ NÐà"×6Ò6°|×7IÒ7IÈaÐ7IÑ7PÔ7PÑQÔQÐØ#/°2Ò#5×">Ò">Ð?RÑ"SÔ"S×"]Ò"]Ð^`Ñ"aÔ"aÐØ1Ð4KÑKÐÝ#œiÐ(;ÀÐDÑDÔDÐØ1Ð4K×4OÒ4OÐTVÐ4OÑ4WÔ4W×4]Ò4]ÐbfÐ4]Ñ4gÔ4gÑgÐà $× :Ò :¸>Ñ JÔ JÐà&Ð)<Ñ<Ð?TÑTˆ
Ø—^’^ JÑ/Ô/ˆ
Ø—\’\ *Ñ-Ô-ˆ
àÐr&   ©N)	r   r   r   r   rN   r!   Ú
LongTensorrr   rx   ry   s   @r'   r{   r{   i  sŒ   ø€ € € € € ð>˜zð >ð >ð >ð >ð >ð >ð$ 37ð	ð àÔ$ðð Ô&ðð Ô(¨4Ñ/ð	ð ð ð ð ð ð ð r&   r{   c                   ó0   ‡ — e Zd Zˆ fd„Zd„ Z	 	 dd„Zˆ xZS )ÚLukeSelfAttentionc                 ób  •— t          ¦   «                              ¦   «          |j        |j        z  dk    r0t	          |d¦  «        s t          d|j        › d|j        › d�¦  «        ‚|j        | _        t          |j        |j        z  ¦  «        | _        | j        | j        z  | _        |j	        | _	        t          j        |j        | j        ¦  «        | _        t          j        |j        | j        ¦  «        | _        t          j        |j        | j        ¦  «        | _        | j	        rlt          j        |j        | j        ¦  «        | _        t          j        |j        | j        ¦  «        | _        t          j        |j        | j        ¦  «        | _        t          j        |j        ¦  «        | _        d S )Nr   Úembedding_sizezThe hidden size z4 is not a multiple of the number of attention heads ú.)rM   rN   rQ   Únum_attention_headsÚhasattrÚ
ValueErrorÚintÚattention_head_sizeÚall_head_sizeÚuse_entity_aware_attentionr   r‚   ÚqueryÚkeyÚvalueÚ	w2e_queryÚ	e2w_queryÚ	e2e_queryrZ   Úattention_probs_dropout_probr\   r]   s     €r'   rN   zLukeSelfAttention.__init__•  s{  ø€ Ý‰Œ×ÒÑÔÐØÔ Ô :Ñ:¸aÒ?Ð?ÍÐPVÐXhÑHiÔHiÐ?Ýð7 6Ô#5ð 7ð 7ØÔ3ð7ð 7ð 7ñô ð ð
 $*Ô#=ˆÔ Ý#& vÔ'9¸FÔ<VÑ'VÑ#WÔ#WˆÔ Ø!Ô5¸Ô8PÑPˆÔØ*0Ô*KˆÔ'å”Y˜vÔ1°4Ô3EÑFÔFˆŒ
Ý”9˜VÔ/°Ô1CÑDÔDˆŒÝ”Y˜vÔ1°4Ô3EÑFÔFˆŒ
àÔ*ð 	OÝœY vÔ'9¸4Ô;MÑNÔNˆDŒNÝœY vÔ'9¸4Ô;MÑNÔNˆDŒNÝœY vÔ'9¸4Ô;MÑNÔNˆDŒNå”z &Ô"EÑFÔFˆŒˆˆr&   c                 óœ   — |                      ¦   «         d d…         | j        | j        fz   } |j        |Ž }|                     dddd¦  «        S )Nrb   r   é   r   r   )ri   r–   rš   ÚviewÚpermute)r^   ÚxÚnew_x_shapes      r'   Útranspose_for_scoresz&LukeSelfAttention.transpose_for_scores­  sM   € Ø—f’f‘h”h˜s ˜s”m tÔ'?ÀÔAYÐ&ZÑZˆØˆAŒF�KÐ ˆØ�yŠy˜˜A˜q !Ñ$Ô$Ð$r&   NFc                 óp  — |                      d¦  «        }|€|}nt          j        ||gd¬¦  «        }|                      |                      |¦  «        ¦  «        }|                      |                      |¦  «        ¦  «        }| j        �rà|��Ý|                      |                      |¦  «        ¦  «        }	|                      |                      |¦  «        ¦  «        }
|                      |  	                    |¦  «        ¦  «        }|                      |  
                    |¦  «        ¦  «        }|d d …d d …d |…d d …f         }|d d …d d …d |…d d …f         }|d d …d d …|d …d d …f         }|d d …d d …|d …d d …f         }t          j        |	|                     dd¦  «        ¦  «        }t          j        |
|                     dd¦  «        ¦  «        }t          j        ||                     dd¦  «        ¦  «        }t          j        ||                     dd¦  «        ¦  «        }t          j        ||gd¬¦  «        }t          j        ||gd¬¦  «        }t          j        ||gd¬¦  «        }nQ|                      |                      |¦  «        ¦  «        }t          j        ||                     dd¦  «        ¦  «        }|t          j        | j        ¦  «        z  }|�||z   }t           j                             |d¬¦  «        }|                      |¦  «        }t          j        ||¦  «        }|                     dddd¦  «                             ¦   «         }|                      ¦   «         d d…         | j        fz   } |j        |Ž }|d d …d |…d d …f         }|€d }n|d d …|d …d d …f         }|r|||f}n||f}|S )Nr   rˆ   rb   r‡   r   r¥   r   )ri   r!   Úcatrª   rž   rŸ   rœ   r�   r    r¡   r¢   ÚmatmulÚ	transposeÚmathÚsqrtrš   r   Ú
functionalÚsoftmaxr\   r§   Ú
contiguousr›   r¦   )r^   Úword_hidden_statesr   Úattention_maskÚoutput_attentionsÚ	word_sizeÚconcat_hidden_statesÚ	key_layerÚvalue_layerÚw2w_query_layerÚw2e_query_layerÚe2w_query_layerÚe2e_query_layerÚw2w_key_layerÚe2w_key_layerÚw2e_key_layerÚe2e_key_layerÚw2w_attention_scoresÚw2e_attention_scoresÚe2w_attention_scoresÚe2e_attention_scoresÚword_attention_scoresÚentity_attention_scoresÚattention_scoresÚquery_layerÚattention_probsÚcontext_layerÚnew_context_layer_shapeÚoutput_word_hidden_statesÚoutput_entity_hidden_statesÚoutputss                                  r'   rr   zLukeSelfAttention.forward²  s  € ð '×+Ò+¨AÑ.Ô.ˆ	àÐ'Ø#5Ð Ð å#(¤9Ð.@ÐBVÐ-WÐ]^Ð#_Ñ#_Ô#_Ð à×-Ò-¨d¯hªhÐ7KÑ.LÔ.LÑMÔMˆ	Ø×/Ò/°·
²
Ð;OÑ0PÔ0PÑQÔQˆàÔ*ñ 	VÐ/CÑ/Oð #×7Ò7¸¿
º
ÐCUÑ8VÔ8VÑWÔWˆOØ"×7Ò7¸¿ºÐGYÑ8ZÔ8ZÑ[Ô[ˆOØ"×7Ò7¸¿ºÐG[Ñ8\Ô8\Ñ]Ô]ˆOØ"×7Ò7¸¿ºÐG[Ñ8\Ô8\Ñ]Ô]ˆOð & a a a¨¨¨¨J¨Y¨J¸¸¸Ð&9Ô:ˆMØ% a a a¨¨¨¨J¨Y¨J¸¸¸Ð&9Ô:ˆMØ% a a a¨¨¨¨I¨J¨J¸¸¸Ð&9Ô:ˆMØ% a a a¨¨¨¨I¨J¨J¸¸¸Ð&9Ô:ˆMõ $)¤<°À×AXÒAXÐY[Ð]_ÑA`ÔA`Ñ#aÔ#aÐ Ý#(¤<°À×AXÒAXÐY[Ð]_ÑA`ÔA`Ñ#aÔ#aÐ Ý#(¤<°À×AXÒAXÐY[Ð]_ÑA`ÔA`Ñ#aÔ#aÐ Ý#(¤<°À×AXÒAXÐY[Ð]_ÑA`ÔA`Ñ#aÔ#aÐ õ %*¤IÐ/CÐEYÐ.ZÐ`aÐ$bÑ$bÔ$bÐ!Ý&+¤iÐ1EÐG[Ð0\ÐbcÐ&dÑ&dÔ&dÐ#Ý$œyÐ*?ÐAXÐ)YÐ_`ÐaÑaÔaÐÐð ×3Ò3°D·J²JÐ?SÑ4TÔ4TÑUÔUˆKÝ$œ|¨K¸×9LÒ9LÈRÐQSÑ9TÔ9TÑUÔUÐà+­d¬i¸Ô8PÑ.QÔ.QÑQÐØÐ%à/°.Ñ@Ðõ œ-×/Ò/Ð0@ÀbÐ/ÑIÔIˆð Ÿ,š, Ñ7Ô7ˆåœ _°kÑBÔBˆà%×-Ò-¨a°°A°qÑ9Ô9×DÒDÑFÔFˆØ"/×"4Ò"4Ñ"6Ô"6°s¸°sÔ";¸tÔ?QÐ>SÑ"SÐØ*˜Ô*Ð,CÐDˆà$1°!°!°!°Z°i°ZÀÀÀÐ2BÔ$CÐ!ØÐ'Ø*.Ð'Ð'à*7¸¸¸¸9¸:¸:ÀqÀqÀqÐ8HÔ*IÐ'àð 	OØ0Ð2MÈÐ_ˆGˆGà0Ð2MÐNˆGàˆr&   ©NF)r   r   r   rN   rª   rr   rx   ry   s   @r'   r’   r’   ”  sp   ø€ € € € € ðGð Gð Gð Gð Gð0%ð %ð %ð ØðKð Kð Kð Kð Kð Kð Kð Kr&   r’   c                   óP   ‡ — e Zd Zˆ fd„Zdej        dej        dej        fd„Zˆ xZS )ÚLukeSelfOutputc                 ó  •— t          ¦   «                              ¦   «          t          j        |j        |j        ¦  «        | _        t          j        |j        |j        ¬¦  «        | _        t          j        |j	        ¦  «        | _
        d S ©NrK   )rM   rN   r   r‚   rQ   ÚdenserX   rY   rZ   r[   r\   r]   s     €r'   rN   zLukeSelfOutput.__init__  sf   ø€ Ý‰Œ×ÒÑÔÐÝ”Y˜vÔ1°6Ô3EÑFÔFˆŒ
Ýœ fÔ&8¸fÔ>SÐTÑTÔTˆŒÝ”z &Ô"<Ñ=Ô=ˆŒˆˆr&   r1   Úinput_tensorÚreturnc                 óŠ   — |                       |¦  «        }|                      |¦  «        }|                      ||z   ¦  «        }|S r�   ©rÖ   r\   rX   ©r^   r1   r×   s      r'   rr   zLukeSelfOutput.forward  ó@   € ØŸ
š
 =Ñ1Ô1ˆØŸš ]Ñ3Ô3ˆØŸš }°|Ñ'CÑDÔDˆØÐr&   ©r   r   r   rN   r!   ÚTensorrr   rx   ry   s   @r'   rÓ   rÓ     ói   ø€ € € € € ð>ð >ð >ð >ð >ð U¤\ð ÀÄð ÐRWÔR^ð ð ð ð ð ð ð ð r&   rÓ   c                   ó*   ‡ — e Zd Zˆ fd„Z	 	 dd„Zˆ xZS )ÚLukeAttentionc                 ó˜   •— t          ¦   «                              ¦   «          t          |¦  «        | _        t	          |¦  «        | _        d S r�   )rM   rN   r’   r^   rÓ   Úoutputr]   s     €r'   rN   zLukeAttention.__init__  s;   ø€ Ý‰Œ×ÒÑÔÐÝ% fÑ-Ô-ˆŒ	Ý$ VÑ,Ô,ˆŒˆˆr&   NFc                 ó~  — |                      d¦  «        }|                      ||||¦  «        }|€|d         }|}n6t          j        |d d…         d¬¦  «        }t          j        ||gd¬¦  «        }|                      ||¦  «        }	|	d d …d |…d d …f         }
|€d }n|	d d …|d …d d …f         }|
|f|dd …         z   }|S )Nr   r   r¥   rˆ   )ri   r^   r!   r¬   rã   )r^   r´   r   rµ   r¶   r·   Úself_outputsÚconcat_self_outputsr¸   Úattention_outputÚword_attention_outputÚentity_attention_outputrÐ   s                r'   rr   zLukeAttention.forward  s  € ð '×+Ò+¨AÑ.Ô.ˆ	Ø—y’yØØ ØØñ	
ô 
ˆð  Ð'Ø".¨q¤/ÐØ#5Ð Ð å"'¤)¨L¸¸!¸Ô,<À!Ð"DÑ"DÔ"DÐÝ#(¤9Ð.@ÐBVÐ-WÐ]^Ð#_Ñ#_Ô#_Ð àŸ;š;Ð':Ð<PÑQÔQÐà 0°°°°J°Y°JÀÀÀÐ1AÔ BÐØÐ'Ø&*Ð#Ð#à&6°q°q°q¸)¸*¸*ÀaÀaÀaÐ7GÔ&HÐ#ð )Ð*AÐBÀ\ÐRSÐRTÐRTÔEUÑUˆàˆr&   rÑ   ©r   r   r   rN   rr   rx   ry   s   @r'   rá   rá     sT   ø€ € € € € ð-ð -ð -ð -ð -ð Øð ð  ð  ð  ð  ð  ð  ð  r&   rá   c                   óB   ‡ — e Zd Zˆ fd„Zdej        dej        fd„Zˆ xZS )ÚLukeIntermediatec                 ó  •— t          ¦   «                              ¦   «          t          j        |j        |j        ¦  «        | _        t          |j        t          ¦  «        rt          |j                 | _        d S |j        | _        d S r�   )rM   rN   r   r‚   rQ   Úintermediate_sizerÖ   Ú
isinstanceÚ
hidden_actÚstrr
   Úintermediate_act_fnr]   s     €r'   rN   zLukeIntermediate.__init__:  sn   ø€ Ý‰Œ×ÒÑÔÐÝ”Y˜vÔ1°6Ô3KÑLÔLˆŒ
Ý�fÔ'­Ñ-Ô-ð 	9Ý'-¨fÔ.?Ô'@ˆDÔ$Ð$Ð$à'-Ô'8ˆDÔ$Ð$Ð$r&   r1   rØ   c                 óZ   — |                       |¦  «        }|                      |¦  «        }|S r�   )rÖ   rò   ©r^   r1   s     r'   rr   zLukeIntermediate.forwardB  s,   € ØŸ
š
 =Ñ1Ô1ˆØ×0Ò0°Ñ?Ô?ˆØÐr&   rÝ   ry   s   @r'   rì   rì   9  s^   ø€ € € € € ð9ð 9ð 9ð 9ð 9ð U¤\ð °e´lð ð ð ð ð ð ð ð r&   rì   c                   óP   ‡ — e Zd Zˆ fd„Zdej        dej        dej        fd„Zˆ xZS )Ú
LukeOutputc                 ó  •— t          ¦   «                              ¦   «          t          j        |j        |j        ¦  «        | _        t          j        |j        |j        ¬¦  «        | _        t          j	        |j
        ¦  «        | _        d S rÕ   )rM   rN   r   r‚   rî   rQ   rÖ   rX   rY   rZ   r[   r\   r]   s     €r'   rN   zLukeOutput.__init__J  sf   ø€ Ý‰Œ×ÒÑÔÐÝ”Y˜vÔ7¸Ô9KÑLÔLˆŒ
Ýœ fÔ&8¸fÔ>SÐTÑTÔTˆŒÝ”z &Ô"<Ñ=Ô=ˆŒˆˆr&   r1   r×   rØ   c                 óŠ   — |                       |¦  «        }|                      |¦  «        }|                      ||z   ¦  «        }|S r�   rÚ   rÛ   s      r'   rr   zLukeOutput.forwardP  rÜ   r&   rÝ   ry   s   @r'   rö   rö   I  rß   r&   rö   c                   ó0   ‡ — e Zd Zˆ fd„Z	 	 dd„Zd„ Zˆ xZS )Ú	LukeLayerc                 óæ   •— t          ¦   «                              ¦   «          |j        | _        d| _        t	          |¦  «        | _        t          |¦  «        | _        t          |¦  «        | _	        d S ©Nr   )
rM   rN   Úchunk_size_feed_forwardÚseq_len_dimrá   Ú	attentionrì   Úintermediaterö   rã   r]   s     €r'   rN   zLukeLayer.__init__X  s^   ø€ Ý‰Œ×ÒÑÔÐØ'-Ô'EˆÔ$ØˆÔÝ& vÑ.Ô.ˆŒÝ,¨VÑ4Ô4ˆÔÝ  Ñ(Ô(ˆŒˆˆr&   NFc                 óf  — |                      d¦  «        }|                      ||||¬¦  «        }|€	|d         }nt          j        |d d…         d¬¦  «        }|dd …         }t	          | j        | j        | j        |¦  «        }	|	d d …d |…d d …f         }
|€d }n|	d d …|d …d d …f         }|
|f|z   }|S )Nr   )r¶   r   r¥   rˆ   )ri   rÿ   r!   r¬   r   Úfeed_forward_chunkrý   rþ   )r^   r´   r   rµ   r¶   r·   Úself_attention_outputsÚconcat_attention_outputrÐ   Úlayer_outputÚword_layer_outputÚentity_layer_outputs               r'   rr   zLukeLayer.forward`  s  € ð '×+Ò+¨AÑ.Ô.ˆ	à!%§¢ØØ ØØ/ð	 "0ñ "
ô "
Ðð  Ð'Ø&<¸QÔ&?Ð#Ð#å&+¤iÐ0FÀrÈÀrÔ0JÐPQÐ&RÑ&RÔ&RÐ#à(¨¨¨Ô,ˆå0ØÔ# TÔ%AÀ4ÔCSÐUlñ
ô 
ˆð )¨¨¨¨J¨Y¨J¸¸¸Ð)9Ô:ÐØÐ'Ø"&ÐÐà".¨q¨q¨q°)°*°*¸a¸a¸aÐ/?Ô"@Ðà$Ð&9Ð:¸WÑDˆàˆr&   c                 ó\   — |                       |¦  «        }|                      ||¦  «        }|S r�   )r   rã   )r^   rç   Úintermediate_outputr  s       r'   r  zLukeLayer.feed_forward_chunkƒ  s2   € Ø"×/Ò/Ð0@ÑAÔAÐØ—{’{Ð#6Ð8HÑIÔIˆØÐr&   rÑ   )r   r   r   rN   rr   r  rx   ry   s   @r'   rú   rú   W  sd   ø€ € € € € ð)ð )ð )ð )ð )ð Øð!ð !ð !ð !ðFð ð ð ð ð ð r&   rú   c                   ó.   ‡ — e Zd Zˆ fd„Z	 	 	 	 dd„Zˆ xZS )ÚLukeEncoderc                 óÔ   •‡— t          ¦   «                              ¦   «          ‰| _        t          j        ˆfd„t          ‰j        ¦  «        D ¦   «         ¦  «        | _        d| _        d S )Nc                 ó.   •— g | ]}t          ‰¦  «        ‘ŒS r%   )rú   )Ú.0Ú_r_   s     €r'   ú
<listcomp>z(LukeEncoder.__init__.<locals>.<listcomp>�  s!   ø€ Ð#_Ð#_Ð#_¸!¥I¨fÑ$5Ô$5Ð#_Ð#_Ð#_r&   F)	rM   rN   r_   r   Ú
ModuleListÚrangeÚnum_hidden_layersÚlayerÚgradient_checkpointingr]   s    `€r'   rN   zLukeEncoder.__init__Š  s`   øø€ Ý‰Œ×ÒÑÔÐØˆŒÝ”]Ð#_Ð#_Ð#_Ð#_½uÀVÔE]Ñ?^Ô?^Ð#_Ñ#_Ô#_Ñ`Ô`ˆŒ
Ø&+ˆÔ#Ð#Ð#r&   NFTc                 óV  — |rdnd }|rdnd }|rdnd }	t          | j        ¦  «        D ]A\  }
}|r||fz   }||fz   } |||||¦  «        }|d         }|�|d         }|r|	|d         fz   }	ŒB|r||fz   }||fz   }|st          d„ |||	||fD ¦   «         ¦  «        S t          |||	||¬¦  «        S )Nr%   r   r   r¥   c              3   ó   K  — | ]}|®|V — Œ	d S r�   r%   ©r  Úvs     r'   ú	<genexpr>z&LukeEncoder.forward.<locals>.<genexpr>¶  ó4   è è € ð 
ð 
àð �=ð ð !�=�=�=ð
ð 
r&   )Úlast_hidden_stater1   r2   r   r   )Ú	enumerater  r$   r)   )r^   r´   r   rµ   r¶   Úoutput_hidden_statesÚreturn_dictÚall_word_hidden_statesÚall_entity_hidden_statesÚall_self_attentionsÚiÚlayer_moduleÚlayer_outputss                r'   rr   zLukeEncoder.forward�  sg  € ð (<Ð!E  ÀÐØ)=Ð#G 2 2À4Ð Ø$5Ð?˜b˜b¸4Ðå(¨¬Ñ4Ô4ð 	Pð 	P‰OˆAˆ|Ø#ð ^Ø)?ÐCUÐBWÑ)WÐ&Ø+CÐG[ÐF]Ñ+]Ð(à(˜LØ"Ø$ØØ!ñ	ô ˆMð "/¨qÔ!1Ðà#Ð/Ø'4°QÔ'7Ð$à ð PØ&9¸]È1Ô=MÐ<OÑ&OÐ#øàð 	ZØ%;Ð?QÐ>SÑ%SÐ"Ø'?ÐCWÐBYÑ'YÐ$àð 	Ýð 
ð 
ð 'Ø*Ø'Ø(Ø,ðð
ñ 
ô 
ñ 
ô 
ð 
õ #Ø0Ø0Ø*Ø%9Ø!9ð
ñ 
ô 
ð 	
r&   )NFFTrê   ry   s   @r'   r  r  ‰  sZ   ø€ € € € € ð,ð ,ð ,ð ,ð ,ð ØØ"Øð7
ð 7
ð 7
ð 7
ð 7
ð 7
ð 7
ð 7
r&   r  c                   óB   ‡ — e Zd Zˆ fd„Zdej        dej        fd„Zˆ xZS )Ú
LukePoolerc                 óÀ   •— t          ¦   «                              ¦   «          t          j        |j        |j        ¦  «        | _        t          j        ¦   «         | _        d S r�   )rM   rN   r   r‚   rQ   rÖ   ÚTanhÚ
activationr]   s     €r'   rN   zLukePooler.__init__Ì  sC   ø€ Ý‰Œ×ÒÑÔÐÝ”Y˜vÔ1°6Ô3EÑFÔFˆŒ
Ýœ'™)œ)ˆŒˆˆr&   r1   rØ   c                 ór   — |d d …df         }|                       |¦  «        }|                      |¦  «        }|S )Nr   )rÖ   r*  )r^   r1   Úfirst_token_tensorÚpooled_outputs       r'   rr   zLukePooler.forwardÑ  s@   € ð +¨1¨1¨1¨a¨4Ô0ÐØŸ
š
Ð#5Ñ6Ô6ˆØŸš¨Ñ6Ô6ˆØÐr&   rÝ   ry   s   @r'   r'  r'  Ë  s^   ø€ € € € € ð$ð $ð $ð $ð $ð
 U¤\ð °e´lð ð ð ð ð ð ð ð r&   r'  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚEntityPredictionHeadTransformc                 óV  •— t          ¦   «                              ¦   «          t          j        |j        |j        ¦  «        | _        t          |j        t          ¦  «        rt          |j                 | _        n|j        | _        t          j        |j        |j        ¬¦  «        | _        d S rÕ   )rM   rN   r   r‚   rQ   r€   rÖ   rï   rð   rñ   r
   Útransform_act_fnrX   rY   r]   s     €r'   rN   z&EntityPredictionHeadTransform.__init__Û  s…   ø€ Ý‰Œ×ÒÑÔÐÝ”Y˜vÔ1°6Ô3IÑJÔJˆŒ
Ý�fÔ'­Ñ-Ô-ð 	6Ý$*¨6Ô+<Ô$=ˆDÔ!Ð!à$*Ô$5ˆDÔ!Ýœ fÔ&<À&ÔBWÐXÑXÔXˆŒˆˆr&   c                 ó„   — |                       |¦  «        }|                      |¦  «        }|                      |¦  «        }|S r�   )rÖ   r1  rX   rô   s     r'   rr   z%EntityPredictionHeadTransform.forwardä  s=   € ØŸ
š
 =Ñ1Ô1ˆØ×-Ò-¨mÑ<Ô<ˆØŸš }Ñ5Ô5ˆØÐr&   rê   ry   s   @r'   r/  r/  Ú  sL   ø€ € € € € ðYð Yð Yð Yð Yðð ð ð ð ð ð r&   r/  c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚEntityPredictionHeadc                 ó*  •— t          ¦   «                              ¦   «          || _        t          |¦  «        | _        t          j        |j        |j        d¬¦  «        | _	        t          j
        t          j        |j        ¦  «        ¦  «        | _        d S )NFr}   )rM   rN   r_   r/  Ú	transformr   r‚   r€   r   ÚdecoderÚ	Parameterr!   rj   r~   r]   s     €r'   rN   zEntityPredictionHead.__init__ì  sp   ø€ Ý‰Œ×ÒÑÔÐØˆŒÝ6°vÑ>Ô>ˆŒÝ”y Ô!7¸Ô9QÐX]Ð^Ñ^Ô^ˆŒÝ”L¥¤¨VÔ-EÑ!FÔ!FÑGÔGˆŒ	ˆ	ˆ	r&   c                 ój   — |                       |¦  «        }|                      |¦  «        | j        z   }|S r�   )r6  r7  r~   rô   s     r'   rr   zEntityPredictionHead.forwardó  s1   € ØŸš }Ñ5Ô5ˆØŸš ]Ñ3Ô3°d´iÑ?ˆàÐr&   rê   ry   s   @r'   r4  r4  ë  sL   ø€ € € € € ðHð Hð Hð Hð Hðð ð ð ð ð ð r&   r4  c                   ót   ‡ — e Zd ZU eed<   dZdZddgZ ej	        ¦   «         de
j        fˆ fd„¦   «         Zˆ xZS )ÚLukePreTrainedModelr_   ÚlukeTrá   r{   Úmodulec                 ó¢  •— t          ¦   «                              |¦  «         t          |t          j        ¦  «        rŽ|j        dk    rt          j        |j        ¦  «         n&t          j	        |j        d| j
        j        ¬¦  «         |j        �>t          |j        dd¦  «        s*t          j        |j        |j                 ¦  «         dS dS dS dS )zInitialize the weightsr   g        )ÚmeanÚstdNÚ_is_hf_initializedF)rM   Ú_init_weightsrï   r   rO   Úembedding_dimÚinitÚzeros_ÚweightÚnormal_r_   Úinitializer_rangerJ   Úgetattr)r^   r=  r`   s     €r'   rB  z!LukePreTrainedModel._init_weights  sÈ   ø€ õ 	‰Œ×Ò˜fÑ%Ô%Ð%Ý�f�bœlÑ+Ô+ð 	?ØÔ# qÒ(Ð(Ý”˜FœMÑ*Ô*Ð*Ð*å”˜Vœ]°¸$¼+Ô:WÐXÑXÔXÐXàÔ!Ð-µg¸f¼mÐMaÐchÑ6iÔ6iÐ-Ý”˜FœM¨&Ô*<Ô=Ñ>Ô>Ð>Ð>Ð>ð	?ð 	?ð .Ð-Ð-Ð-r&   )r   r   r   r   r#   Úbase_model_prefixÚsupports_gradient_checkpointingÚ_no_split_modulesr!   Úno_gradr   ÚModulerB  rx   ry   s   @r'   r;  r;  ú  s~   ø€ € € € € € àÐÐÑØÐØ&*Ð#Ø(Ð*@ÐAÐà€U„]�_„_ð
? B¤Ið 
?ð 
?ð 
?ð 
?ð 
?ñ „_ð
?ð 
?ð 
?ð 
?ð 
?r&   r;  zt
    The bare LUKE model transformer outputting raw hidden-states for both word tokens and entities without any
    c                   óP  ‡ — e Zd Zddedefˆ fd„Zd„ Zd„ Zd„ Zd„ Z	e
	 	 	 	 	 	 	 	 	 	 	 	 dd
ej        d	z  dej        d	z  dej        d	z  dej        d	z  dej        d	z  dej        d	z  dej        d	z  dej        d	z  dej        d	z  ded	z  ded	z  ded	z  deez  fd„¦   «         Zˆ xZS )Ú	LukeModelTr_   Úadd_pooling_layerc                 ó(  •— t          ¦   «                              |¦  «         || _        t          |¦  «        | _        t          |¦  «        | _        t          |¦  «        | _        |rt          |¦  «        nd| _
        |                      ¦   «          dS )zv
        add_pooling_layer (bool, *optional*, defaults to `True`):
            Whether to add a pooling layer
        N)rM   rN   r_   rG   rq   r{   r�   r  Úencoderr'  ÚpoolerÚ	post_init)r^   r_   rQ  r`   s      €r'   rN   zLukeModel.__init__  sƒ   ø€ õ
 	‰Œ×Ò˜Ñ Ô Ð ØˆŒå(¨Ñ0Ô0ˆŒÝ!5°fÑ!=Ô!=ˆÔÝ" 6Ñ*Ô*ˆŒà,=ÐG•j Ñ(Ô(Ð(À4ˆŒð 	�ŠÑÔÐÐÐr&   c                 ó   — | j         j        S r�   ©rq   rS   ©r^   s    r'   Úget_input_embeddingszLukeModel.get_input_embeddings&  s   € ØŒÔ.Ð.r&   c                 ó   — || j         _        d S r�   rW  ©r^   rŸ   s     r'   Úset_input_embeddingszLukeModel.set_input_embeddings)  s   € Ø*/ˆŒÔ'Ð'Ð'r&   c                 ó   — | j         j         S r�   ©r�   rX  s    r'   Úget_entity_embeddingszLukeModel.get_entity_embeddings,  s   € ØÔ%Ô7Ð7r&   c                 ó   — || j         _         d S r�   r^  r[  s     r'   Úset_entity_embeddingszLukeModel.set_entity_embeddings/  s   € Ø38ˆÔÔ0Ð0Ð0r&   Nrm   rµ   rn   rl   r„   Úentity_attention_maskÚentity_token_type_idsÚentity_position_idsro   r¶   r  r  rØ   c                 óŒ  — |
�|
n| j         j        }
|�|n| j         j        }|�|n| j         j        }|�|	�t	          d¦  «        ‚|�+|                      ||¦  «         |                     ¦   «         }n.|	�|	                     ¦   «         dd…         }nt	          d¦  «        ‚|\  }}|�|j        n|	j        }|€t          j	        ||f|¬¦  «        }|€!t          j
        |t          j        |¬¦  «        }|�T|                     d¦  «        }|€t          j	        ||f|¬¦  «        }|€#t          j
        ||ft          j        |¬¦  «        }|                      ||||	¬¦  «        }|�t          j        ||gd¬	¦  «        }t          | j         |d
                              |j        ¦  «        |¬¦  «        }|€d}n|                      |||¦  «        }|                      ||||
||¬¦  «        }|d         }| j        �|                      |¦  «        nd}|s||f|dd…         z   S t)          |||j        |j        |j        |j        ¬¦  «        S )uz  
        entity_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`):
            Indices of entity tokens in the entity vocabulary.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.
        entity_attention_mask (`torch.FloatTensor` of shape `(batch_size, entity_length)`, *optional*):
            Mask to avoid performing attention on padding entity token indices. Mask values selected in `[0, 1]`:

            - 1 for entity tokens that are **not masked**,
            - 0 for entity tokens that are **masked**.
        entity_token_type_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`, *optional*):
            Segment token indices to indicate first and second portions of the entity token inputs. Indices are
            selected in `[0, 1]`:

            - 0 corresponds to a *portion A* entity token,
            - 1 corresponds to a *portion B* entity token.
        entity_position_ids (`torch.LongTensor` of shape `(batch_size, entity_length, max_mention_length)`, *optional*):
            Indices of positions of each input entity in the position embeddings. Selected in the range `[0,
            config.max_position_embeddings - 1]`.

        Examples:

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

        >>> tokenizer = AutoTokenizer.from_pretrained("studio-ousia/luke-base")
        >>> model = LukeModel.from_pretrained("studio-ousia/luke-base")
        # Compute the contextualized entity representation corresponding to the entity mention "BeyoncÃ©"

        >>> text = "BeyoncÃ© lives in Los Angeles."
        >>> entity_spans = [(0, 7)]  # character-based entity span corresponding to "BeyoncÃ©"

        >>> encoding = tokenizer(text, entity_spans=entity_spans, add_prefix_space=True, return_tensors="pt")
        >>> outputs = model(**encoding)
        >>> word_last_hidden_state = outputs.last_hidden_state
        >>> entity_last_hidden_state = outputs.entity_last_hidden_state
        # Input Wikipedia entities to obtain enriched contextualized representations of word tokens

        >>> text = "BeyoncÃ© lives in Los Angeles."
        >>> entities = [
        ...     "BeyoncÃ©",
        ...     "Los Angeles",
        ... ]  # Wikipedia entity titles corresponding to the entity mentions "BeyoncÃ©" and "Los Angeles"
        >>> entity_spans = [
        ...     (0, 7),
        ...     (17, 28),
        ... ]  # character-based entity spans corresponding to "BeyoncÃ©" and "Los Angeles"

        >>> encoding = tokenizer(
        ...     text, entities=entities, entity_spans=entity_spans, add_prefix_space=True, return_tensors="pt"
        ... )
        >>> outputs = model(**encoding)
        >>> word_last_hidden_state = outputs.last_hidden_state
        >>> entity_last_hidden_state = outputs.entity_last_hidden_state
        ```NzDYou cannot specify both input_ids and inputs_embeds at the same timerb   z5You have to specify either input_ids or inputs_embeds)re   rc   r   )rm   rl   rn   ro   rˆ   ).N)r_   ro   rµ   )rµ   r¶   r  r  r   )r  Úpooler_outputr1   r2   r   r   )r_   r¶   r  r  r˜   Ú%warn_if_padding_and_no_attention_maskri   re   r!   Úonesrj   rk   rq   r¬   r   rg   rd   r�   rS  rT  r   r1   r2   r   r   )r^   rm   rµ   rn   rl   r„   rb  rc  rd  ro   r¶   r  r  Úkwargsrp   Ú
batch_sizeÚ
seq_lengthre   Úentity_seq_lengthÚword_embedding_outputÚentity_embedding_outputÚencoder_outputsÚsequence_outputr-  s                           r'   rr   zLukeModel.forward2  sÛ  € ðR 2CÐ1NÐ-Ð-ÐTXÔT_ÔTqÐà$8Ð$DÐ Ð È$Ì+ÔJjð 	ð &1Ð%<�k�kÀ$Ä+ÔBYˆàÐ  ]Ð%>ÝÐcÑdÔdÐdØÐ"Ø×6Ò6°yÀ.ÑQÔQÐQØ#Ÿ.š.Ñ*Ô*ˆKˆKØÐ&Ø'×,Ò,Ñ.Ô.¨s°¨sÔ3ˆKˆKåÐTÑUÔUÐUà!,Ñˆ
�JØ%.Ð%:�Ô!Ð!ÀÔ@TˆàÐ!Ý"œZ¨°ZÐ(@ÈÐPÑPÔPˆNØÐ!Ý"œ[¨½E¼JÈvÐVÑVÔVˆNØÐ!Ø *§¢°Ñ 2Ô 2ÐØ$Ð,Ý(-¬
°JÐ@QÐ3RÐ[aÐ(bÑ(bÔ(bÐ%Ø$Ð,Ý(-¬°ZÐARÐ4SÕ[`Ô[eÐntÐ(uÑ(uÔ(uÐ%ð !%§¢ØØ%Ø)Ø'ð	 !0ñ !
ô !
Ðð !Ð,Ý"œY¨Ð8MÐ'NÐTVÐWÑWÔWˆNå2Ø”;à(¨Ô3×6Ò6Ð7LÔ7RÑSÔSØ)ð	
ñ 
ô 
ˆð ÐØ&*Ð#Ð#à&*×&<Ò&<¸ZÐI\Ð^sÑ&tÔ&tÐ#ð Ÿ,š,Ø!Ø#Ø)Ø/Ø!5Ø#ð 'ñ 
ô 
ˆð *¨!Ô,ˆð 9=¼Ð8O˜Ÿš OÑ4Ô4Ð4ÐUYˆàð 	JØ# ]Ð3°oÀaÀbÀbÔ6IÑIÐIå-Ø-Ø'Ø)Ô7Ø&Ô1Ø%4Ô%MØ!0Ô!Eð
ñ 
ô 
ð 	
r&   )T)NNNNNNNNNNNN)r   r   r   r   ÚboolrN   rY  r\  r_  ra  r   r!   r�   r"   r$   r   rr   rx   ry   s   @r'   rP  rP    sÇ  ø€ € € € € ðð ˜zð ¸dð ð ð ð ð ð ð"/ð /ð /ð0ð 0ð 0ð8ð 8ð 8ð9ð 9ð 9ð ð .2Ø37Ø26Ø04Ø.2Ø:>Ø9=Ø7;Ø26Ø)-Ø,0Ø#'ðX
ð X
àÔ# dÑ*ðX
ð Ô)¨DÑ0ðX
ð Ô(¨4Ñ/ð	X
ð
 Ô&¨Ñ-ðX
ð Ô$ tÑ+ðX
ð  %Ô0°4Ñ7ðX
ð  %Ô/°$Ñ6ðX
ð #Ô-°Ñ4ðX
ð Ô(¨4Ñ/ðX
ð   $™;ðX
ð # T™kðX
ð ˜D‘[ðX
ð 
Ð/Ñ	/ðX
ð X
ð X
ñ „^ðX
ð X
ð X
ð X
ð X
r&   rP  c                 óÖ   — |                       |¦  «                             ¦   «         }t          j        |d¬¦  «                             |¦  «        |z  }|                     ¦   «         |z   S )a  
    Replace non-padding symbols with their position numbers. Position numbers begin at padding_idx+1. Padding symbols
    are ignored. This is modified from fairseq's `utils.make_positions`.

    Args:
        x: torch.Tensor x:

    Returns: torch.Tensor
    r   rˆ   )Úner™   r!   ÚcumsumrŒ   rk   )rm   rJ   ÚmaskÚincremental_indicess       r'   rf   rf   Î  s`   € ð �<Š<˜Ñ$Ô$×(Ò(Ñ*Ô*€DÝ œ<¨°!Ð4Ñ4Ô4×<Ò<¸TÑBÔBÀdÑJÐØ×#Ò#Ñ%Ô%¨Ñ3Ð3r&   c                   ó(   ‡ — e Zd ZdZˆ fd„Zd„ Zˆ xZS )Ú
LukeLMHeadz*Roberta Head for masked language modeling.c                 ó‚  •— t          ¦   «                              ¦   «          t          j        |j        |j        ¦  «        | _        t          j        |j        |j        ¬¦  «        | _        t          j        |j        |j	        ¦  «        | _
        t          j        t          j        |j	        ¦  «        ¦  «        | _        d S rÕ   )rM   rN   r   r‚   rQ   rÖ   rX   rY   Ú
layer_normrP   r7  r8  r!   rj   r~   r]   s     €r'   rN   zLukeLMHead.__init__â  s‰   ø€ Ý‰Œ×ÒÑÔÐÝ”Y˜vÔ1°6Ô3EÑFÔFˆŒ
Ýœ, vÔ'9¸vÔ?TÐUÑUÔUˆŒå”y Ô!3°VÔ5FÑGÔGˆŒÝ”L¥¤¨VÔ->Ñ!?Ô!?Ñ@Ô@ˆŒ	ˆ	ˆ	r&   c                 ó¢   — |                       |¦  «        }t          |¦  «        }|                      |¦  «        }|                      |¦  «        }|S r�   )rÖ   r   rz  r7  )r^   Úfeaturesri  r¨   s       r'   rr   zLukeLMHead.forwardê  sE   € Ø�JŠJ�xÑ Ô ˆÝ�‰GŒGˆØ�OŠO˜AÑÔˆð �LŠL˜‰OŒOˆàˆr&   )r   r   r   r    rN   rr   rx   ry   s   @r'   rx  rx  ß  sR   ø€ € € € € Ø4Ð4ðAð Að Að Að Aðð ð ð ð ð ð r&   rx  z—
    The LUKE model with a language modeling head and entity prediction head on top for masked language modeling and
    masked entity prediction.
    c            !       ón  ‡ — e Zd ZdddœZˆ fd„Zd„ Zd„ Ze	 	 	 	 	 	 	 	 	 	 	 	 	 	 ddej	        dz  d	ej
        dz  d
ej	        dz  dej	        dz  dej	        dz  dej	        dz  dej	        dz  dej	        dz  dej	        dz  dej	        dz  dej
        dz  dedz  dedz  dedz  deez  fd„¦   «         Zˆ xZS )ÚLukeForMaskedLMz/luke.entity_embeddings.entity_embeddings.weightzlm_head.decoder.bias)z!entity_predictions.decoder.weightzlm_head.biasc                 ó  •— t          ¦   «                              |¦  «         t          |¦  «        | _        t	          |¦  «        | _        t          |¦  «        | _        t          j	        ¦   «         | _
        |                      ¦   «          d S r�   )rM   rN   rP  r<  rx  Úlm_headr4  Úentity_predictionsr   r   Úloss_fnrU  r]   s     €r'   rN   zLukeForMaskedLM.__init__  sq   ø€ Ý‰Œ×Ò˜Ñ Ô Ð å˜fÑ%Ô%ˆŒ	å! &Ñ)Ô)ˆŒÝ"6°vÑ">Ô">ˆÔåÔ*Ñ,Ô,ˆŒð 	�ŠÑÔÐÐÐr&   c                 ó   — | j         j        S r�   ©r€  r7  rX  s    r'   Úget_output_embeddingsz%LukeForMaskedLM.get_output_embeddings  s   € ØŒ|Ô#Ð#r&   c                 ó   — || j         _        d S r�   r„  )r^   Únew_embeddingss     r'   Úset_output_embeddingsz%LukeForMaskedLM.set_output_embeddings  s   € Ø-ˆŒÔÐÐr&   Nrm   rµ   rn   rl   r„   rb  rc  rd  ÚlabelsÚentity_labelsro   r¶   r  r  rØ   c                 ó  — |�|n| j         j        }|                      |||||||||||d¬¦  «        }d}d}|                      |j        ¦  «        }|	�e|	                     |j        ¦  «        }	|                      |                     d| j         j	        ¦  «        |	                     d¦  «        ¦  «        }|€|}d}d}|j
        �m|                      |j
        ¦  «        }|
�Q|                      |                     d| j         j        ¦  «        |
                     d¦  «        ¦  «        }|€|}n||z   }|s0t          d„ ||||||j        |j        |j        fD ¦   «         ¦  «        S t#          ||||||j        |j        |j        ¬¦  «        S )aC  
        entity_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`):
            Indices of entity tokens in the entity vocabulary.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.
        entity_attention_mask (`torch.FloatTensor` of shape `(batch_size, entity_length)`, *optional*):
            Mask to avoid performing attention on padding entity token indices. Mask values selected in `[0, 1]`:

            - 1 for entity tokens that are **not masked**,
            - 0 for entity tokens that are **masked**.
        entity_token_type_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`, *optional*):
            Segment token indices to indicate first and second portions of the entity token inputs. Indices are
            selected in `[0, 1]`:

            - 0 corresponds to a *portion A* entity token,
            - 1 corresponds to a *portion B* entity token.
        entity_position_ids (`torch.LongTensor` of shape `(batch_size, entity_length, max_mention_length)`, *optional*):
            Indices of positions of each input entity in the position embeddings. Selected in the range `[0,
            config.max_position_embeddings - 1]`.
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the masked language modeling loss. Indices should be in `[-100, 0, ...,
            config.vocab_size]` (see `input_ids` docstring) Tokens with indices set to `-100` are ignored (masked), the
            loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`
        entity_labels (`torch.LongTensor` of shape `(batch_size, entity_length)`, *optional*):
            Labels for computing the masked language modeling loss. Indices should be in `[-100, 0, ...,
            config.vocab_size]` (see `input_ids` docstring) Tokens with indices set to `-100` are ignored (masked), the
            loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`
        NT©rm   rµ   rn   rl   r„   rb  rc  rd  ro   r¶   r  r  rb   c              3   ó   K  — | ]}|®|V — Œ	d S r�   r%   r  s     r'   r  z*LukeForMaskedLM.forward.<locals>.<genexpr>m  s4   è è € ð ð àð �=ð ð !�=�=�=ðð r&   )r,   r-   r.   r/   r0   r1   r   r2   )r_   r  r<  r€  r  rg   re   r‚  r¦   rP   r   r�  r   r$   r1   r   r2   r+   )r^   rm   rµ   rn   rl   r„   rb  rc  rd  r‰  rŠ  ro   r¶   r  r  ri  rÐ   r,   r-   r/   r.   r0   s                         r'   rr   zLukeForMaskedLM.forward  sÙ  € ðb &1Ð%<�k�kÀ$Ä+ÔBYˆà—)’)ØØ)Ø)Ø%Ø!Ø"7Ø"7Ø 3Ø'Ø/Ø!5Øð ñ 
ô 
ˆð ˆàˆØ—’˜gÔ7Ñ8Ô8ˆØÐà—Y’Y˜vœ}Ñ-Ô-ˆFØ—|’| F§K¢K°°D´KÔ4JÑ$KÔ$KÈVÏ[Ê[ÐY[É_Ì_Ñ]Ô]ˆHØˆ|Ø�àˆØˆØÔ+Ð7Ø ×3Ò3°GÔ4TÑUÔUˆMØÐ(ØŸ<š<¨×(:Ò(:¸2¸t¼{Ô?\Ñ(]Ô(]Ð_l×_qÒ_qÐrtÑ_uÔ_uÑvÔv�Ø�<Ø#�D�Dà (™?�Dàð 	Ýð ð ð ØØØØ!ØÔ)ØÔ0ØÔ&ð	ðñ ô ñ ô ð õ "ØØØØØ'Ø!Ô/Ø!(Ô!=ØÔ)ð	
ñ 	
ô 	
ð 		
r&   ©NNNNNNNNNNNNNN)r   r   r   Ú_tied_weights_keysrN   r…  rˆ  r   r!   r�   r"   rq  r$   r+   rr   rx   ry   s   @r'   r~  r~  õ  sÓ  ø€ € € € € ð ._Ø.ðð Ðð
ð ð ð ð ð$ð $ð $ð.ð .ð .ð ð .2Ø37Ø26Ø04Ø.2Ø9=Ø9=Ø7;Ø*.Ø15Ø26Ø)-Ø,0Ø#'ðp
ð p
àÔ# dÑ*ðp
ð Ô)¨DÑ0ðp
ð Ô(¨4Ñ/ð	p
ð
 Ô&¨Ñ-ðp
ð Ô$ tÑ+ðp
ð  %Ô/°$Ñ6ðp
ð  %Ô/°$Ñ6ðp
ð #Ô-°Ñ4ðp
ð Ô  4Ñ'ðp
ð Ô'¨$Ñ.ðp
ð Ô(¨4Ñ/ðp
ð   $™;ðp
ð # T™kðp
ð ˜D‘[ðp
ð" 
Ð#Ñ	#ð#p
ð p
ð p
ñ „^ðp
ð p
ð p
ð p
ð p
r&   r~  zº
    The LUKE model with a classification head on top (a linear layer on top of the hidden state of the first entity
    token) for entity classification tasks, such as Open Entity.
    c                   óB  ‡ — e Zd Zˆ fd„Ze	 	 	 	 	 	 	 	 	 	 	 	 	 ddej        dz  dej        dz  dej        dz  dej        dz  dej        dz  dej        dz  d	ej        dz  d
ej        dz  dej        dz  dej        dz  dedz  dedz  dedz  de	e
z  fd„¦   «         Zˆ xZS )ÚLukeForEntityClassificationc                 ó6  •— t          ¦   «                              |¦  «         t          |¦  «        | _        |j        | _        t          j        |j        ¦  «        | _        t          j	        |j
        |j        ¦  «        | _        |                      ¦   «          d S r�   ©rM   rN   rP  r<  Ú
num_labelsr   rZ   r[   r\   r‚   rQ   Ú
classifierrU  r]   s     €r'   rN   z$LukeForEntityClassification.__init__�  sy   ø€ Ý‰Œ×Ò˜Ñ Ô Ð å˜fÑ%Ô%ˆŒ	à Ô+ˆŒÝ”z &Ô"<Ñ=Ô=ˆŒÝœ) FÔ$6¸Ô8IÑJÔJˆŒð 	�ŠÑÔÐÐÐr&   Nrm   rµ   rn   rl   r„   rb  rc  rd  ro   r‰  r¶   r  r  rØ   c                 óÊ  — |�|n| j         j        }|                      |||||||||	||d¬¦  «        }|j        dd…ddd…f         }|                      |¦  «        }|                      |¦  «        }d}|
�Ÿ|
                     |j        ¦  «        }
|
j        dk    r!t          j
                             ||
¦  «        }nYt          j
                             |                     d¦  «        |
                     d¦  «                             |¦  «        ¦  «        }|s-t          d„ |||j        |j        |j        fD ¦   «         ¦  «        S t'          |||j        |j        |j        ¬¦  «        S )	u¸
  
        entity_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`):
            Indices of entity tokens in the entity vocabulary.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.
        entity_attention_mask (`torch.FloatTensor` of shape `(batch_size, entity_length)`, *optional*):
            Mask to avoid performing attention on padding entity token indices. Mask values selected in `[0, 1]`:

            - 1 for entity tokens that are **not masked**,
            - 0 for entity tokens that are **masked**.
        entity_token_type_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`, *optional*):
            Segment token indices to indicate first and second portions of the entity token inputs. Indices are
            selected in `[0, 1]`:

            - 0 corresponds to a *portion A* entity token,
            - 1 corresponds to a *portion B* entity token.
        entity_position_ids (`torch.LongTensor` of shape `(batch_size, entity_length, max_mention_length)`, *optional*):
            Indices of positions of each input entity in the position embeddings. Selected in the range `[0,
            config.max_position_embeddings - 1]`.
        labels (`torch.LongTensor` of shape `(batch_size,)` or `(batch_size, num_labels)`, *optional*):
            Labels for computing the classification loss. If the shape is `(batch_size,)`, the cross entropy loss is
            used for the single-label classification. In this case, labels should contain the indices that should be in
            `[0, ..., config.num_labels - 1]`. If the shape is `(batch_size, num_labels)`, the binary cross entropy
            loss is used for the multi-label classification. In this case, labels should only contain `[0, 1]`, where 0
            and 1 indicate false and true, respectively.

        Examples:

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

        >>> tokenizer = AutoTokenizer.from_pretrained("studio-ousia/luke-large-finetuned-open-entity")
        >>> model = LukeForEntityClassification.from_pretrained("studio-ousia/luke-large-finetuned-open-entity")

        >>> text = "BeyoncÃ© lives in Los Angeles."
        >>> entity_spans = [(0, 7)]  # character-based entity span corresponding to "BeyoncÃ©"
        >>> inputs = tokenizer(text, entity_spans=entity_spans, return_tensors="pt")
        >>> outputs = model(**inputs)
        >>> logits = outputs.logits
        >>> predicted_class_idx = logits.argmax(-1).item()
        >>> print("Predicted class:", model.config.id2label[predicted_class_idx])
        Predicted class: person
        ```NTrŒ  r   r   rb   c              3   ó   K  — | ]}|®|V — Œ	d S r�   r%   r  s     r'   r  z6LukeForEntityClassification.forward.<locals>.<genexpr>ú  ó0   è è € ð ð àØ�=ð à �=�=�=ðð r&   ©r,   r/   r1   r   r2   )r_   r  r<  r   r\   r•  rg   re   Úndimr   r±   Úcross_entropyÚ binary_cross_entropy_with_logitsr¦   rŒ   r$   r1   r   r2   r4   ©r^   rm   rµ   rn   rl   r„   rb  rc  rd  ro   r‰  r¶   r  r  ri  rÐ   Úfeature_vectorr/   r,   s                      r'   rr   z#LukeForEntityClassification.forward›  s–  € ð| &1Ð%<�k�kÀ$Ä+ÔBYˆà—)’)ØØ)Ø)Ø%Ø!Ø"7Ø"7Ø 3Ø'Ø/Ø!5Øð ñ 
ô 
ˆð !Ô9¸!¸!¸!¸QÀÀÀ¸'ÔBˆØŸš nÑ5Ô5ˆØ—’ Ñ0Ô0ˆàˆØÐð —Y’Y˜vœ}Ñ-Ô-ˆFØŒ{˜aÒÐÝ”}×2Ò2°6¸6ÑBÔB��å”}×EÒEÀfÇkÂkÐRTÁoÄoÐW]×WbÒWbÐceÑWfÔWf×WnÒWnÐouÑWvÔWvÑwÔw�àð 	Ýð ð à ¨Ô(=¸wÔ?[Ð]dÔ]oÐpðñ ô ñ ô ð õ *ØØØ!Ô/Ø!(Ô!=ØÔ)ð
ñ 
ô 
ð 	
r&   ©NNNNNNNNNNNNN)r   r   r   rN   r   r!   r�   r"   rq  r$   r4   rr   rx   ry   s   @r'   r‘  r‘  ˆ  s‰  ø€ € € € € ð
ð 
ð 
ð 
ð 
ð ð .2Ø37Ø26Ø04Ø.2Ø:>Ø9=Ø7;Ø26Ø+/Ø)-Ø,0Ø#'ðj
ð j
àÔ# dÑ*ðj
ð Ô)¨DÑ0ðj
ð Ô(¨4Ñ/ð	j
ð
 Ô&¨Ñ-ðj
ð Ô$ tÑ+ðj
ð  %Ô0°4Ñ7ðj
ð  %Ô/°$Ñ6ðj
ð #Ô-°Ñ4ðj
ð Ô(¨4Ñ/ðj
ð Ô! DÑ(ðj
ð   $™;ðj
ð # T™kðj
ð ˜D‘[ðj
ð  
Ð+Ñ	+ð!j
ð j
ð j
ñ „^ðj
ð j
ð j
ð j
ð j
r&   r‘  zº
    The LUKE model with a classification head on top (a linear layer on top of the hidden states of the two entity
    tokens) for entity pair classification tasks, such as TACRED.
    c                   óB  ‡ — e Zd Zˆ fd„Ze	 	 	 	 	 	 	 	 	 	 	 	 	 ddej        dz  dej        dz  dej        dz  dej        dz  dej        dz  dej        dz  d	ej        dz  d
ej        dz  dej        dz  dej        dz  dedz  dedz  dedz  de	e
z  fd„¦   «         Zˆ xZS )ÚLukeForEntityPairClassificationc                 ó>  •— t          ¦   «                              |¦  «         t          |¦  «        | _        |j        | _        t          j        |j        ¦  «        | _        t          j	        |j
        dz  |j        d¦  «        | _        |                      ¦   «          d S )Nr¥   Fr“  r]   s     €r'   rN   z(LukeForEntityPairClassification.__init__  s€   ø€ Ý‰Œ×Ò˜Ñ Ô Ð å˜fÑ%Ô%ˆŒ	à Ô+ˆŒÝ”z &Ô"<Ñ=Ô=ˆŒÝœ) FÔ$6¸Ñ$:¸FÔ<MÈuÑUÔUˆŒð 	�ŠÑÔÐÐÐr&   Nrm   rµ   rn   rl   r„   rb  rc  rd  ro   r‰  r¶   r  r  rØ   c                 ó  — |�|n| j         j        }|                      |||||||||	||d¬¦  «        }t          j        |j        dd…ddd…f         |j        dd…ddd…f         gd¬¦  «        }|                      |¦  «        }|                      |¦  «        }d}|
�Ÿ|
                     |j	        ¦  «        }
|
j
        dk    r!t          j                             ||
¦  «        }nYt          j                             |                     d¦  «        |
                     d¦  «                             |¦  «        ¦  «        }|s-t#          d„ |||j        |j        |j        fD ¦   «         ¦  «        S t+          |||j        |j        |j        ¬	¦  «        S )
u  
        entity_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`):
            Indices of entity tokens in the entity vocabulary.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.
        entity_attention_mask (`torch.FloatTensor` of shape `(batch_size, entity_length)`, *optional*):
            Mask to avoid performing attention on padding entity token indices. Mask values selected in `[0, 1]`:

            - 1 for entity tokens that are **not masked**,
            - 0 for entity tokens that are **masked**.
        entity_token_type_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`, *optional*):
            Segment token indices to indicate first and second portions of the entity token inputs. Indices are
            selected in `[0, 1]`:

            - 0 corresponds to a *portion A* entity token,
            - 1 corresponds to a *portion B* entity token.
        entity_position_ids (`torch.LongTensor` of shape `(batch_size, entity_length, max_mention_length)`, *optional*):
            Indices of positions of each input entity in the position embeddings. Selected in the range `[0,
            config.max_position_embeddings - 1]`.
        labels (`torch.LongTensor` of shape `(batch_size,)` or `(batch_size, num_labels)`, *optional*):
            Labels for computing the classification loss. If the shape is `(batch_size,)`, the cross entropy loss is
            used for the single-label classification. In this case, labels should contain the indices that should be in
            `[0, ..., config.num_labels - 1]`. If the shape is `(batch_size, num_labels)`, the binary cross entropy
            loss is used for the multi-label classification. In this case, labels should only contain `[0, 1]`, where 0
            and 1 indicate false and true, respectively.

        Examples:

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

        >>> tokenizer = AutoTokenizer.from_pretrained("studio-ousia/luke-large-finetuned-tacred")
        >>> model = LukeForEntityPairClassification.from_pretrained("studio-ousia/luke-large-finetuned-tacred")

        >>> text = "BeyoncÃ© lives in Los Angeles."
        >>> entity_spans = [
        ...     (0, 7),
        ...     (17, 28),
        ... ]  # character-based entity spans corresponding to "BeyoncÃ©" and "Los Angeles"
        >>> inputs = tokenizer(text, entity_spans=entity_spans, return_tensors="pt")
        >>> outputs = model(**inputs)
        >>> logits = outputs.logits
        >>> predicted_class_idx = logits.argmax(-1).item()
        >>> print("Predicted class:", model.config.id2label[predicted_class_idx])
        Predicted class: per:cities_of_residence
        ```NTrŒ  r   r   rˆ   rb   c              3   ó   K  — | ]}|®|V — Œ	d S r�   r%   r  s     r'   r  z:LukeForEntityPairClassification.forward.<locals>.<genexpr>€  r˜  r&   r™  )r_   r  r<  r!   r¬   r   r\   r•  rg   re   rš  r   r±   r›  rœ  r¦   rŒ   r$   r1   r   r2   r9   r�  s                      r'   rr   z'LukeForEntityPairClassification.forward  sÒ  € ðB &1Ð%<�k�kÀ$Ä+ÔBYˆà—)’)ØØ)Ø)Ø%Ø!Ø"7Ø"7Ø 3Ø'Ø/Ø!5Øð ñ 
ô 
ˆõ œØÔ-¨a¨a¨a°°A°A°A¨gÔ6¸Ô8XÐYZÐYZÐYZÐ\]Ð_`Ð_`Ð_`ÐY`Ô8aÐbÐhið
ñ 
ô 
ˆð Ÿš nÑ5Ô5ˆØ—’ Ñ0Ô0ˆàˆØÐð —Y’Y˜vœ}Ñ-Ô-ˆFØŒ{˜aÒÐÝ”}×2Ò2°6¸6ÑBÔB��å”}×EÒEÀfÇkÂkÐRTÁoÄoÐW]×WbÒWbÐceÑWfÔWf×WnÒWnÐouÑWvÔWvÑwÔw�àð 	Ýð ð à ¨Ô(=¸wÔ?[Ð]dÔ]oÐpðñ ô ñ ô ð õ .ØØØ!Ô/Ø!(Ô!=ØÔ)ð
ñ 
ô 
ð 	
r&   rŸ  )r   r   r   rN   r   r!   r�   r"   rq  r$   r9   rr   rx   ry   s   @r'   r¡  r¡  	  s‰  ø€ € € € € ð
ð 
ð 
ð 
ð 
ð ð .2Ø37Ø26Ø04Ø.2Ø:>Ø9=Ø7;Ø26Ø*.Ø)-Ø,0Ø#'ðo
ð o
àÔ# dÑ*ðo
ð Ô)¨DÑ0ðo
ð Ô(¨4Ñ/ð	o
ð
 Ô&¨Ñ-ðo
ð Ô$ tÑ+ðo
ð  %Ô0°4Ñ7ðo
ð  %Ô/°$Ñ6ðo
ð #Ô-°Ñ4ðo
ð Ô(¨4Ñ/ðo
ð Ô  4Ñ'ðo
ð   $™;ðo
ð # T™kðo
ð ˜D‘[ðo
ð  
Ð/Ñ	/ð!o
ð o
ð o
ñ „^ðo
ð o
ð o
ð o
ð o
r&   r¡  z£
    The LUKE model with a span classification head on top (a linear layer on top of the hidden states output) for tasks
    such as named entity recognition.
    c            #       ón  ‡ — e Zd Zˆ fd„Ze	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 ddej        dz  dej        dz  dej        dz  dej        dz  dej        dz  dej        dz  d	ej        dz  d
ej        dz  dej        dz  dej        dz  dej        dz  dej        dz  dedz  dedz  dedz  de	e
z  f d„¦   «         Zˆ xZS )ÚLukeForEntitySpanClassificationc                 ó<  •— t          ¦   «                              |¦  «         t          |¦  «        | _        |j        | _        t          j        |j        ¦  «        | _        t          j	        |j
        dz  |j        ¦  «        | _        |                      ¦   «          d S )Nr   r“  r]   s     €r'   rN   z(LukeForEntitySpanClassification.__init__–  s~   ø€ Ý‰Œ×Ò˜Ñ Ô Ð å˜fÑ%Ô%ˆŒ	à Ô+ˆŒÝ”z &Ô"<Ñ=Ô=ˆŒÝœ) FÔ$6¸Ñ$:¸FÔ<MÑNÔNˆŒð 	�ŠÑÔÐÐÐr&   Nrm   rµ   rn   rl   r„   rb  rc  rd  Úentity_start_positionsÚentity_end_positionsro   r‰  r¶   r  r  rØ   c                 óN  — |�|n| j         j        }|                      |||||||||||d¬¦  «        }|j                             d¦  «        }|	                     d¦  «                             dd|¦  «        }	|	j        |j        j        k    r|	                     |j        j        ¦  «        }	t          j
        |j        d|	¦  «        }|
                     d¦  «                             dd|¦  «        }
|
j        |j        j        k    r|
                     |j        j        ¦  «        }
t          j
        |j        d|
¦  «        }t          j        |||j        gd¬¦  «        }|                      |¦  «        }|                      |¦  «        }d}|�Ë|                     |j        ¦  «        }|j        dk    rMt           j                             |                     d| j        ¦  «        |                     d¦  «        ¦  «        }nYt           j                             |                     d¦  «        |                     d¦  «                             |¦  «        ¦  «        }|s-t/          d„ |||j        |j        |j        fD ¦   «         ¦  «        S t7          |||j        |j        |j        ¬	¦  «        S )
u  
        entity_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`):
            Indices of entity tokens in the entity vocabulary.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.
        entity_attention_mask (`torch.FloatTensor` of shape `(batch_size, entity_length)`, *optional*):
            Mask to avoid performing attention on padding entity token indices. Mask values selected in `[0, 1]`:

            - 1 for entity tokens that are **not masked**,
            - 0 for entity tokens that are **masked**.
        entity_token_type_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`, *optional*):
            Segment token indices to indicate first and second portions of the entity token inputs. Indices are
            selected in `[0, 1]`:

            - 0 corresponds to a *portion A* entity token,
            - 1 corresponds to a *portion B* entity token.
        entity_position_ids (`torch.LongTensor` of shape `(batch_size, entity_length, max_mention_length)`, *optional*):
            Indices of positions of each input entity in the position embeddings. Selected in the range `[0,
            config.max_position_embeddings - 1]`.
        entity_start_positions (`torch.LongTensor`):
            The start positions of entities in the word token sequence.
        entity_end_positions (`torch.LongTensor`):
            The end positions of entities in the word token sequence.
        labels (`torch.LongTensor` of shape `(batch_size, entity_length)` or `(batch_size, entity_length, num_labels)`, *optional*):
            Labels for computing the classification loss. If the shape is `(batch_size, entity_length)`, the cross
            entropy loss is used for the single-label classification. In this case, labels should contain the indices
            that should be in `[0, ..., config.num_labels - 1]`. If the shape is `(batch_size, entity_length,
            num_labels)`, the binary cross entropy loss is used for the multi-label classification. In this case,
            labels should only contain `[0, 1]`, where 0 and 1 indicate false and true, respectively.

        Examples:

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

        >>> tokenizer = AutoTokenizer.from_pretrained("studio-ousia/luke-large-finetuned-conll-2003")
        >>> model = LukeForEntitySpanClassification.from_pretrained("studio-ousia/luke-large-finetuned-conll-2003")

        >>> text = "BeyoncÃ© lives in Los Angeles"
        # List all possible entity spans in the text

        >>> word_start_positions = [0, 8, 14, 17, 21]  # character-based start positions of word tokens
        >>> word_end_positions = [7, 13, 16, 20, 28]  # character-based end positions of word tokens
        >>> entity_spans = []
        >>> for i, start_pos in enumerate(word_start_positions):
        ...     for end_pos in word_end_positions[i:]:
        ...         entity_spans.append((start_pos, end_pos))

        >>> inputs = tokenizer(text, entity_spans=entity_spans, return_tensors="pt")
        >>> outputs = model(**inputs)
        >>> logits = outputs.logits
        >>> predicted_class_indices = logits.argmax(-1).squeeze().tolist()
        >>> for span, predicted_class_idx in zip(entity_spans, predicted_class_indices):
        ...     if predicted_class_idx != 0:
        ...         print(text[span[0] : span[1]], model.config.id2label[predicted_class_idx])
        BeyoncÃ© PER
        Los Angeles LOC
        ```NTrŒ  rb   r‡   r¥   rˆ   c              3   ó   K  — | ]}|®|V — Œ	d S r�   r%   r  s     r'   r  z:LukeForEntitySpanClassification.forward.<locals>.<genexpr>  r˜  r&   r™  )r_   r  r<  r  ri   ru   rv   re   rg   r!   Úgatherr¬   r   r\   r•  rš  r   r±   r›  r¦   r”  rœ  rŒ   r$   r1   r   r2   r;   )r^   rm   rµ   rn   rl   r„   rb  rc  rd  r¨  r©  ro   r‰  r¶   r  r  ri  rÐ   rQ   Ústart_statesÚ
end_statesrž  r/   r,   s                           r'   rr   z'LukeForEntitySpanClassification.forward¢  s©  € ð^ &1Ð%<�k�kÀ$Ä+ÔBYˆà—)’)ØØ)Ø)Ø%Ø!Ø"7Ø"7Ø 3Ø'Ø/Ø!5Øð ñ 
ô 
ˆð Ô/×4Ò4°RÑ8Ô8ˆà!7×!AÒ!AÀ"Ñ!EÔ!E×!LÒ!LÈRÐQSÐU`Ñ!aÔ!aÐØ!Ô(¨GÔ,EÔ,LÒLÐLØ%;×%>Ò%>¸wÔ?XÔ?_Ñ%`Ô%`Ð"Ý”| GÔ$=¸rÐCYÑZÔZˆà3×=Ò=¸bÑAÔA×HÒHÈÈRÐQ\Ñ]Ô]ÐØÔ&¨'Ô*CÔ*JÒJÐJØ#7×#:Ò#:¸7Ô;TÔ;[Ñ#\Ô#\Ð Ý”\ 'Ô";¸RÐAUÑVÔVˆ
åœ L°*¸gÔ>^Ð#_ÐefÐgÑgÔgˆàŸš nÑ5Ô5ˆØ—’ Ñ0Ô0ˆàˆØÐà—Y’Y˜vœ}Ñ-Ô-ˆFð Œ{˜aÒÐÝ”}×2Ò2°6·;²;¸rÀ4Ä?Ñ3SÔ3SÐU[×U`ÒU`ÐacÑUdÔUdÑeÔe��å”}×EÒEÀfÇkÂkÐRTÁoÄoÐW]×WbÒWbÐceÑWfÔWf×WnÒWnÐouÑWvÔWvÑwÔw�àð 	Ýð ð à ¨Ô(=¸wÔ?[Ð]dÔ]oÐpðñ ô ñ ô ð õ .ØØØ!Ô/Ø!(Ô!=ØÔ)ð
ñ 
ô 
ð 	
r&   )NNNNNNNNNNNNNNN)r   r   r   rN   r   r!   r�   r"   rq  r$   r;   rr   rx   ry   s   @r'   r¦  r¦  �  sµ  ø€ € € € € ð
ð 
ð 
ð 
ð 
ð ð .2Ø37Ø26Ø04Ø.2Ø9=Ø9=Ø7;Ø:>Ø8<Ø26Ø*.Ø)-Ø,0Ø#'ð!G
ð G
àÔ# dÑ*ðG
ð Ô)¨DÑ0ðG
ð Ô(¨4Ñ/ð	G
ð
 Ô&¨Ñ-ðG
ð Ô$ tÑ+ðG
ð  %Ô/°$Ñ6ðG
ð  %Ô/°$Ñ6ðG
ð #Ô-°Ñ4ðG
ð !&Ô 0°4Ñ 7ðG
ð $Ô.°Ñ5ðG
ð Ô(¨4Ñ/ðG
ð Ô  4Ñ'ðG
ð   $™;ðG
ð # T™kðG
ð  ˜D‘[ð!G
ð$ 
Ð/Ñ	/ð%G
ð G
ð G
ñ „^ðG
ð G
ð G
ð G
ð G
r&   r¦  z 
    The LUKE Model transformer with a sequence classification/regression head on top (a linear layer on top of the
    pooled output) e.g. for GLUE tasks.
    c                   óB  ‡ — e Zd Zˆ fd„Ze	 	 	 	 	 	 	 	 	 	 	 	 	 ddej        dz  dej        dz  dej        dz  dej        dz  dej        dz  dej        dz  d	ej        dz  d
ej        dz  dej        dz  dej        dz  dedz  dedz  dedz  de	e
z  fd„¦   «         Zˆ xZS )ÚLukeForSequenceClassificationc                 óR  •— t          ¦   «                              |¦  «         |j        | _        t          |¦  «        | _        t          j        |j        �|j        n|j        ¦  «        | _	        t          j
        |j        |j        ¦  «        | _        |                      ¦   «          d S r�   ©rM   rN   r”  rP  r<  r   rZ   Úclassifier_dropoutr[   r\   r‚   rQ   r•  rU  r]   s     €r'   rN   z&LukeForSequenceClassification.__init__4  s‘   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ø Ô+ˆŒÝ˜fÑ%Ô%ˆŒ	Ý”zØ)/Ô)BÐ)NˆFÔ%Ð%ÐTZÔTnñ
ô 
ˆŒõ œ) FÔ$6¸Ô8IÑJÔJˆŒð 	�ŠÑÔÐÐÐr&   Nrm   rµ   rn   rl   r„   rb  rc  rd  ro   r‰  r¶   r  r  rØ   c                 ó\  — |�|n| j         j        }|                      |||||||||	||d¬¦  «        }|j        }|                      |¦  «        }|                      |¦  «        }d}|
��t|
                     |j        ¦  «        }
| j         j        €f| j	        dk    rd| j         _        nN| j	        dk    r7|
j
        t          j        k    s|
j
        t          j        k    rd| j         _        nd| j         _        | j         j        dk    rWt          ¦   «         }| j	        dk    r1 ||                     ¦   «         |
                     ¦   «         ¦  «        }nŽ |||
¦  «        }n�| j         j        dk    rGt!          ¦   «         } ||                     d| j	        ¦  «        |
                     d¦  «        ¦  «        }n*| j         j        dk    rt%          ¦   «         } |||
¦  «        }|s-t'          d	„ |||j        |j        |j        fD ¦   «         ¦  «        S t/          |||j        |j        |j        ¬
¦  «        S )a�  
        entity_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`):
            Indices of entity tokens in the entity vocabulary.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.
        entity_attention_mask (`torch.FloatTensor` of shape `(batch_size, entity_length)`, *optional*):
            Mask to avoid performing attention on padding entity token indices. Mask values selected in `[0, 1]`:

            - 1 for entity tokens that are **not masked**,
            - 0 for entity tokens that are **masked**.
        entity_token_type_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`, *optional*):
            Segment token indices to indicate first and second portions of the entity token inputs. Indices are
            selected in `[0, 1]`:

            - 0 corresponds to a *portion A* entity token,
            - 1 corresponds to a *portion B* entity token.
        entity_position_ids (`torch.LongTensor` of shape `(batch_size, entity_length, max_mention_length)`, *optional*):
            Indices of positions of each input entity in the position embeddings. Selected in the range `[0,
            config.max_position_embeddings - 1]`.
        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).
        NTrŒ  r   Ú
regressionÚsingle_label_classificationÚmulti_label_classificationrb   c              3   ó   K  — | ]}|®|V — Œ	d S r�   r%   r  s     r'   r  z8LukeForSequenceClassification.forward.<locals>.<genexpr>›  r˜  r&   r™  )r_   r  r<  rf  r\   r•  rg   re   Úproblem_typer”  rd   r!   rk   r™   r   Úsqueezer   r¦   r   r$   r1   r   r2   r=   )r^   rm   rµ   rn   rl   r„   rb  rc  rd  ro   r‰  r¶   r  r  ri  rÐ   r-  r/   r,   Úloss_fcts                       r'   rr   z%LukeForSequenceClassification.forward@  sW  € ðV &1Ð%<�k�kÀ$Ä+ÔBYˆà—)’)ØØ)Ø)Ø%Ø!Ø"7Ø"7Ø 3Ø'Ø/Ø!5Øð ñ 
ô 
ˆð  Ô-ˆàŸš ]Ñ3Ô3ˆØ—’ Ñ/Ô/ˆàˆØÑà—Y’Y˜vœ}Ñ-Ô-ˆFØŒ{Ô'Ð/Ø”? aÒ'Ð'Ø/;�D”KÔ,Ð,Ø”_ qÒ(Ð(¨f¬l½e¼jÒ.HÐ.HÈFÌLÕ\aÔ\eÒLeÐLeØ/L�D”KÔ,Ð,à/K�D”KÔ,àŒ{Ô'¨<Ò7Ð7Ý"™9œ9�Ø”? aÒ'Ð'Ø#˜8 F§N¢NÑ$4Ô$4°f·n²nÑ6FÔ6FÑGÔG�D�Dà#˜8 F¨FÑ3Ô3�D�DØ”Ô)Ð-JÒJÐJÝ+Ñ-Ô-�Ø�x §¢¨B°´Ñ @Ô @À&Ç+Â+ÈbÁ/Ä/ÑRÔR��Ø”Ô)Ð-IÒIÐIÝ,Ñ.Ô.�Ø�x ¨Ñ/Ô/�àð 	Ýð ð à ¨Ô(=¸wÔ?[Ð]dÔ]oÐpðñ ô ñ ô ð õ ,ØØØ!Ô/Ø!(Ô!=ØÔ)ð
ñ 
ô 
ð 	
r&   rŸ  )r   r   r   rN   r   r!   r�   r"   rq  r$   r=   rr   rx   ry   s   @r'   r°  r°  -  s‰  ø€ € € € € ð
ð 
ð 
ð 
ð 
ð ð .2Ø37Ø26Ø04Ø.2Ø:>Ø9=Ø7;Ø26Ø+/Ø)-Ø,0Ø#'ðf
ð f
àÔ# dÑ*ðf
ð Ô)¨DÑ0ðf
ð Ô(¨4Ñ/ð	f
ð
 Ô&¨Ñ-ðf
ð Ô$ tÑ+ðf
ð  %Ô0°4Ñ7ðf
ð  %Ô/°$Ñ6ðf
ð #Ô-°Ñ4ðf
ð Ô(¨4Ñ/ðf
ð Ô! DÑ(ðf
ð   $™;ðf
ð # T™kðf
ð ˜D‘[ðf
ð  
Ð-Ñ	-ð!f
ð f
ð f
ñ „^ðf
ð f
ð f
ð f
ð f
r&   r°  zú
    The LUKE Model with a token classification head on top (a linear layer on top of the hidden-states output). To
    solve Named-Entity Recognition (NER) task using LUKE, `LukeForEntitySpanClassification` is more suitable than this
    class.
    c                   óB  ‡ — e Zd Zˆ fd„Ze	 	 	 	 	 	 	 	 	 	 	 	 	 ddej        dz  dej        dz  dej        dz  dej        dz  dej        dz  dej        dz  d	ej        dz  d
ej        dz  dej        dz  dej        dz  dedz  dedz  dedz  de	e
z  fd„¦   «         Zˆ xZS )ÚLukeForTokenClassificationc                 óV  •— t          ¦   «                              |¦  «         |j        | _        t          |d¬¦  «        | _        t          j        |j        �|j        n|j        ¦  «        | _	        t          j
        |j        |j        ¦  «        | _        |                      ¦   «          d S ©NF)rQ  r²  r]   s     €r'   rN   z#LukeForTokenClassification.__init__²  s–   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ø Ô+ˆŒå˜f¸Ð>Ñ>Ô>ˆŒ	Ý”zØ)/Ô)BÐ)NˆFÔ%Ð%ÐTZÔTnñ
ô 
ˆŒõ œ) FÔ$6¸Ô8IÑJÔJˆŒð 	�ŠÑÔÐÐÐr&   Nrm   rµ   rn   rl   r„   rb  rc  rd  ro   r‰  r¶   r  r  rØ   c                 ó2  — |�|n| j         j        }|                      |||||||||	||d¬¦  «        }|j        }|                      |¦  «        }|                      |¦  «        }d}|
�`|
                     |j        ¦  «        }
t          ¦   «         } || 	                    d| j
        ¦  «        |
 	                    d¦  «        ¦  «        }|s-t          d„ |||j        |j        |j        fD ¦   «         ¦  «        S t          |||j        |j        |j        ¬¦  «        S )aM  
        entity_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`):
            Indices of entity tokens in the entity vocabulary.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.
        entity_attention_mask (`torch.FloatTensor` of shape `(batch_size, entity_length)`, *optional*):
            Mask to avoid performing attention on padding entity token indices. Mask values selected in `[0, 1]`:

            - 1 for entity tokens that are **not masked**,
            - 0 for entity tokens that are **masked**.
        entity_token_type_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`, *optional*):
            Segment token indices to indicate first and second portions of the entity token inputs. Indices are
            selected in `[0, 1]`:

            - 0 corresponds to a *portion A* entity token,
            - 1 corresponds to a *portion B* entity token.
        entity_position_ids (`torch.LongTensor` of shape `(batch_size, entity_length, max_mention_length)`, *optional*):
            Indices of positions of each input entity in the position embeddings. Selected in the range `[0,
            config.max_position_embeddings - 1]`.
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the multiple choice classification loss. Indices should be in `[0, ...,
            num_choices-1]` where `num_choices` is the size of the second dimension of the input tensors. (See
            `input_ids` above)
        NTrŒ  rb   c              3   ó   K  — | ]}|®|V — Œ	d S r�   r%   r  s     r'   r  z5LukeForTokenClassification.forward.<locals>.<genexpr>  r˜  r&   r™  )r_   r  r<  r  r\   r•  rg   re   r   r¦   r”  r$   r1   r   r2   r?   )r^   rm   rµ   rn   rl   r„   rb  rc  rd  ro   r‰  r¶   r  r  ri  rÐ   rp  r/   r,   r»  s                       r'   rr   z"LukeForTokenClassification.forward¿  sP  € ðV &1Ð%<�k�kÀ$Ä+ÔBYˆà—)’)ØØ)Ø)Ø%Ø!Ø"7Ø"7Ø 3Ø'Ø/Ø!5Øð ñ 
ô 
ˆð "Ô3ˆàŸ,š, Ñ7Ô7ˆØ—’ Ñ1Ô1ˆàˆØÐà—Y’Y˜vœ}Ñ-Ô-ˆFÝ'Ñ)Ô)ˆHØ�8˜FŸKšK¨¨D¬OÑ<Ô<¸f¿kºkÈ"¹o¼oÑNÔNˆDàð 	Ýð ð à ¨Ô(=¸wÔ?[Ð]dÔ]oÐpðñ ô ñ ô ð õ )ØØØ!Ô/Ø!(Ô!=ØÔ)ð
ñ 
ô 
ð 	
r&   rŸ  )r   r   r   rN   r   r!   r�   r"   rq  r$   r?   rr   rx   ry   s   @r'   r½  r½  ª  s‰  ø€ € € € € ðð ð ð ð ð ð .2Ø37Ø26Ø04Ø.2Ø:>Ø9=Ø7;Ø26Ø+/Ø)-Ø,0Ø#'ðT
ð T
àÔ# dÑ*ðT
ð Ô)¨DÑ0ðT
ð Ô(¨4Ñ/ð	T
ð
 Ô&¨Ñ-ðT
ð Ô$ tÑ+ðT
ð  %Ô0°4Ñ7ðT
ð  %Ô/°$Ñ6ðT
ð #Ô-°Ñ4ðT
ð Ô(¨4Ñ/ðT
ð Ô! DÑ(ðT
ð   $™;ðT
ð # T™kðT
ð ˜D‘[ðT
ð  
Ð*Ñ	*ð!T
ð T
ð T
ñ „^ðT
ð T
ð T
ð T
ð T
r&   r½  c            !       óX  ‡ — e Zd Zˆ fd„Ze	 	 	 	 	 	 	 	 	 	 	 	 	 	 ddej        dz  dej        dz  dej        dz  dej        dz  dej        dz  dej        dz  d	ej        dz  d
ej        dz  dej        dz  dej        dz  dej        dz  dedz  dedz  dedz  de	e
z  fd„¦   «         Zˆ xZS )ÚLukeForQuestionAnsweringc                 óþ   •— t          ¦   «                              |¦  «         |j        | _        t          |d¬¦  «        | _        t          j        |j        |j        ¦  «        | _        |  	                    ¦   «          d S r¿  )
rM   rN   r”  rP  r<  r   r‚   rQ   Ú
qa_outputsrU  r]   s     €r'   rN   z!LukeForQuestionAnswering.__init__  sj   ø€ Ý‰Œ×Ò˜Ñ Ô Ð à Ô+ˆŒå˜f¸Ð>Ñ>Ô>ˆŒ	Ýœ) FÔ$6¸Ô8IÑJÔJˆŒð 	�ŠÑÔÐÐÐr&   Nrm   rµ   rn   rl   r„   rb  rc  rd  ro   Ústart_positionsÚend_positionsr¶   r  r  rØ   c                 ó¢  — |�|n| j         j        }|                      |||||||||	||d¬¦  «        }|j        }|                      |¦  «        }|                     dd¬¦  «        \  }}|                     d¦  «        }|                     d¦  «        }d}|
�ç|�åt          |
                     ¦   «         ¦  «        dk    r|
                     d¦  «        }
t          |                     ¦   «         ¦  «        dk    r|                     d¦  «        }|                     d¦  «        }|
 	                    d|¦  «         | 	                    d|¦  «         t          |¬¦  «        } |||
¦  «        } |||¦  «        }||z   d	z  }|s.t          d
„ ||||j        |j        |j        fD ¦   «         ¦  «        S t          ||||j        |j        |j        ¬¦  «        S )a  
        entity_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`):
            Indices of entity tokens in the entity vocabulary.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.
        entity_attention_mask (`torch.FloatTensor` of shape `(batch_size, entity_length)`, *optional*):
            Mask to avoid performing attention on padding entity token indices. Mask values selected in `[0, 1]`:

            - 1 for entity tokens that are **not masked**,
            - 0 for entity tokens that are **masked**.
        entity_token_type_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`, *optional*):
            Segment token indices to indicate first and second portions of the entity token inputs. Indices are
            selected in `[0, 1]`:

            - 0 corresponds to a *portion A* entity token,
            - 1 corresponds to a *portion B* entity token.
        entity_position_ids (`torch.LongTensor` of shape `(batch_size, entity_length, max_mention_length)`, *optional*):
            Indices of positions of each input entity in the position embeddings. Selected in the range `[0,
            config.max_position_embeddings - 1]`.
        NTrŒ  r   rb   rˆ   r   )Úignore_indexr¥   c              3   ó   K  — | ]}|®|V — Œ	d S r�   r%   r  s     r'   r  z3LukeForQuestionAnswering.forward.<locals>.<genexpr>v  s4   è è € ð ð àð �=ð ð !�=�=�=ðð r&   )r,   rB   rC   r1   r   r2   )r_   r  r<  r  rÅ  Úsplitrº  Úlenri   Úclamp_r   r$   r1   r   r2   rA   )r^   rm   rµ   rn   rl   r„   rb  rc  rd  ro   rÆ  rÇ  r¶   r  r  ri  rÐ   rp  r/   rB   rC   Ú
total_lossÚignored_indexr»  Ú
start_lossÚend_losss                             r'   rr   z LukeForQuestionAnswering.forward$  s,  € ðP &1Ð%<�k�kÀ$Ä+ÔBYˆà—)’)ØØ)Ø)Ø%Ø!Ø"7Ø"7Ø 3Ø'Ø/Ø!5Øð ñ 
ô 
ˆð "Ô3ˆà—’ Ñ1Ô1ˆØ#)§<¢<°°r <Ñ#:Ô#:Ñ ˆ�jØ#×+Ò+¨BÑ/Ô/ˆØ×'Ò'¨Ñ+Ô+ˆ
àˆ
ØÐ&¨=Ð+Då�?×'Ò'Ñ)Ô)Ñ*Ô*¨QÒ.Ð.Ø"1×"9Ò"9¸"Ñ"=Ô"=�Ý�=×%Ò%Ñ'Ô'Ñ(Ô(¨1Ò,Ð,Ø -× 5Ò 5°bÑ 9Ô 9�à(×-Ò-¨aÑ0Ô0ˆMØ×"Ò" 1 mÑ4Ô4Ð4Ø× Ò   MÑ2Ô2Ð2å'°]ÐCÑCÔCˆHØ!˜ ,°Ñ@Ô@ˆJØ�x 
¨MÑ:Ô:ˆHØ$ xÑ/°1Ñ4ˆJàð 	Ýð ð ð Ø ØØÔ)ØÔ0ØÔ&ððñ ô ñ ô ð õ 0ØØ%Ø!Ø!Ô/Ø!(Ô!=ØÔ)ð
ñ 
ô 
ð 	
r&   rŽ  )r   r   r   rN   r   r!   r�   r"   rq  r$   rA   rr   rx   ry   s   @r'   rÃ  rÃ    sŸ  ø€ € € € € ð	ð 	ð 	ð 	ð 	ð ð .2Ø37Ø26Ø15Ø.2Ø:>Ø9=Ø7;Ø26Ø37Ø15Ø)-Ø,0Ø#'ðe
ð e
àÔ# dÑ*ðe
ð Ô)¨DÑ0ðe
ð Ô(¨4Ñ/ð	e
ð
 Ô'¨$Ñ.ðe
ð Ô$ tÑ+ðe
ð  %Ô0°4Ñ7ðe
ð  %Ô/°$Ñ6ðe
ð #Ô-°Ñ4ðe
ð Ô(¨4Ñ/ðe
ð Ô)¨DÑ0ðe
ð Ô'¨$Ñ.ðe
ð   $™;ðe
ð # T™kðe
ð ˜D‘[ðe
ð" 
Ð1Ñ	1ð#e
ð e
ð e
ñ „^ðe
ð e
ð e
ð e
ð e
r&   rÃ  c                   óB  ‡ — e Zd Zˆ fd„Ze	 	 	 	 	 	 	 	 	 	 	 	 	 ddej        dz  dej        dz  dej        dz  dej        dz  dej        dz  dej        dz  d	ej        dz  d
ej        dz  dej        dz  dej        dz  dedz  dedz  dedz  de	e
z  fd„¦   «         Zˆ xZS )ÚLukeForMultipleChoicec                 ó0  •— t          ¦   «                              |¦  «         t          |¦  «        | _        t	          j        |j        �|j        n|j        ¦  «        | _        t	          j	        |j
        d¦  «        | _        |                      ¦   «          d S rü   )rM   rN   rP  r<  r   rZ   r³  r[   r\   r‚   rQ   r•  rU  r]   s     €r'   rN   zLukeForMultipleChoice.__init__�  s„   ø€ Ý‰Œ×Ò˜Ñ Ô Ð å˜fÑ%Ô%ˆŒ	Ý”zØ)/Ô)BÐ)NˆFÔ%Ð%ÐTZÔTnñ
ô 
ˆŒõ œ) FÔ$6¸Ñ:Ô:ˆŒð 	�ŠÑÔÐÐÐr&   Nrm   rµ   rn   rl   r„   rb  rc  rd  ro   r‰  r¶   r  r  rØ   c                 ó¸  — |�|n| j         j        }|�|j        d         n|	j        d         }|�)|                     d|                     d¦  «        ¦  «        nd}|�)|                     d|                     d¦  «        ¦  «        nd}|�)|                     d|                     d¦  «        ¦  «        nd}|�)|                     d|                     d¦  «        ¦  «        nd}|	�=|	                     d|	                     d¦  «        |	                     d¦  «        ¦  «        nd}	|�)|                     d|                     d¦  «        ¦  «        nd}|�)|                     d|                     d¦  «        ¦  «        nd}|�)|                     d|                     d¦  «        ¦  «        nd}|�=|                     d|                     d¦  «        |                     d¦  «        ¦  «        nd}|                      |||||||||	||d¬¦  «        }|j        }|                      |¦  «        }|                      |¦  «        }|                     d|¦  «        }d}|
�4|
 	                    |j
        ¦  «        }
t          ¦   «         } |||
¦  «        }|s-t          d„ |||j        |j        |j        fD ¦   «         ¦  «        S t!          |||j        |j        |j        ¬¦  «        S )	a^  
        input_ids (`torch.LongTensor` of shape `(batch_size, num_choices, sequence_length)`):
            Indices of input sequence tokens in the vocabulary.

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

            [What are input IDs?](../glossary#input-ids)
        token_type_ids (`torch.LongTensor` of shape `(batch_size, num_choices, sequence_length)`, *optional*):
            Segment token indices to indicate first and second portions of the inputs. Indices are selected in `[0,
            1]`:

            - 0 corresponds to a *sentence A* token,
            - 1 corresponds to a *sentence B* token.

            [What are token type IDs?](../glossary#token-type-ids)
        position_ids (`torch.LongTensor` of shape `(batch_size, num_choices, sequence_length)`, *optional*):
            Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
            config.max_position_embeddings - 1]`.

            [What are position IDs?](../glossary#position-ids)
        entity_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`):
            Indices of entity tokens in the entity vocabulary.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.
        entity_attention_mask (`torch.FloatTensor` of shape `(batch_size, entity_length)`, *optional*):
            Mask to avoid performing attention on padding entity token indices. Mask values selected in `[0, 1]`:

            - 1 for entity tokens that are **not masked**,
            - 0 for entity tokens that are **masked**.
        entity_token_type_ids (`torch.LongTensor` of shape `(batch_size, entity_length)`, *optional*):
            Segment token indices to indicate first and second portions of the entity token inputs. Indices are
            selected in `[0, 1]`:

            - 0 corresponds to a *portion A* entity token,
            - 1 corresponds to a *portion B* entity token.
        entity_position_ids (`torch.LongTensor` of shape `(batch_size, entity_length, max_mention_length)`, *optional*):
            Indices of positions of each input entity in the position embeddings. Selected in the range `[0,
            config.max_position_embeddings - 1]`.
        inputs_embeds (`torch.FloatTensor` of shape `(batch_size, num_choices, 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.
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the multiple choice classification loss. Indices should be in `[0, ...,
            num_choices-1]` where `num_choices` is the size of the second dimension of the input tensors. (See
            `input_ids` above)
        Nr   rb   r‡   TrŒ  c              3   ó   K  — | ]}|®|V — Œ	d S r�   r%   r  s     r'   r  z0LukeForMultipleChoice.forward.<locals>.<genexpr>  r  r&   r™  )r_   r  Úshaper¦   ri   r<  rf  r\   r•  rg   re   r   r$   r1   r   r2   rE   )r^   rm   rµ   rn   rl   r„   rb  rc  rd  ro   r‰  r¶   r  r  ri  Únum_choicesrÐ   r-  r/   Úreshaped_logitsr,   r»  s                         r'   rr   zLukeForMultipleChoice.forward›  sQ  € ðF &1Ð%<�k�kÀ$Ä+ÔBYˆØ,5Ð,A�i”o aÔ(Ð(À}ÔGZÐ[\ÔG]ˆà>GÐ>S�I—N’N 2 y§~¢~°bÑ'9Ô'9Ñ:Ô:Ð:ÐY]ˆ	ØM[ÐMg˜×,Ò,¨R°×1DÒ1DÀRÑ1HÔ1HÑIÔIÐIÐmqˆØM[ÐMg˜×,Ò,¨R°×1DÒ1DÀRÑ1HÔ1HÑIÔIÐIÐmqˆØGSÐG_�|×(Ò(¨¨\×->Ò->¸rÑ-BÔ-BÑCÔCÐCÐeiˆð Ð(ð ×Ò˜r =×#5Ò#5°bÑ#9Ô#9¸=×;MÒ;MÈbÑ;QÔ;QÑRÔRÐRàð 	ð BLÐAW�Z—_’_ R¨¯ª¸Ñ)<Ô)<Ñ=Ô=Ð=Ð]aˆ
ð %Ð0ð "×&Ò& rÐ+@×+EÒ+EÀbÑ+IÔ+IÑJÔJÐJàð 	ð %Ð0ð "×&Ò& rÐ+@×+EÒ+EÀbÑ+IÔ+IÑJÔJÐJàð 	ð #Ð.ð  ×$Ò$ RÐ)<×)AÒ)AÀ"Ñ)EÔ)EÐGZ×G_ÒG_Ð`bÑGcÔGcÑdÔdÐdàð 	ð —)’)ØØ)Ø)Ø%Ø!Ø"7Ø"7Ø 3Ø'Ø/Ø!5Øð ñ 
ô 
ˆð  Ô-ˆàŸš ]Ñ3Ô3ˆØ—’ Ñ/Ô/ˆØ Ÿ+š+ b¨+Ñ6Ô6ˆàˆØÐà—Y’Y˜Ô5Ñ6Ô6ˆFÝ'Ñ)Ô)ˆHØ�8˜O¨VÑ4Ô4ˆDàð 	Ýð 
ð 
ð Ø#ØÔ)ØÔ0ØÔ&ðð
ñ 
ô 
ñ 
ô 
ð 
õ -ØØ"Ø!Ô/Ø!(Ô!=ØÔ)ð
ñ 
ô 
ð 	
r&   rŸ  )r   r   r   rN   r   r!   r�   r"   rq  r$   rE   rr   rx   ry   s   @r'   rÓ  rÓ  �  s‰  ø€ € € € € ð
ð 
ð 
ð 
ð 
ð ð .2Ø37Ø26Ø04Ø.2Ø:>Ø9=Ø7;Ø26Ø+/Ø)-Ø,0Ø#'ðO
ð O
àÔ# dÑ*ðO
ð Ô)¨DÑ0ðO
ð Ô(¨4Ñ/ð	O
ð
 Ô&¨Ñ-ðO
ð Ô$ tÑ+ðO
ð  %Ô0°4Ñ7ðO
ð  %Ô/°$Ñ6ðO
ð #Ô-°Ñ4ðO
ð Ô(¨4Ñ/ðO
ð Ô! DÑ(ðO
ð   $™;ðO
ð # T™kðO
ð ˜D‘[ðO
ð  
Ð.Ñ	.ð!O
ð O
ð O
ñ „^ðO
ð O
ð O
ð O
ð O
r&   rÓ  )
r‘  r¡  r¦  rÓ  rÃ  r°  r½  r~  rP  r;  )Hr    r¯   Údataclassesr   r!   r   Útorch.nnr   r   r   Ú r	   rD  Úactivationsr
   r   Úmasking_utilsr   Úmodeling_layersr   Úmodeling_outputsr   r   Úmodeling_utilsr   Úpytorch_utilsr   Úutilsr   r   r   Úconfiguration_luker   Ú
get_loggerr   Úloggerr   r)   r+   r4   r9   r;   r=   r?   rA   rE   rN  rG   r{   r’   rÓ   rá   rì   rö   rú   r  r'  r/  r4  r;  rP  rf   rx  r~  r‘  r¡  r¦  r°  r½  rÃ  rÓ  Ú__all__r%   r&   r'   ú<module>rè     sh
  ðð Ð à €€€Ø !Ð !Ð !Ð !Ð !Ð !à €€€Ø Ð Ð Ð Ð Ð Ø AÐ AÐ AÐ AÐ AÐ AÐ AÐ AÐ AÐ Aà &Ð &Ð &Ð &Ð &Ð &Ø 'Ð 'Ð 'Ð 'Ð 'Ð 'Ð 'Ð 'Ø 6Ð 6Ð 6Ð 6Ð 6Ð 6Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø KÐ KÐ KÐ KÐ KÐ KÐ KÐ KØ -Ð -Ð -Ð -Ð -Ð -Ø 6Ð 6Ð 6Ð 6Ð 6Ð 6Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø *Ð *Ð *Ð *Ð *Ð *ð 
ˆÔ	˜HÑ	%Ô	%€ð €ððñ ô ð
 ðFð Fð Fð Fð FÐ%?ñ Fô Fñ „ñô ðFð" €ððñ ô ð
 ðFð Fð Fð Fð F˜/ñ Fô Fñ „ñô ðFð €ððñ ô ð
 ð<ð <ð <ð <ð <˜ñ <ô <ñ „ñô ð<ð8 €ððñ ô ð
 ð<ð <ð <ð <ð < ñ <ô <ñ „ñô ð<ð& €ððñ ô ð
 ð<ð <ð <ð <ð < [ñ <ô <ñ „ñô ð<ð& €ððñ ô ð
 ð<ð <ð <ð <ð < [ñ <ô <ñ „ñô ð<ð& €ððñ ô ð
 ð<ð <ð <ð <ð < ;ñ <ô <ñ „ñô ð<ð& €ððñ ô ð
 ð<ð <ð <ð <ð < ñ <ô <ñ „ñô ð<ð& €ððñ ô ð
 ð<ð <ð <ð <ð < {ñ <ô <ñ „ñô ð<ð$ €ððñ ô ð
 ð<ð <ð <ð <ð < Kñ <ô <ñ „ñô ð<ð*D=ð D=ð D=ð D=ð D=�R”Yñ D=ô D=ð D=ðN(ð (ð (ð (ð (˜2œ9ñ (ô (ð (ðVið ið ið ið i˜œ	ñ iô ið iðZð ð ð ð �R”Yñ ô ð ð&ð &ð &ð &ð &�B”Iñ &ô &ð &ðTð ð ð ð �r”yñ ô ð ð ð ð ð ð �”ñ ô ð ð/ð /ð /ð /ð /Ð*ñ /ô /ð /ðd>
ð >
ð >
ð >
ð >
�"”)ñ >
ô >
ð >
ðDð ð ð ð �”ñ ô ð ðð ð ð ð  B¤Iñ ô ð ð"ð ð ð ð ˜2œ9ñ ô ð ð ð?ð ?ð ?ð ?ð ?˜/ñ ?ô ?ñ „ð?ð( €ððñ ô ð
w
ð w
ð w
ð w
ð w
Ð#ñ w
ô w
ñô ð
w
ðt4ð 4ð 4ð"ð ð ð ð �”ñ ô ð ð, €ððñ ô ðJ
ð J
ð J
ð J
ð J
Ð)ñ J
ô J
ñô ðJ
ðZ €ððñ ô ðx
ð x
ð x
ð x
ð x
Ð"5ñ x
ô x
ñô ðx
ðv €ððñ ô ð}
ð }
ð }
ð }
ð }
Ð&9ñ }
ô }
ñô ð}
ð@ €ððñ ô ðU
ð U
ð U
ð U
ð U
Ð&9ñ U
ô U
ñô ðU
ðp €ððñ ô ðt
ð t
ð t
ð t
ð t
Ð$7ñ t
ô t
ñô ðt
ðn €ððñ ô ðc
ð c
ð c
ð c
ð c
Ð!4ñ c
ô c
ñô ðc
ðL ðr
ð r
ð r
ð r
ð r
Ð2ñ r
ô r
ñ „ðr
ðj ð]
ð ]
ð ]
ð ]
ð ]
Ð/ñ ]
ô ]
ñ „ð]
ð@ð ð €€€r&   