§
    ‚ŠtjZ ã                   ól  — d Z ddlmZ ddlmZ ddlZddlmZ ddlmZm	Z	 ddl
mZ dd	lmZmZmZmZmZ dd
lmZ ddlmZ ddlmZ ddlmZmZ ddlmZ ddlmZ  ej        e ¦  «        Z! ed¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z"ee G d„ de¦  «        ¦   «         ¦   «         Z# ed¬¦  «        e G d„ de¦  «        ¦   «         ¦   «         Z$e G d„ de$¦  «        ¦   «         Z% ed¬¦  «         G d„ de$¦  «        ¦   «         Z& ed¬¦  «         G d „ d!e$e¦  «        ¦   «         Z'g d"¢Z(dS )#zRAG model implementation.é    )ÚCallable)Ú	dataclassN)Únné   )ÚCacheÚEncoderDecoderCache)ÚPreTrainedConfig)ÚGenerationConfigÚGenerationMixinÚGenerationModeÚLogitsProcessorListÚStoppingCriteriaList)ÚGENERATION_MODES_MAPPING)ÚModelOutput)ÚPreTrainedModel)Úauto_docstringÚloggingé   )Ú	RagConfig)ÚRagRetrieverzI
    Base class for retriever augmented marginalized models outputs.
    )Úcustom_introc                   óx  — 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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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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Zeej        df         dz  ed<   dZeej        df         dz  ed<   dS )ÚRetrievAugLMMarginOutputa¢  
    loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided):
        Language modeling loss.
    logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):
        Prediction scores of the language modeling head. The score is possibly marginalized over all documents for
        each vocabulary token.
    doc_scores (`torch.FloatTensor` of shape `(batch_size, config.n_docs)`):
        Score between each retrieved document embeddings (see `retrieved_doc_embeds`) and
        `question_encoder_last_hidden_state`.
    past_key_values (`Cache`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
        It is a [`~cache_utils.Cache`] instance. For more details, see our [kv cache guide](https://huggingface.co/docs/transformers/en/kv_cache).

        Contains precomputed hidden-states (key and values in the attention blocks) of the decoder that can be used
        (see `past_key_values` input) to speed up sequential decoding.
    retrieved_doc_embeds (`torch.FloatTensor` of shape `(batch_size, config.n_docs, hidden_size)`, *optional*, returned when *output_retrieved=True*):
        Embedded documents retrieved by the retriever. Is used with `question_encoder_last_hidden_state` to compute
        the `doc_scores`.
    retrieved_doc_ids (`torch.LongTensor` of shape `(batch_size, config.n_docs)`, *optional*, returned when *output_retrieved=True*):
        The indexes of the embedded documents retrieved by the retriever.
    context_input_ids (`torch.LongTensor` of shape `(batch_size * config.n_docs, config.max_combined_length)`, *optional*, returned when *output_retrieved=True*):
        Input ids post-processed from the retrieved documents and the question encoder input_ids by the retriever.
    context_attention_mask (`torch.LongTensor` of shape `(batch_size * config.n_docs, config.max_combined_length)`, *optional*, returned when *output_retrieved=True*):
        Attention mask post-processed from the retrieved documents and the question encoder `input_ids` by the
        retriever.
    question_encoder_last_hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
        Sequence of hidden states at the output of the last layer of the question encoder pooled output of the
        model.
    question_enc_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 and one for the output of each layer) of
        shape `(batch_size, sequence_length, hidden_size)`.

        Hidden states of the question encoder at the output of each layer plus the initial embedding outputs.
    question_enc_attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
        Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
        sequence_length)`.

        Attentions weights of the question encoder, after the attention softmax, used to compute the weighted
        average in the self-attention heads.
    generator_enc_last_hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
        Sequence of hidden-states at the output of the last layer of the generator encoder of the model.
    generator_enc_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 and one for the output of each layer) of
        shape `(batch_size, sequence_length, hidden_size)`.

        Hidden states of the generator encoder at the output of each layer plus the initial embedding outputs.
    generator_enc_attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
        Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
        sequence_length)`.

        Attentions weights of the generator encoder, after the attention softmax, used to compute the weighted
        average in the self-attention heads.
    generator_dec_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 and one for the output of each layer) of
        shape `(batch_size, sequence_length, hidden_size)`.

        Hidden states of the generator decoder at the output of each layer plus the initial embedding outputs.
    generator_dec_attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
        Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
        sequence_length)`.

        Attentions weights of the generator decoder, after the attention softmax, used to compute the weighted
        average in the self-attention heads.
    generator_cross_attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
        Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
        sequence_length)`.

        Cross-attentions weights of the generator decoder, after the attention softmax, used to compute the
        weighted average in the cross-attention heads.
    NÚlossÚlogitsÚ
doc_scoresÚpast_key_valuesÚretrieved_doc_embedsÚretrieved_doc_idsÚcontext_input_idsÚcontext_attention_maskÚ"question_encoder_last_hidden_state.Úquestion_enc_hidden_statesÚquestion_enc_attentionsÚgenerator_enc_last_hidden_stateÚgenerator_enc_hidden_statesÚgenerator_enc_attentionsÚgenerator_dec_hidden_statesÚgenerator_dec_attentionsÚgenerator_cross_attentions)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚtorchÚFloatTensorÚ__annotations__r   r   r   r   r   r   Ú
LongTensorr    r!   r"   r#   Útupler$   r%   r&   r'   r(   r)   r*   © ó    úb/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/rag/modeling_rag.pyr   r   $   s  € € € € € € ðDð DðL &*€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø'+€FˆEÔ Ñ$Ð+Ð+Ñ+Ø+/€J�Ô! DÑ(Ð/Ð/Ñ/Ø$(€O�U˜T‘\Ð(Ð(Ñ(Ø59Ð˜%Ô+¨dÑ2Ð9Ð9Ñ9Ø15Ð�uÔ'¨$Ñ.Ð5Ð5Ñ5Ø15Ð�uÔ'¨$Ñ.Ð5Ð5Ñ5Ø6:Ð˜EÔ,¨tÑ3Ð:Ð:Ñ:ØCGÐ&¨Ô(9¸DÑ(@ÐGÐGÑGØGKÐ  eÔ&7¸Ð&<Ô =ÀÑ DÐKÐKÑKØDHÐ˜U 5Ô#4°cÐ#9Ô:¸TÑAÐHÐHÑHØ@DÐ# UÔ%6¸Ñ%=ÐDÐDÑDØHLÐ  uÔ'8¸#Ð'=Ô!>ÀÑ!EÐLÐLÑLØEIÐ˜e EÔ$5°sÐ$:Ô;¸dÑBÐIÐIÑIØHLÐ  uÔ'8¸#Ð'=Ô!>ÀÑ!EÐLÐLÑLØEIÐ˜e EÔ$5°sÐ$:Ô;¸dÑBÐIÐIÑIØGKÐ  eÔ&7¸Ð&<Ô =ÀÑ DÐKÐKÑKÐKÐKr5   r   c                   óZ  — 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
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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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Zeej        df         dz  ed<   dZeej        df         dz  ed<   dS )ÚRetrievAugLMOutputa"  
    logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):
        Prediction scores of the language modeling head. The score is possibly marginalized over all documents for
        each vocabulary token.
    doc_scores (`torch.FloatTensor` of shape `(batch_size, config.n_docs)`):
        Score between each retrieved document embeddings (see `retrieved_doc_embeds`) and
        `question_encoder_last_hidden_state`.
    past_key_values (`Cache`, *optional*, returned when `use_cache=True` is passed or when `config.use_cache=True`):
        It is a [`~cache_utils.Cache`] instance. For more details, see our [kv cache guide](https://huggingface.co/docs/transformers/en/kv_cache).

        Contains precomputed hidden-states (key and values in the attention blocks) of the decoder that can be used
        (see `past_key_values` input) to speed up sequential decoding.
    retrieved_doc_embeds (`torch.FloatTensor` of shape `(batch_size, config.n_docs, hidden_size)`, *optional*, returned when *output_retrieved=True*):
        Embedded documents retrieved by the retriever. Is used with `question_encoder_last_hidden_state` to compute
        the `doc_scores`.
    retrieved_doc_ids (`torch.LongTensor` of shape `(batch_size, config.n_docs)`, *optional*, returned when *output_retrieved=True*):
        The indexes of the embedded documents retrieved by the retriever.
    context_input_ids (`torch.LongTensor` of shape `(batch_size * config.n_docs, config.max_combined_length)`, *optional*, returned when *output_retrieved=True*):
        Input ids post-processed from the retrieved documents and the question encoder input_ids by the retriever.
    context_attention_mask (`torch.LongTensor` of shape `(batch_size * config.n_docs, config.max_combined_length)`, *optional*, returned when *output_retrieved=True*):
        Attention mask post-processed from the retrieved documents and the question encoder `input_ids` by the
        retriever.
    question_encoder_last_hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
        Sequence of hidden states at the output of the last layer of the question encoder pooled output of the
        model.
    question_enc_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 and one for the output of each layer) of
        shape `(batch_size, sequence_length, hidden_size)`.

        Hidden states of the question encoder at the output of each layer plus the initial embedding outputs.
    question_enc_attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
        Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
        sequence_length)`.

        Attentions weights of the question encoder, after the attention softmax, used to compute the weighted
        average in the self-attention heads.
    generator_enc_last_hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
        Sequence of hidden-states at the output of the last layer of the generator encoder of the model.
    generator_enc_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 and one for the output of each layer) of
        shape `(batch_size, sequence_length, hidden_size)`.

        Hidden states of the generator encoder at the output of each layer plus the initial embedding outputs.
    generator_enc_attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
        Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
        sequence_length)`.

        Attentions weights of the generator encoder, after the attention softmax, used to compute the weighted
        average in the self-attention heads.
    generator_dec_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 and one for the output of each layer) of
        shape `(batch_size, sequence_length, hidden_size)`.

        Hidden states of the generator decoder at the output of each layer plus the initial embedding outputs.
    generator_dec_attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
        Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
        sequence_length)`.

        Attentions weights of the generator decoder, after the attention softmax, used to compute the weighted
        average in the self-attention heads.
    generator_cross_attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
        Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
        sequence_length)`.

        Cross-attentions weights of the generator decoder, after the attention softmax, used to compute the
        weighted average in the cross-attention heads.
    Nr   r   r   r   r   r    r!   r"   .r#   r$   r%   r&   r'   r(   r)   r*   )r+   r,   r-   r.   r   r/   r0   r1   r   r   r   r   r   r2   r    r!   r"   r#   r3   r$   r%   r&   r'   r(   r)   r*   r4   r5   r6   r8   r8   „   sð  € € € € € € ðBð BðH (,€FˆEÔ Ñ$Ð+Ð+Ñ+Ø+/€J�Ô! DÑ(Ð/Ð/Ñ/Ø$(€O�U˜T‘\Ð(Ð(Ñ(Ø59Ð˜%Ô+¨dÑ2Ð9Ð9Ñ9Ø15Ð�uÔ'¨$Ñ.Ð5Ð5Ñ5Ø15Ð�uÔ'¨$Ñ.Ð5Ð5Ñ5Ø6:Ð˜EÔ,¨tÑ3Ð:Ð:Ñ:ØCGÐ&¨Ô(9¸DÑ(@ÐGÐGÑGØGKÐ  eÔ&7¸Ð&<Ô =ÀÑ DÐKÐKÑKØDHÐ˜U 5Ô#4°cÐ#9Ô:¸TÑAÐHÐHÑHØ@DÐ# UÔ%6¸Ñ%=ÐDÐDÑDØHLÐ  uÔ'8¸#Ð'=Ô!>ÀÑ!EÐLÐLÑLØEIÐ˜e EÔ$5°sÐ$:Ô;¸dÑBÐIÐIÑIØHLÐ  uÔ'8¸#Ð'=Ô!>ÀÑ!EÐLÐLÑLØEIÐ˜e EÔ$5°sÐ$:Ô;¸dÑBÐIÐIÑIØGKÐ  eÔ&7¸Ð&<Ô =ÀÑ DÐKÐKÑKÐKÐKr5   r8   a¹  
    RAG models were released with the paper [Retrieval-Augmented Generation for Knowledge-Intensive NLP
    Tasks](https://huggingface.co/papers/2005.11401) by Patrick Lewis, Ethan Perez, Aleksandra Piktus et al.

    RAG is a retriever augmented model and encapsulate three components: a question encoder, a dataset retriever and a
    generator, the encoder and generator are trainable while the retriever is just an indexed dataset.
    c            
       óh   — e Zd ZU eed<   dZdZdZe	 	 	 d
de	dz  de	dz  de
dz  defd	„¦   «         ZdS )ÚRagPreTrainedModelÚconfigÚragTNÚ.question_encoder_pretrained_model_name_or_pathÚ'generator_pretrained_model_name_or_pathÚ	retrieverÚreturnc                 óœ  — d„ |                      ¦   «         D ¦   «         }d„ |                      ¦   «         D ¦   «         }|D ]}|d|z   = Œ	|D ]}|d|z   = Œ	|                     dd¦  «        }|€D|€
J d¦   «         ‚dd	lm}	 d
|vr ddlm}
  |
j        |fi |¤ddi¤Ž\  }}||d
<    |	j        |fi |¤Ž}|                     dd¦  «        }|€D|€
J d¦   «         ‚ddlm} d
|vr ddlm}
  |
j        |fi |¤ddi¤Ž\  }}||d
<    |j        |fi |¤Ž}|                     d
¦  «        }|€t          j
        |j        |j        fi |¤Ž} | ||||¬¦  «        S )a  
        Instantiates an question encoder and a generator from one or two base classes of the library from pretrained
        model checkpoints.

        The model is set in evaluation mode by default using `model.eval()` (Dropout modules are deactivated). To train
        the model, you need to first set it back in training mode with `model.train()`.

        Params:
            question_encoder_pretrained_model_name_or_path (`str`, *optional*, defaults to `None`):
                Information necessary to initiate the question encoder. Can be either:

                    - A string, the *model id* of a pretrained model hosted inside a model repo on huggingface.co.
                    - A path to a *directory* containing model weights saved using
                      [`~PreTrainedModel.save_pretrained`], e.g., `./my_model_directory/`.

            generator_pretrained_model_name_or_path (`str`, *optional*, defaults to `None`):
                Information necessary to initiate the generator. Can be either:

                    - A string, the *model id* of a pretrained model hosted inside a model repo on huggingface.co.
                    - A path to a *directory* containing model weights saved using
                      [`~PreTrainedModel.save_pretrained`], e.g., `./my_model_directory/`.

            model_args (remaining positional arguments, *optional*):
                All remaining positional arguments will be passed to the underlying model's `__init__` method.
            retriever ([`RagRetriever`], *optional*):
                The retriever to use.
            kwwargs (remaining dictionary of keyword arguments, *optional*):
                Can be used to update the configuration object (after it being loaded) and initiate the model (e.g.,
                `output_attentions=True`).

                - To update the question_encoder configuration, use the prefix *question_encoder_* for each
                  configuration parameter.
                - To update the generator configuration, use the prefix *generator_* for each configuration parameter.
                - To update the parent model configuration, do not use a prefix for each configuration parameter.

                Behaves differently depending on whether a `config` is provided or automatically loaded.

        Example:

        ```python
        >>> from transformers import RagModel

        >>> # initialize a RAG from two pretrained models.
        >>> model = RagModel.from_pretrained_question_encoder_generator(
        ...     "facebook/dpr-question_encoder-single-nq-base", "google-t5/t5-small"
        ... )
        >>> # saving model after fine-tuning
        >>> model.save_pretrained("./rag")
        >>> # load fine-tuned model
        >>> model = RagModel.from_pretrained("./rag")
        ```c                 ón   — i | ]2\  }}|                      d ¦  «        ¯|t          d ¦  «        d…         |“Œ3S )Úquestion_encoder_N©Ú
startswithÚlen©Ú.0ÚargumentÚvalues      r6   ú
<dictcomp>zQRagPreTrainedModel.from_pretrained_question_encoder_generator.<locals>.<dictcomp>)  sW   € ð #
ð #
ð #
á�˜%Ø×"Ò"Ð#6Ñ7Ô7ð#
Ø•SÐ,Ñ-Ô-Ð/Ð/Ô0°%ð#
ð #
ð #
r5   c                 ón   — i | ]2\  }}|                      d ¦  «        ¯|t          d ¦  «        d…         |“Œ3S )Ú
generator_NrD   rG   s      r6   rK   zQRagPreTrainedModel.from_pretrained_question_encoder_generator.<locals>.<dictcomp>/  sU   € ð 
ð 
ð 
á�˜%Ø×"Ò" <Ñ0Ô0ð
Ø•S˜Ñ&Ô&Ð(Ð(Ô)¨5ð
ð 
ð 
r5   rC   rM   ÚmodelNznIf `model` is not defined as an argument, a `question_encoder_pretrained_model_name_or_path` has to be definedé   ©Ú	AutoModelr;   )Ú
AutoConfigÚreturn_unused_kwargsTzqIf `generator_model` is not defined as an argument, a `generator_pretrained_model_name_or_path` has to be defined©ÚAutoModelForSeq2SeqLM)Úquestion_encoderÚ	generatorr;   r?   )ÚitemsÚpopÚauto.modeling_autorQ   Úauto.configuration_autorR   Úfrom_pretrainedrU   Úgetr   Ú'from_question_encoder_generator_configsr;   )Úclsr=   r>   r?   ÚkwargsÚkwargs_question_encoderÚkwargs_generatorÚkeyrV   rQ   rR   Úquestion_encoder_configrW   rU   Úgenerator_configr;   s                   r6   Ú*from_pretrained_question_encoder_generatorz=RagPreTrainedModel.from_pretrained_question_encoder_generatorí   s®  € ðx#
ð #
à#)§<¢<¡>¤>ð#
ñ #
ô #
Ðð
ð 
à#)§<¢<¡>¤>ð
ñ 
ô 
Ðð +ð 	2ð 	2ˆCØÐ*¨SÑ0Ð1Ð1Ø#ð 	+ð 	+ˆCØ�| cÑ)Ð*Ð*ð
 3×6Ò6°wÀÑEÔEÐØÐ#ØAÐMÐMðñ NÔMÐMð 7Ð6Ð6Ð6Ð6Ð6àÐ6Ð6Ð6Ø@Ð@Ð@Ð@Ð@Ð@àC]À:ÔC]ØBðDð Dà-ðDð Dð *.ðDð Dð DÑ@Ð'Ð)@ð
 5LÐ'¨Ñ1à8˜yÔ8Ø>ð ð  ØBYð ð  Ðð %×(Ò(¨°$Ñ7Ô7ˆ	ØÐØ:ÐFÐFð!ñ GÔFÐFð CÐBÐBÐBÐBÐBàÐ/Ð/Ð/Ø@Ð@Ð@Ð@Ð@Ð@à5O°ZÔ5OØ;ð6ð 6Ø?Oð6ð 6Øfjð6ð 6ð 6Ñ2Ð Ð"2ð .>Ð  Ñ*à=Ð-Ô=Ø7ðð Ø;Kðð ˆIð
 —’˜HÑ%Ô%ˆØˆ>ÝÔFØ Ô'¨Ô)9ðð Ø=Cðð ˆFð ˆsÐ$4À	ÐRXÐdmÐnÑnÔnÐnr5   )NNN)r+   r,   r-   r   r1   Úbase_model_prefixÚ_supports_flash_attnÚ_supports_sdpaÚclassmethodÚstrr   r   rf   r4   r5   r6   r:   r:   Ý   s®   € € € € € € ð ÐÐÑØÐØÐØ€Nàð FJØ>BØ)-ð	Boð Boà8;¸d¹
ðBoð 25°t±ðBoð   $Ñ&ð	Boð 
ðBoð Boð Boñ „[ðBoð Boð Bor5   r:   c            !       óœ  ‡ — e Zd Z	 	 	 	 ddedz  dedz  dedz  dedz  fˆ fd„Ze	 	 	 	 	 	 	 	 	 	 	 	 	 	 ddej	        dz  dej
        dz  d	eeej                          dz  d
ej	        dz  dej        dz  dedz  dej        dz  dej	        dz  dej	        dz  dedz  dedz  dedz  dedz  dedz  deej
                 ez  fd„¦   «         Zˆ xZS )ÚRagModelNr;   rV   rW   r?   c                 óz  •— |€|�|€
J d¦   «         ‚|€t          j        |j        |j        fi |¤Ž}n*t          || j        ¦  «        sJ d|› d| j        › �¦   «         ‚t          ¦   «                              |¦  «         |€ ddlm} | 	                    |j
        ¦  «        }|€ ddlm} | 	                    |j        ¦  «        }|| _        | j        �<t          |t          ¦  «        s J dt          | j        ¦  «        › d	�¦   «         ‚|| _        || _
        || _        d| _        d
| _        |                      ¦   «          dS )áÉ  
        question_encoder (`PreTrainedModel`, *optional*):
            The model responsible for encoding the question into hidden states for retrieval.
        generator (`PreTrainedModel`, *optional*):
            The model responsible for generating text based on retrieved documents.
        retriever (`RagRetriever`, *optional*):
            The component responsible for retrieving documents from a knowledge base given the encoded question.
        NzQEither a configuration or an question_encoder and a generator has to be provided.zconfig: z has to be of type rO   rP   rT   z`self.retriever` is of type z&, but should be of type `RagRetriever`F)r   r^   r;   Ú
isinstanceÚconfig_classÚsuperÚ__init__rZ   rQ   Úfrom_configrV   rU   rW   r?   r   ÚtypeÚctx_encoderÚcontext_encoder_trainingÚ	post_init)	Úselfr;   rV   rW   r?   r`   rQ   rU   Ú	__class__s	           €r6   rs   zRagModel.__init__u  s’  ø€ ð  Ð!Ð&6Ð&BÀyÐG\ÐG\Ø_ñ H]ÔG\Ð]ð ˆ>ÝÔFØ Ô'¨Ô)9ðð Ø=Cðð ˆFˆFõ ˜f dÔ&7Ñ8Ô8ÐsÐsÐ:sÀVÐ:sÐ:sÐ`dÔ`qÐ:sÐ:sÑsÔsÐ8Ý‰Œ×Ò˜Ñ Ô Ð ØÐ#Ø6Ð6Ð6Ð6Ð6Ð6à(×4Ò4°VÔ5LÑMÔMÐàÐØBÐBÐBÐBÐBÐBà-×9Ò9¸&Ô:JÑKÔKˆIà"ˆŒØŒ>Ð%Ý˜i­Ñ6Ô6ð ð Øk­t°D´NÑ/CÔ/CÐkÐkÐkñô Ð6ð 'ˆDŒNà 0ˆÔØ"ˆŒàˆÔØ(-ˆÔ%à�ŠÑÔÐÐÐr5   Ú	input_idsÚattention_maskÚencoder_outputsÚdecoder_input_idsÚdecoder_attention_maskr   r   r    r!   Ú	use_cacheÚoutput_attentionsÚoutput_hidden_statesÚoutput_retrievedÚn_docsr@   c                 ó  — |�|n| j         j        }|
�|
n| j         j        }
|�|n| j         j        }|�|n| j         j        }|�|n| j         j        }| j        duo|du p|	du p|du o|du }|�€�|�rf|                      ||d¬¦  «        }|d         }|                      ||                     ¦   «          	                    dt          j        ¬¦  «                             ¦   «         t          | j        j         dd¦  «        |d¬	¦  «        }| j        �r|d
         |d         |d         |d         |d         |d         f\  }}	}}}}| 	                    |¦  «        }|	 	                    |¦  «        }	| 	                    |¦  «        }| 	                    |¦  «        }|                      ||d¬¦  «        j        }|                     d||j        d         ¦  «        }t          j        |                     d¦  «        |                     dd¦  «        ¦  «                             d¦  «        }nÖ|d
         |d         |d         |d         f\  }}	}}| 	                    |¦  «        }| 	                    |¦  «        }|	 	                    |¦  «        }	t          j        |                     d¦  «        |                     dd¦  «        ¦  «                             d¦  «        }n$|€
J d¦   «         ‚|	€
J d¦   «         ‚|€
J d¦   «         ‚|€
J d¦   «         ‚|j        d         |z  dk    sJ d|› d|j        d         › d�¦   «         ‚|�|                     |d¬¦  «        }|�|                     |d¬¦  «        }|                      ||	|||||
|d¬¦	  «	        }|sd}d}d}d}d}n|j        }|j        }|r|sd}d}	d}d}t7          d*i d|j        “d|“d|j        “d
|“d|	“d|“d |“d!|“d"|“d#|“d$|j        “d%|j        “d&|j         “d'|j!        “d(|j"        “d)|j#        “ŽS )+ay  
        input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
            Indices of input sequence tokens in the vocabulary. [`RagConfig`], used to initialize the model, specifies
            which generator to use, it also specifies a compatible generator tokenizer. Use that tokenizer class to
            obtain the indices.

            [What are input IDs?](../glossary#input-ids)
        encoder_outputs (`tuple(tuple(torch.FloatTensor)`, *optional*)
            Tuple consists of (`generator_enc_last_hidden_state`, *optional*: `generator_enc_hidden_states`,
            *optional*: `generator_enc_attentions`). `generator_enc_last_hidden_state` of shape `(batch_size, n_docs *
            sequence_length, hidden_size)` is a sequence of hidden-states at the output of the last layer of the
            generator's encoder.

            Used by the ([`RagModel`]) model during decoding.
        decoder_input_ids (`torch.LongTensor` of shape `(batch_size, target_sequence_length)`, *optional*):
            Provide for generation tasks. `None` by default, construct as per instructions for the generator model
            you're using with your RAG instance.
        decoder_attention_mask (`torch.BoolTensor` of shape `(batch_size,  target_sequence_length)`, *optional*):
            Default behavior: generate a tensor that ignores pad tokens in `decoder_input_ids`. Causal mask will also
            be used by default.
        doc_scores (`torch.FloatTensor` of shape `(batch_size, config.n_docs)`):
            Score between each retrieved document embeddings (see `retrieved_doc_embeds`) and
            `question_encoder_last_hidden_state`. If the model has is not initialized with a `retriever` `doc_scores`
            has to be provided to the forward pass. `doc_scores` can be computed via
            `question_encoder_last_hidden_state` and `retrieved_doc_embeds`, see examples for more information.
        context_input_ids (`torch.LongTensor` of shape `(batch_size * config.n_docs, config.max_combined_length)`, *optional*, returned when *output_retrieved=True*):
            Input IDs post-processed from the retrieved documents and the question encoder `input_ids` by the
            retriever. If the model was not initialized with a `retriever` ``context_input_ids` has to be provided to
            the forward pass. `context_input_ids` are returned by [`~RagRetriever.__call__`].
        context_attention_mask (`torch.LongTensor` of shape `(batch_size * config.n_docs, config.max_combined_length)`,*optional*, returned when *output_retrieved=True*):
            Attention mask post-processed from the retrieved documents and the question encoder `input_ids` by the
            retriever. If the model has is not initialized with a `retriever` `context_attention_mask` has to be
            provided to the forward pass. `context_attention_mask` are returned by [`~RagRetriever.__call__`].
        output_retrieved (`bool`, *optional*):
            Whether or not to return the `retrieved_doc_embeds`, `retrieved_doc_ids`, `context_input_ids` and
            `context_attention_mask`. See returned tensors for more detail.
        n_docs (`int`, *optional*):
            The number of documents to retrieve.

        Example:

        ```python
        >>> from transformers import AutoTokenizer, RagRetriever, RagModel
        >>> import torch

        >>> tokenizer = AutoTokenizer.from_pretrained("facebook/rag-token-base")
        >>> retriever = RagRetriever.from_pretrained(
        ...     "facebook/rag-token-base", index_name="exact", use_dummy_dataset=True
        ... )
        >>> # initialize with RagRetriever to do everything in one forward call
        >>> model = RagModel.from_pretrained("facebook/rag-token-base", retriever=retriever)

        >>> inputs = tokenizer("How many people live in Paris?", return_tensors="pt")
        >>> outputs = model(input_ids=inputs["input_ids"])
        ```NT)r|   Úreturn_dictr   Úcpu©ÚdeviceÚdtypeÚprefixÚpt©r‹   r„   Úreturn_tensorsr    r!   r   Útokenized_doc_idsÚtokenized_doc_attention_maskÚdoc_idséÿÿÿÿr   rO   z˜Make sure that `context_input_ids` are passed, if no `retriever` is set. Alternatively, you can set a retriever using the `set_retriever(...)` function.z�Make sure that `context_attention_mask` are passed, if no `retriever` is set. Alternatively, you can set a retriever using the `set_retriever(...)` function.z‘Make sure that `doc_scores` are passed, if no `retriever` is set. Alternatively, you can set a retriever using the `set_retriever(...)` function.z^Make sure that `doc_scores` are passed when passing `encoder_outputs` to the forward function.úM The first dimension of `context_input_ids` should be a multiple of `n_docs`=ú	, but is ú.©Údim)	r{   r|   r}   r~   r   r   r€   r�   r†   ©Nr   r   r   r   r"   r#   r$   r%   r&   r'   r(   r)   r*   r4   )$r;   r„   r€   r�   r‚   rƒ   r?   rV   ÚdetachÚtor/   Úfloat32ÚnumpyÚgetattrrW   rw   rv   Úpooler_outputÚviewÚshapeÚbmmÚ	unsqueezeÚ	transposeÚsqueezeÚrepeat_interleaveÚhidden_statesÚ
attentionsr8   r   r   Úencoder_last_hidden_stateÚencoder_hidden_statesÚencoder_attentionsÚdecoder_hidden_statesÚdecoder_attentionsÚcross_attentions)ry   r{   r|   r}   r~   r   r   r   r    r!   r€   r�   r‚   rƒ   r„   r`   Úhas_to_retrieveÚquestion_enc_outputsr"   Úretriever_outputsr   Úretrieved_doc_input_idsÚretrieved_doc_attention_maskr   Úgen_outputsr#   r$   s                              r6   ÚforwardzRagModel.forward©  s
  € ðT "Ð-��°4´;Ô3EˆØ!*Ð!6�I�I¸D¼KÔ<Qˆ	Ø1BÐ1NÐ-Ð-ÐTXÔT_ÔTqÐà$8Ð$DÐ Ð È$Ì+ÔJjð 	ð 0@Ð/KÐ+Ð+ÐQUÔQ\ÔQmÐð ŒN $Ð&ð (Ø" dÐ*ÐbÐ.DÈÐ.LÐbÐPZÐ^bÐPbð(à 4Ð'ð 	ð Ñ"Øñ LØ'+×'<Ò'<Ø¨nÈ$ð (=ñ (ô (Ð$ð 6JÈ!Ô5LÐ2à$(§N¢NØØ6×=Ò=Ñ?Ô?×BÒBÈ%ÕW\ÔWdÐBÑeÔe×kÒkÑmÔmÝ" 4¤>Ô#8¸(ÀDÑIÔIØ!Ø#'ð %3ñ %ô %Ð!ð Ô0ñ 2!ð *Ð*=Ô>Ø)Ð*BÔCØ)Ð*@ÔAØ)Ð*=Ô>Ø)Ð*HÔIØ)¨)Ô4ðñØ)Ø.Ø,Ø/Ø4Ø)ð ):×(<Ò(<¸YÑ(GÔ(GÐ%Ø-C×-FÒ-FÀyÑ-QÔ-QÐ*à.E×.HÒ.HÈÑ.SÔ.SÐ+Ø3O×3RÒ3RÐS\Ñ3]Ô3]Ð0Ø+/×+;Ò+;Ø/Ð@\Ðjnð ,<ñ ,ô ,ä#ð )ð ,@×+DÒ+DØ˜FÐ$FÔ$LÈQÔ$Oñ,ô ,Ð(õ
 "'¤Ø:×DÒDÀQÑGÔGÐI]×IgÒIgÐhiÐklÑImÔImñ"ô "ç’g˜a‘j”jð �Jð *Ð*=Ô>Ø)Ð*BÔCØ)Ð*@ÔAØ)¨)Ô4ð	jÑfÐ%Ð'=Ð?SÐUfð ,@×+BÒ+BÐCeÑ+fÔ+fÐ(Ø(9×(<Ò(<¸YÑ(GÔ(GÐ%Ø-C×-FÒ-FÀyÑ-QÔ-QÐ*õ "'¤Ø:×DÒDÀQÑGÔGÐI]×IgÒIgÐhiÐklÑImÔImñ"ô "ç’g˜a‘j”jð �Jð )Ð4Ð4ðPñ 5Ô4Ð4ð .Ð9Ð9ðTñ :Ô9Ð9ð "Ð-Ð-ðJñ .Ô-Ð-ð
 Ð%Ð%Ølñ &Ô%Ð%ð Ô  Ô# fÑ,°Ò2Ð2Ð2ð.Ð\bð .ð .Ø!Ô'¨Ô*ð.ð .ð .ñ 3Ô2Ð2ð Ð(Ø 1× CÒ CÀFÐPQÐ CÑ RÔ RÐà!Ð-Ø%;×%MÒ%MÈfÐZ[Ð%MÑ%\Ô%\Ð"à—n’nØ'Ø1Ø+Ø/Ø#9Ø+ØØ/Øð %ñ 

ô 

ˆð ð 	FØ15Ð.Ø)-Ð&Ø&*Ð#Ø#'Ð Ø $ÐÐà)=Ô)KÐ&Ø&:Ô&EÐ#àð 	%Ð&6ð 	%à 'ÐØ%)Ð"Ø#'Ð Ø $Ðå!ð 
ð 
ð 
ØÔ%Ð%ð
à!�zð
ð (Ô7Ð7ð
ð 0Ð/ð	
ð
 $:Ð#9ð
ð "6Ð!5ð
ð 0Ð/ð
ð 0RÐ/Qð
ð (BÐ'Að
ð %<Ð$;ð
ð -8Ô,QÐ,Qð
ð )4Ô(IÐ(Ið
ð &1Ô%CÐ%Cð
ð )4Ô(IÐ(Ið
ð &1Ô%CÐ%Cð
ð  (3Ô'CÐ'Cð!
ð 	
r5   ©NNNN)NNNNNNNNNNNNNN)r+   r,   r-   r	   r   r   rs   r   r/   r2   ÚTensorr3   r0   Ú
BoolTensorr   ÚboolÚintr8   r´   Ú__classcell__©rz   s   @r6   rm   rm   s  sö  ø€ € € € € ð +/Ø37Ø,0Ø)-ð2ð 2à  4Ñ'ð2ð *¨DÑ0ð2ð # TÑ)ð	2ð
   $Ñ&ð2ð 2ð 2ð 2ð 2ð 2ðh ð .2Ø.2ØBFØ59Ø:>Ø(,Ø/3Ø59Ø:>Ø!%Ø)-Ø,0Ø(,Ø!ðe
ð e
àÔ# dÑ*ðe
ð œ tÑ+ðe
ð ˜u UÔ%6Ô7Ô8¸4Ñ?ð	e
ð
 !Ô+¨dÑ2ðe
ð !&Ô 0°4Ñ 7ðe
ð  ™ðe
ð Ô%¨Ñ,ðe
ð !Ô+¨dÑ2ðe
ð !&Ô 0°4Ñ 7ðe
ð ˜$‘;ðe
ð   $™;ðe
ð # T™kðe
ð  ™+ðe
ð �d‘
ðe
ð" 
ˆuŒ|Ô	Ð1Ñ	1ð#e
ð e
ð e
ñ „^ðe
ð e
ð e
ð e
ð e
r5   rm   zu
    A RAG-sequence model implementation. It performs RAG-sequence specific marginalization in the forward pass.
    c            &       ó  ‡ — e Zd Z	 	 	 	 d(dedz  dedz  dedz  dedz  fˆ fd„Zdefd„Zdefd	„Ze		 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d)d
e
j        dz  de
j        dz  deee
j                          dz  de
j        dz  de
j        dz  dedz  de
j        dz  de
j        dz  de
j        dz  dedz  dedz  dedz  dedz  dedz  dedz  de
j        dz  dedz  def$d„¦   «         Zed„ ¦   «         Zed„ ¦   «         Zed„ ¦   «         Z e
j        ¦   «         	 	 	 	 	 	 	 	 	 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dz  d!edz  d"edz  dedz  de
j        fd#„¦   «         Z	 d+d&„Zed'„ ¦   «         Zˆ xZS ),ÚRagSequenceForGenerationNr;   rV   rW   r?   c                 ó   •— |€|�|€
J d¦   «         ‚|€t          j        |j        |j        fi |¤Ž}t          ¦   «                              |¦  «         t          ||||¬¦  «        | _        |                      ¦   «          dS ©ro   NzHEither a configuration or an encoder and a generator has to be provided.)r;   rV   rW   r?   ©r   r^   r;   rr   rs   rm   r<   rx   ©ry   r;   rV   rW   r?   r`   rz   s         €r6   rs   z!RagSequenceForGeneration.__init__˜  s§   ø€ ð  Ð!Ð&6Ð&BÀyÐG\ÐG\ØVñ H]ÔG\Ð]ð ˆ>ÝÔFØ Ô'¨Ô)9ðð Ø=Cðð ˆFõ 	‰Œ×Ò˜Ñ Ô Ð õ  6Ð<LÐXaÐmvÐwÑwÔwˆŒà�ŠÑÔÐÐÐr5   c                 ó   — || j         _        d S r˜   ©r<   r?   ©ry   r?   s     r6   Úset_retrieverz&RagSequenceForGeneration.set_retriever·  ó   € Ø&ˆŒÔÐÐr5   rv   c                 ó6   — d| j         _        || j         _        d S ©NT©r<   rw   rv   ©ry   rv   s     r6   Ú set_context_encoder_for_trainingz9RagSequenceForGeneration.set_context_encoder_for_trainingº  ó   € Ø,0ˆŒÔ)Ø*ˆŒÔÐÐr5   r{   r|   r}   r~   r   r   r    r!   r   r€   r�   r‚   rƒ   Úexclude_bos_scoreÚreduce_lossÚlabelsr„   r@   c                 ó:  — |�|n| j         j        }|�|n| j         j        }|�|n| j         j        }|�|€|}d}
|                      ||||||||	||
||||¬¦  «        }d}|�0|                      |j        |j        ||| j         j        ||¬¦  «        }t          di d|“d|j        “d|j        “d|j
        “d	|j        “d
|j        “d|j        “d|j        “d|j        “d|j        “d|j        “d|j        “d|j        “d|j        “d|j        “d|j        “d|j        “ŽS )a3  
        input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
            Indices of input sequence tokens in the vocabulary. [`RagConfig`], used to initialize the model, specifies
            which generator to use, it also specifies a compatible generator tokenizer. Use that tokenizer class to
            obtain the indices.

            [What are input IDs?](../glossary#input-ids)
        encoder_outputs (`tuple(tuple(torch.FloatTensor)`, *optional*)
            Tuple consists of (`generator_enc_last_hidden_state`, *optional*: `generator_enc_hidden_states`,
            *optional*: `generator_enc_attentions`). `generator_enc_last_hidden_state` of shape `(batch_size, n_docs *
            sequence_length, hidden_size)` is a sequence of hidden-states at the output of the last layer of the
            generator's encoder.

            Used by the ([`RagModel`]) model during decoding.
        decoder_input_ids (`torch.LongTensor` of shape `(batch_size, target_sequence_length)`, *optional*):
            Provide for generation tasks. `None` by default, construct as per instructions for the generator model
            you're using with your RAG instance.
        decoder_attention_mask (`torch.BoolTensor` of shape `(batch_size,  target_sequence_length)`, *optional*):
            Default behavior: generate a tensor that ignores pad tokens in `decoder_input_ids`. Causal mask will also
            be used by default.
        context_input_ids (`torch.LongTensor` of shape `(batch_size * config.n_docs, config.max_combined_length)`, *optional*, returned when *output_retrieved=True*):
            Input IDs post-processed from the retrieved documents and the question encoder `input_ids` by the
            retriever. If the model was not initialized with a `retriever` ``context_input_ids` has to be provided to
            the forward pass. `context_input_ids` are returned by [`~RagRetriever.__call__`].
        context_attention_mask (`torch.LongTensor` of shape `(batch_size * config.n_docs, config.max_combined_length)`,*optional*, returned when *output_retrieved=True*):
            Attention mask post-processed from the retrieved documents and the question encoder `input_ids` by the
            retriever. If the model has is not initialized with a `retriever` `context_attention_mask` has to be
            provided to the forward pass. `context_attention_mask` are returned by [`~RagRetriever.__call__`].
        doc_scores (`torch.FloatTensor` of shape `(batch_size, config.n_docs)`):
            Score between each retrieved document embeddings (see `retrieved_doc_embeds`) and
            `question_encoder_last_hidden_state`. If the model has is not initialized with a `retriever` `doc_scores`
            has to be provided to the forward pass. `doc_scores` can be computed via
            `question_encoder_last_hidden_state` and `retrieved_doc_embeds`, see examples for more information.
        output_retrieved (`bool`, *optional*):
            Whether or not to return the `retrieved_doc_embeds`, `retrieved_doc_ids`, `context_input_ids` and
            `context_attention_mask`. See returned tensors for more detail.
        exclude_bos_score (`bool`, *optional*):
            Only relevant if `labels` is passed. If `True`, the score of the BOS token is disregarded when computing
            the loss.
        reduce_loss (`bool`, *optional*):
            Only relevant if `labels` is passed. If `True`, the NLL loss is reduced using the `torch.Tensor.sum`
            operation.
        n_docs (`int`, *optional*):
            The number of documents to retrieve.

        Example:

        ```python
        >>> from transformers import AutoTokenizer, RagRetriever, RagSequenceForGeneration
        >>> import torch

        >>> tokenizer = AutoTokenizer.from_pretrained("facebook/rag-sequence-nq")
        >>> retriever = RagRetriever.from_pretrained(
        ...     "facebook/rag-sequence-nq", index_name="exact", use_dummy_dataset=True
        ... )
        >>> # initialize with RagRetriever to do everything in one forward call
        >>> model = RagSequenceForGeneration.from_pretrained("facebook/rag-token-nq", retriever=retriever)

        >>> inputs = tokenizer("How many people live in Paris?", return_tensors="pt")
        >>> targets = tokenizer(text_target="In Paris, there are 10 million people.", return_tensors="pt")
        >>> input_ids = inputs["input_ids"]
        >>> labels = targets["input_ids"]
        >>> outputs = model(input_ids=input_ids, labels=labels)

        >>> # or use retriever separately
        >>> model = RagSequenceForGeneration.from_pretrained("facebook/rag-sequence-nq", use_dummy_dataset=True)
        >>> # 1. Encode
        >>> question_hidden_states = model.question_encoder(input_ids)[0]
        >>> # 2. Retrieve
        >>> docs_dict = retriever(input_ids.numpy(), question_hidden_states.detach().numpy(), return_tensors="pt")
        >>> doc_scores = torch.bmm(
        ...     question_hidden_states.unsqueeze(1), docs_dict["retrieved_doc_embeds"].float().transpose(1, 2)
        ... ).squeeze(1)
        >>> # 3. Forward to generator
        >>> outputs = model(
        ...     context_input_ids=docs_dict["context_input_ids"],
        ...     context_attention_mask=docs_dict["context_attention_mask"],
        ...     doc_scores=doc_scores,
        ...     decoder_input_ids=labels,
        ... )
        ```NF©r{   r|   r}   r~   r   r    r!   r   r   r€   r�   r‚   rƒ   r„   )rÎ   ÚepsilonrÍ   r„   r   r   r   r   r    r!   r   r   r"   r#   r$   r%   r&   r'   r(   r)   r*   r4   )r;   r„   rÍ   rÎ   r<   Úget_nllr   r   Úlabel_smoothingr   r   r    r!   r   r   r"   r#   r$   r%   r&   r'   r(   r)   r*   )ry   r{   r|   r}   r~   r   r   r    r!   r   r€   r�   r‚   rƒ   rÍ   rÎ   rÏ   r„   r`   Úoutputsr   s                        r6   r´   z RagSequenceForGeneration.forward¾  sÿ  € ðN "Ð-��°4´;Ô3EˆØ1BÐ1NÐ-Ð-ÐTXÔT_ÔTqÐØ%0Ð%<�k�kÀ$Ä+ÔBYˆàÐØ Ð(Ø$*Ð!ØˆIà—(’(ØØ)Ø+Ø/Ø#9Ø/Ø#9Ø!Ø+ØØ/Ø!5Ø-Øð ñ 
ô 
ˆð" ˆØÐØ—<’<Ø”ØÔ"Ø!Ø'ØœÔ3Ø"3Øð  ñ ô ˆDõ (ð 
ð 
ð 
Ø�ð
à”>�>ð
ð Ô)Ð)ð
ð $Ô3Ð3ð	
ð
 &Ô7Ð7ð
ð $+Ô#AÐ#Að
ð ")Ô!=Ð!=ð
ð &Ô7Ð7ð
ð 07Ô/YÐ/Yð
ð (/Ô'IÐ'Ið
ð %,Ô$CÐ$Cð
ð -4Ô,SÐ,Sð
ð )0Ô(KÐ(Kð
ð &-Ô%EÐ%Eð
ð )0Ô(KÐ(Kð
ð  &-Ô%EÐ%Eð!
ð" (/Ô'IÐ'Ið#
ð 	
r5   c                 ó   — | j         j        S r˜   rÃ   ©ry   s    r6   r?   z"RagSequenceForGeneration.retriever_  ó   € àŒxÔ!Ð!r5   c                 ó   — | j         j        S r˜   ©r<   rW   r×   s    r6   rW   z"RagSequenceForGeneration.generatorc  rØ   r5   c                 ó   — | j         j        S r˜   ©r<   rV   r×   s    r6   rV   z)RagSequenceForGeneration.question_encoderg  ó   € àŒxÔ(Ð(r5   Údo_deduplicationÚnum_return_sequencesÚ	num_beamsc
           	      óæ  — |	�|	n| j         j        }	|�|n| j         j        }|�|nt          | j         dd¦  «        }|�|nt          | j         dd¦  «        }|€|€
J d¦   «         ‚| j        �°|€®|                      ||¬¦  «        d         }|                      ||                     ¦   «                              dt          j	        ¬	¦  «         
                    ¦   «         t          | j        j         d
d¦  «        |	d¬¦  «        d         }|                     |¦  «        }g }||
d<   ||
d<   d|
d<   |�|j        d         n|j        d         |	z  }t          |¦  «        D �]r}|||	z  |dz   |	z  …         } | j        j        |fi |
¤Ž}|r=t          j        t!          d„ |D ¦   «                              ¦   «         ¦  «        ¦  «        }|j        d         }|�0|||dz   …                              |d¦  «        } | ||d¬¦  «        }nŽ|€
J d¦   «         ‚|€
J d¦   «         ‚|                     |d¦  «        }|||	z  |dz   |	z  …         }|                     |d¦  «        }|||dz   …dd…f         }|                     |d¦  «        } | ||||d¬¦  «        }|d                               |¦  «        d         }|                     ||         ¦  «         �Œt|                      || j         j        j        ¬¦  «        S )a  
        Implements RAG sequence "thorough" decoding. Read the [`~generation.GenerationMixin.generate`]` documentation
        for more information on how to set other generate input parameters.

        Args:
            input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
                The sequence used as a prompt for the generation. If `input_ids` is not passed, then
                `context_input_ids` has to be provided.
            attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
                Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:

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

                [What are attention masks?](../glossary#attention-mask)
            context_input_ids (`torch.LongTensor` of shape `(batch_size * config.n_docs, config.max_combined_length)`, *optional*, returned when *output_retrieved=True*):
                Input IDs post-processed from the retrieved documents and the question encoder input_ids by the
                retriever.
            context_attention_mask (`torch.LongTensor` of shape `(batch_size * config.n_docs, config.max_combined_length)`, *optional*, returned when *output_retrieved=True*):
                Attention mask post-processed from the retrieved documents and the question encoder `input_ids` by the
                retriever.

                If the model is not initialized with a `retriever` or `input_ids` is not given, `context_input_ids` and
                `context_attention_mask` have to be provided to the forward pass. They are returned by
                [`~RagRetriever.__call__`].
            doc_scores (`torch.FloatTensor` of shape `(batch_size, config.n_docs)`):
                Score between each retrieved document embeddings (see `retrieved_doc_embeds`) and
                `question_encoder_last_hidden_state`.

                If the model is not initialized with a `retriever` or `input_ids` is not given, `doc_scores` has to be
                provided to the forward pass. `doc_scores` are returned by [`~RagRetriever.__call__`].
            do_deduplication (`bool`, *optional*):
                Whether or not to deduplicate the generations from different context documents for a given input. Has
                to be set to `False` if used while training with distributed backend.
            num_return_sequences(`int`, *optional*, defaults to 1):
                The number of independently computed returned sequences for each element in the batch. Note that this
                is not the value we pass to the `generator`'s `[`~generation.GenerationMixin.generate`]` function,
                where we set `num_return_sequences` to `num_beams`.
            num_beams (`int`, *optional*, defaults to 1):
                Number of beams for beam search. 1 means no beam search.
            n_docs (`int`, *optional*, defaults to `config.n_docs`)
                Number of documents to retrieve and/or number of documents for which to generate an answer.
            kwargs (`dict[str, Any]`, *optional*):
                Additional kwargs will be passed to [`~generation.GenerationMixin.generate`].

        Return:
            `torch.LongTensor` of shape `(batch_size * num_return_sequences, sequence_length)`: The generated
            sequences. The second dimension (sequence length) is either equal to `max_length` or shorter if all batches
            finished early due to the `eos_token_id`.
        Nrß   r   rà   z= At least one of input_ids or context_input_ids must be given©r|   r   r‡   rˆ   r‹   rŒ   r�   r    r|   c                 óR   — i | ]$}t          |                     ¦   «         ¦  «        |“Œ%S r4   )rk   Útolist)rH   Úks     r6   rK   z5RagSequenceForGeneration.generate.<locals>.<dictcomp>×  s(   € Ð4bÐ4bÐ4bÈAµS¸¿º¹¼±_´_ÀaÐ4bÐ4bÐ4br5   T)rÏ   rÍ   z�Make sure that `context_attention_mask` are passed, if no `input_ids` is set. Alternatively, you can set a retriever using the `set_retriever(...)` function.z‘Make sure that `doc_scores` are passed, if no `input_ids` is set. Alternatively, you can set a retriever using the `set_retriever(...)` function.)r    r!   r   rÏ   rÍ   r   )Úpad_token_id)r;   r„   rÞ   r�   r?   rV   r™   rš   r/   r›   rœ   rW   r    ÚrangeÚgenerateÚstackÚlistÚvaluesÚrepeatÚtopkÚappendÚ_cat_and_padræ   )ry   r{   r|   r    r!   r   rÞ   rß   rà   r„   Úmodel_kwargsÚnum_doc_return_sequencesÚquestion_hidden_statesÚhyposÚ
batch_sizeÚindexÚgenerator_input_idsÚoutput_sequencesÚnum_candidatesÚnew_input_idsrÕ   Úindividual_input_idsÚindividual_attention_maskÚindividual_doc_scoresÚtop_cand_indss                            r6   rè   z!RagSequenceForGeneration.generatek  sÎ  € ðB "Ð-��°4´;Ô3EˆØ/?Ð/KÐ+Ð+ÐQUÔQ\ÔQmÐð $Ð/ð !Ð å˜œÐ&<¸aÑ@Ô@ð 	!ð
 "+Ð!6�I�I½GÀDÄKÐQ\Ð^_Ñ<`Ô<`ˆ	àÐ$Ð(9Ð(EÐ(EØKñ )FÔ(EÐEð Œ>Ð%Ð*;Ð*CØ%)×%:Ò%:¸9ÐUcÐ%:Ñ%dÔ%dÐefÔ%gÐ"Ø $§¢ØØ&×-Ò-Ñ/Ô/×2Ò2¸%ÅuÄ}Ð2ÑUÔU×[Ò[Ñ]Ô]Ý˜tœ~Ô4°hÀÑEÔEØØ#ð !/ñ !ô !ð "ô!#Ðð !2× 4Ò 4°YÑ ?Ô ?ÐàˆØ$-ˆ�[Ñ!Ø/8ˆÐ+Ñ,Ø)-ˆÐ%Ñ&à+4Ð+@�Y”_ QÔ'Ð'ÐFWÔF]Ð^_ÔF`ÐdjÑFjˆ
å˜:Ñ&Ô&ð 3	:ñ 3	:ˆEà"3°E¸F±NÀeÈaÁiÐSYÑEYÐ4YÔ"ZÐà6˜tœ~Ô6Ø#ð ð  àð ð  Ðð  ð nå#(¤;­tÐ4bÐ4bÐQaÐ4bÑ4bÔ4b×4iÒ4iÑ4kÔ4kÑ/lÔ/lÑ#mÔ#mÐ à-Ô3ØôˆNð
 Ð$Ø )¨%°%¸!±)Ð*;Ô <× CÒ CÀNÐTUÑ VÔ V�Ø˜$˜}Ð5EÐY]Ð^Ñ^Ô^��à-Ð9Ð9ðTñ :Ô9Ð9ð "Ð-Ð-ðJñ .Ô-Ð-ð
 (;×'AÒ'AØ" Añ(ô (Ð$ð -CÀ5È6Á>ÐUZÐ]^ÑU^ÐbhÑThÐChÔ,iÐ)Ø,E×,LÒ,LÈ^Ð]^Ñ,_Ô,_Ð)à(2°5¸EÀA¹IÐ3FÈÈÈÐ3IÔ(JÐ%Ø(=×(DÒ(DÀ^ÐUVÑ(WÔ(WÐ%à˜$Ø&:Ø+DØ4Ø+Ø&*ðñ ô �ð & fœoÐ-×3Ò3Ð4LÑMÔMÈaÔPˆMð �LŠLÐ)¨-Ô8Ñ9Ô9Ð9Ñ9à× Ò  °T´[Ô5JÔ5WÐ ÑXÔXÐXr5   Fç        c                 ó‚  ‡ ‡— t          j        ‰d d …dd …f         ‰                     ‰j        d         d¦  «                             ‰ j        j        j        ¦  «        gd¦  «        Š|�|n‰ j        j        }‰ j        j	        p‰ j        j        j	        }|d uo0‰d d …df          
                    |¦  «                             ¦   «         }	ˆ ˆfd„}
t          j                             |d¬¦  «                             |j        d         |z  |d|                     d¦  «        ¦  «        }t          j                             |d¬¦  «                             d¦  «                             d¦  «        }|d d …d d …d d…d d …f         }|d d …d d …dd…d d …f         }|d d …d d …dd …d d …f         }t          j        |||z   |gd¬¦  «        }‰                     d¦  «                             d¦  «                             d|dd¦  «        Š‰                     ¦   «         |                     ¦   «         k    sJ ‚|                     d‰¬¦  «        }|                     dd¬	¦  «        } |
||¦  «        \  }}|r&|	r$|d d …d d …dd …f                              d¦  «        n|                     d¦  «        }|                     d¦  «        }|                     d¦  «        }|                     d¦  «        }| }| }|r(|                     ¦   «         }|                     ¦   «         }||                     d¦  «        z  }d
|z
  |z  ||z  z   }|S )Nr   r   c                 ó   •— ‰                      ‰j        j        j        ¦  «        }|                     ¦   «         r,|                      |d¦  «         |                     |d¦  «         |                      d¦  «        |                     d¦  «        fS ©Nrþ   r’   ©Úeqr;   rW   ræ   ÚanyÚmasked_fill_r¤   ©ÚllÚ
smooth_objÚpad_maskry   Útargets      €€r6   Ú
_mask_padsz4RagSequenceForGeneration.get_nll.<locals>._mask_pads  óy   ø€ Ø—y’y ¤Ô!6Ô!CÑDÔDˆHØ�|Š|‰~Œ~ð 7Ø—’ ¨#Ñ.Ô.Ð.Ø×'Ò'¨°#Ñ6Ô6Ð6Ø—:’:˜b‘>”> :×#5Ò#5°bÑ#9Ô#9Ð9Ð9r5   r’   r–   rO   ©r—   rõ   T©r—   Úkeepdimç      ð?)r/   ÚcatÚnewr    Úfill_r;   rW   ræ   r„   Úbos_token_idr  Úallr   Ú
functionalÚlog_softmaxrŸ   Úsizer¢   rì   r—   ÚgatherÚsumÚ	logsumexp)ry   Ú
seq_logitsr   r
  rÎ   rÒ   rÍ   r„   r  Úuse_bosr  Úseq_logprobsÚdoc_logprobsÚfirst_token_scoresÚsecond_token_scoresÚ	remainderÚrag_logprobsr  r  Únll_lossÚsmooth_lossÚeps_ir   s   `  `                   r6   rÓ   z RagSequenceForGeneration.get_nll  so  øø€ õ ”Ø�A�A�A�q�r�r�EŒ]˜FŸJšJ v¤|°A¤¸Ñ:Ô:×@Ò@ÀÄÔAVÔAcÑdÔdÐeÐghñ
ô 
ˆð "Ð-��°4´;Ô3Eˆð ”{Ô/ÐU°4´;Ô3HÔ3UˆØ dÐ*ÐR¨v°a°a°a¸°d¬|¯ª¸|Ñ/LÔ/L×/PÒ/PÑ/RÔ/Rˆð	:ð 	:ð 	:ð 	:ð 	:ð 	:õ ”}×0Ò0°ÀÐ0ÑDÔD×IÒIØÔ˜QÔ 6Ñ)¨6°2°z·²ÀrÑ7JÔ7Jñ
ô 
ˆõ ”}×0Ò0°ÀÐ0ÑCÔC×MÒMÈbÑQÔQ×[Ò[Ð\^Ñ_Ô_ˆð *¨!¨!¨!¨Q¨Q¨Q°°°°A°A°A¨+Ô6ÐØ*¨1¨1¨1¨a¨a¨a°°1°°a°a°a¨<Ô8ÐØ     A A A q r r¨1¨1¨1 Ô-ˆ	Ý”yÐ"4Ð6IÈLÑ6XÐZcÐ!dÐjkÐlÑlÔlˆð ×!Ò! !Ñ$Ô$×.Ò.¨rÑ2Ô2×9Ò9¸!¸VÀQÈÑJÔJˆØ�zŠz‰|Œ|˜|×/Ò/Ñ1Ô1Ò1Ð1Ð1Ð1à× Ò  R¨vÐ Ñ6Ô6ˆØ!×%Ò%¨"°dÐ%Ñ;Ô;ˆ
à#˜ B¨
Ñ3Ô3‰ˆˆJð %6ÐP¸'ÐPˆR����1�1�1�a�b�b�Œ\×Ò˜aÑ Ô Ð ÀrÇvÂvÈaÁyÄyˆØ—^’^ AÑ&Ô&ˆ
Ø�\Š\˜!‰_Œ_ˆØ×)Ò)¨!Ñ,Ô,ˆ
à�3ˆØ!�kˆàð 	,Ø—|’|‘~”~ˆHØ%Ÿ/š/Ñ+Ô+ˆKà˜,×+Ò+¨BÑ/Ô/Ñ/ˆØ�g‘ Ñ)¨E°KÑ,?Ñ?ˆØˆr5   c                 ó6  — | d                               t          d„ | D ¦   «         ¦  «        t          d„ | D ¦   «         ¦  «        ¦  «                             |¦  «        }d}| D ]6}|||||j        d         z   …d |j        d         …f<   ||j        d         z  }Œ7|S )Nr   c              3   ó0   K  — | ]}|j         d          V — ŒdS )r   N©r    ©rH   Úts     r6   ú	<genexpr>z8RagSequenceForGeneration._cat_and_pad.<locals>.<genexpr>A  s(   è è € Ð#@Ð#@°1 A¤G¨A¤JÐ#@Ð#@Ð#@Ð#@Ð#@Ð#@r5   c              3   ó0   K  — | ]}|j         d          V — ŒdS )r   Nr)  r*  s     r6   r,  z8RagSequenceForGeneration._cat_and_pad.<locals>.<genexpr>A  s)   è è € ÐEbÐEbÐUVÀaÄgÈaÄjÐEbÐEbÐEbÐEbÐEbÐEbr5   r   )r  r  Úmaxr  r    )Útensorsræ   ÚoutputÚindr+  s        r6   rï   z%RagSequenceForGeneration._cat_and_pad?  s«   € à˜”—’¥Ð#@Ð#@¸Ð#@Ñ#@Ô#@Ñ @Ô @Å#ÐEbÐEbÐZaÐEbÑEbÔEbÑBbÔBbÑcÔc×iÒiÐjvÑwÔwˆØˆØð 	ð 	ˆAØ;<ˆF�3˜˜qœw qœzÑ)Ð)¨<¨Q¬W°Q¬Z¨<Ð7Ñ8Ø�1”7˜1”:ÑˆCˆCØˆr5   rµ   ©NNNNNNNNNNNNNNNNN)	NNNNNNNNN)Frþ   FN)r+   r,   r-   r	   r   r   rs   rÅ   rË   r   r/   r2   r¶   r3   r·   r   r0   r¸   r¹   r   r´   Úpropertyr?   rW   rV   Úno_gradrè   rÓ   Ústaticmethodrï   rº   r»   s   @r6   r½   r½   ’  sâ  ø€ € € € € ð +/Ø37Ø,0Ø)-ðð à  4Ñ'ðð *¨DÑ0ðð # TÑ)ð	ð
   $Ñ&ðð ð ð ð ð ð>' |ð 'ð 'ð 'ð 'ð+¸Oð +ð +ð +ð +ð ð .2Ø.2Ø=AØ59Ø:>Ø(,Ø59Ø:>Ø/3Ø!%Ø)-Ø,0Ø(,Ø)-Ø#'Ø*.Ø!ð%^
ð ^
àÔ# dÑ*ð^
ð œ tÑ+ð^
ð ˜u U¤\Ô2Ô3°dÑ:ð	^
ð
 !Ô+¨dÑ2ð^
ð !&Ô 0°4Ñ 7ð^
ð  ™ð^
ð !Ô+¨dÑ2ð^
ð !&Ô 0°4Ñ 7ð^
ð Ô%¨Ñ,ð^
ð ˜$‘;ð^
ð   $™;ð^
ð # T™kð^
ð  ™+ð^
ð   $™;ð^
ð  ˜D‘[ð!^
ð" Ô  4Ñ'ð#^
ð$ �d‘
ð%^
ð( 
"ð)^
ð ^
ð ^
ñ „^ð^
ð@ ð"ð "ñ „Xð"ð ð"ð "ñ „Xð"ð ð)ð )ñ „Xð)ð €U„]�_„_ð .2Ø26Ø59Ø:>Ø/3Ø(,Ø+/Ø $Ø!ðVYð VYàÔ# dÑ*ðVYð Ô(¨4Ñ/ðVYð !Ô+¨dÑ2ð	VYð
 !&Ô 0°4Ñ 7ðVYð Ô%¨Ñ,ðVYð  ™+ðVYð " D™jðVYð ˜‘:ðVYð �d‘
ðVYð 
Ô	ðVYð VYð VYñ „_ðVYðr osð9ð 9ð 9ð 9ðv ðð ñ „\ðð ð ð ð r5   r½   zo
    A RAG-token model implementation. It performs RAG-token specific marginalization in the forward pass.
    c            &       ó˜  ‡ — e Zd Z	 	 	 	 d0dedz  dedz  dedz  dedz  fˆ fd„Zdefd„Zdefd	„Z	 	 	 	 	 	 d1d
„Z	e
d„ ¦   «         Ze
d„ ¦   «         Ze
d„ ¦   «         Zed„ ¦   «         Zd2d„Ze	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d3dej        dz  dej        dz  deeej                          dz  dej        dz  dej        dz  dedz  dej        dz  dej        dz  dej        dz  dedz  dedz  dedz  dedz  dedz  dedz  dej        dz  d edz  d!ef$d"„¦   «         Z ej        ¦   «         dddddddd e¦   «          e¦   «         f
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!eej        ge"e         f         dz  d%edz  d&edz  d!ej        fd'„¦   «         Z#d(„ Z$d)„ Z%d*„ Z&d+„ Z'd2d,„Z(d4d/„Z)ˆ xZ*S )5ÚRagTokenForGenerationNr;   rV   rW   r?   c                 ó   •— |€|�|€
J d¦   «         ‚|€t          j        |j        |j        fi |¤Ž}t          ¦   «                              |¦  «         t          ||||¬¦  «        | _        |                      ¦   «          dS r¿   rÀ   rÁ   s         €r6   rs   zRagTokenForGeneration.__init__O  s§   ø€ ð  Ð!Ð&6Ð&BÀyÐG\ÐG\ØVñ H]ÔG\Ð]ð ˆ>ÝÔFØ Ô'¨Ô)9ðð Ø=Cðð ˆFõ 	‰Œ×Ò˜Ñ Ô Ð õ  6Ð<LÐXaÐmvÐwÑwÔwˆŒà�ŠÑÔÐÐÐr5   c                 ó   — || j         _        d S r˜   rÃ   rÄ   s     r6   rÅ   z#RagTokenForGeneration.set_retrievero  rÆ   r5   rv   c                 ó6   — d| j         _        || j         _        d S rÈ   rÉ   rÊ   s     r6   rË   z6RagTokenForGeneration.set_context_encoder_for_trainingr  rÌ   r5   c           
      ó:   — |�|d d …dd …f         }d ||||||d|dœ	S )Nr’   T)	r{   r}   r   r!   r~   r   r€   Údo_marginalizer„   r4   )	ry   r~   r   r|   r€   r}   r   r„   r`   s	            r6   Úprepare_inputs_for_generationz3RagTokenForGeneration.prepare_inputs_for_generationv  sM   € ð Ð&à 1°!°!°!°R°S°S°&Ô 9Ðð Ø.Ø$Ø&4Ø!2Ø.Ø"Ø"Øð

ð 

ð 
	
r5   c                 ó   — | j         j        S r˜   rÃ   r×   s    r6   r?   zRagTokenForGeneration.retriever“  rØ   r5   c                 ó   — | j         j        S r˜   rÚ   r×   s    r6   rW   zRagTokenForGeneration.generator—  rØ   r5   c                 ó   — | j         j        S r˜   rÜ   r×   s    r6   rV   z&RagTokenForGeneration.question_encoder›  rÝ   r5   c                 ó
  ‡‡	— d„ Š	d}t          t          | ¦  «        ¦  «        D ]È}t          | t          ¦  «        rsˆ	ˆfd„| j        j        |         j        | j        j        |         j        | j        j        |         j        | j        j        |         j        fD ¦   «         \  }}}}||||f}n8ˆ	ˆfd„| j        |         j        | j        |         j        fD ¦   «         \  }}||f}||fz  }ŒÉ t          | ¦  «        |¦  «        S )zeReorders cache for generation. BART-inspired but we need to take care of the extra dimension for docsc                 óÖ   — | j         d         |j         d         z  } | j        d|g| j         dd …         ¢R Ž } |                      d|¦  «        }  | j        dg| j         dd …         ¢R Ž }|S )Nr   r’   r   rO   )r    rŸ   Úindex_select)r¦   Ú	new_orderr„   Úresults       r6   Ú_reorder_stackedz>RagTokenForGeneration._reorder_cache.<locals>._reorder_stacked£  s†   € Ø"Ô(¨Ô+¨y¬¸qÔ/AÑAˆFØ.˜MÔ.¨r°6ÐT¸MÔ<OÐPQÐPRÐPRÔ<SÐTÐTÐTˆMØ)×6Ò6°q¸)ÑDÔDˆMØ'�]Ô'¨ÐE¨]Ô-@ÀÀÀÔ-DÐEÐEÐEˆFØˆMr5   r4   c              3   ó`   •K  — | ](} ‰|‰                      |j        ¦  «        ¦  «        V — Œ)d S r˜   ©rš   r‰   ©rH   ÚxrF  Úbeam_idxs     €€r6   r,  z7RagTokenForGeneration._reorder_cache.<locals>.<genexpr>­  sZ   øè è € ð \ð \àð %Ð$ Q¨¯ª°A´HÑ(=Ô(=Ñ>Ô>ð\ð \ð \ð \ð \ð \r5   c              3   ó`   •K  — | ](} ‰|‰                      |j        ¦  «        ¦  «        V — Œ)d S r˜   rH  rI  s     €€r6   r,  z7RagTokenForGeneration._reorder_cache.<locals>.<genexpr>¸  sR   øè è € ð 6ð 6àð %Ð$ Q¨¯ª°A´HÑ(=Ô(=Ñ>Ô>ð6ð 6ð 6ð 6ð 6ð 6r5   )
rç   rF   rp   r   Úself_attention_cacheÚlayersÚkeysrë   Úcross_attention_cacheru   )
r   rK  Úreordered_pastÚidxÚself_attention_kÚself_attention_vÚcross_attention_kÚcross_attention_vÚ	new_tuplerF  s
    `       @r6   Ú_reorder_cachez$RagTokenForGeneration._reorder_cacheŸ  so  øø€ ð	ð 	ð 	ð ˆÝ�˜_Ñ-Ô-Ñ.Ô.ð 	+ð 	+ˆCÝ˜/Õ+>Ñ?Ô?ð Að\ð \ð \ð \ð \ð (Ô<ÔCÀCÔHÔMØ'Ô<ÔCÀCÔHÔOØ'Ô=ÔDÀSÔIÔNØ'Ô=ÔDÀSÔIÔPð	ð\ñ \ô \ÑXÐ Ð"2Ð4EÐGXð .Ð/?ÐARÐTeÐf�	�	ð6ð 6ð 6ð 6ð 6à-Ô4°SÔ9Ô>ÀÔ@VÐWZÔ@[Ô@bÐcð6ñ 6ô 6Ñ2Ð Ð"2ð .Ð/?Ð@�	Ø˜y˜lÑ*ˆNˆNØ$�t�OÑ$Ô$ ^Ñ4Ô4Ð4r5   c                 ó€  — |�|n| j         j        }t          j                             |d¬¦  «                             |j        d         |z  |d|                     d¦  «        ¦  «        }t          j        |d¬¦  «        }|| 	                    d¦  «         	                    d¦  «        z   }t          j
        |d¬¦  «        S )Nr’   r–   r   r   )r;   r„   r   r  r  rŸ   r    r  r/   r¢   r  )ry   r  r   r„   r  r  Úlog_prob_sums          r6   Úmarginalizez!RagTokenForGeneration.marginalizeÀ  sµ   € Ø!Ð-��°4´;Ô3Eˆõ ”}×0Ò0°ÀÐ0ÑDÔD×IÒIØÔ˜QÔ 6Ñ)¨6°2°z·²ÀrÑ7JÔ7Jñ
ô 
ˆõ Ô(¨¸Ð;Ñ;Ô;ˆØ# l×&<Ò&<¸RÑ&@Ô&@×&JÒ&JÈ2Ñ&NÔ&NÑNˆÝŒ˜|°Ð3Ñ3Ô3Ð3r5   r{   r|   r}   r~   r   r   r    r!   r   r€   r�   r‚   rƒ   r<  rÎ   rÏ   r„   r@   c                 ó€  — |�|n| j         j        }|�|n| j         j        }|�|n| j         j        }|�|€|}d}
|                      ||||||||	||
||||¬¦  «        }d}|j        }|�3|€J ‚|                      |j        |j        ||| j         j        |¬¦  «        }|r|  	                    ||j        |¦  «        }t          di d|“d|“d|j        “d|j        “d	|j        “d
|j        “d|j        “d|j        “d|j        “d|j        “d|j        “d|j        “d|j        “d|j        “d|j        “d|j        “d|j        “ŽS )aƒ  
        input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
            Indices of input sequence tokens in the vocabulary. [`RagConfig`], used to initialize the model, specifies
            which generator to use, it also specifies a compatible generator tokenizer. Use that tokenizer class to
            obtain the indices.

            [What are input IDs?](../glossary#input-ids)
        encoder_outputs (`tuple(tuple(torch.FloatTensor)`, *optional*)
            Tuple consists of (`generator_enc_last_hidden_state`, *optional*: `generator_enc_hidden_states`,
            *optional*: `generator_enc_attentions`). `generator_enc_last_hidden_state` of shape `(batch_size, n_docs *
            sequence_length, hidden_size)` is a sequence of hidden-states at the output of the last layer of the
            generator's encoder.

            Used by the ([`RagModel`]) model during decoding.
        decoder_input_ids (`torch.LongTensor` of shape `(batch_size, target_sequence_length)`, *optional*):
            Provide for generation tasks. `None` by default, construct as per instructions for the generator model
            you're using with your RAG instance.
        decoder_attention_mask (`torch.BoolTensor` of shape `(batch_size,  target_sequence_length)`, *optional*):
            Default behavior: generate a tensor that ignores pad tokens in `decoder_input_ids`. Causal mask will also
            be used by default.
        context_input_ids (`torch.LongTensor` of shape `(batch_size * config.n_docs, config.max_combined_length)`, *optional*, returned when *output_retrieved=True*):
            Input IDs post-processed from the retrieved documents and the question encoder `input_ids` by the
            retriever. If the model was not initialized with a `retriever` ``context_input_ids` has to be provided to
            the forward pass. `context_input_ids` are returned by [`~RagRetriever.__call__`].
        context_attention_mask (`torch.LongTensor` of shape `(batch_size * config.n_docs, config.max_combined_length)`,*optional*, returned when *output_retrieved=True*):
            Attention mask post-processed from the retrieved documents and the question encoder `input_ids` by the
            retriever. If the model has is not initialized with a `retriever` `context_attention_mask` has to be
            provided to the forward pass. `context_attention_mask` are returned by [`~RagRetriever.__call__`].
        doc_scores (`torch.FloatTensor` of shape `(batch_size, config.n_docs)`):
            Score between each retrieved document embeddings (see `retrieved_doc_embeds`) and
            `question_encoder_last_hidden_state`. If the model has is not initialized with a `retriever` `doc_scores`
            has to be provided to the forward pass. `doc_scores` can be computed via
            `question_encoder_last_hidden_state` and `retrieved_doc_embeds`, see examples for more information.
        output_retrieved (`bool`, *optional*):
            Whether or not to return the `retrieved_doc_embeds`, `retrieved_doc_ids`, `context_input_ids` and
            `context_attention_mask`. See returned tensors for more detail.
        do_marginalize (`bool`, *optional*):
            If `True`, the logits are marginalized over all documents by making use of
            `torch.nn.functional.log_softmax`.
        reduce_loss (`bool`, *optional*):
            Only relevant if `labels` is passed. If `True`, the NLL loss is reduced using the `torch.Tensor.sum`
            operation.
        n_docs (`int`, *optional*):
            The number of documents to retrieve.

        Example:

        ```python
        >>> from transformers import AutoTokenizer, RagRetriever, RagTokenForGeneration
        >>> import torch

        >>> tokenizer = AutoTokenizer.from_pretrained("facebook/rag-token-nq")
        >>> retriever = RagRetriever.from_pretrained(
        ...     "facebook/rag-token-nq", index_name="exact", use_dummy_dataset=True
        ... )
        >>> # initialize with RagRetriever to do everything in one forward call
        >>> model = RagTokenForGeneration.from_pretrained("facebook/rag-token-nq", retriever=retriever)

        >>> inputs = tokenizer("How many people live in Paris?", return_tensors="pt")
        >>> targets = tokenizer(text_target="In Paris, there are 10 million people.", return_tensors="pt")
        >>> input_ids = inputs["input_ids"]
        >>> labels = targets["input_ids"]
        >>> outputs = model(input_ids=input_ids, labels=labels)

        >>> # or use retriever separately
        >>> model = RagTokenForGeneration.from_pretrained("facebook/rag-token-nq", use_dummy_dataset=True)
        >>> # 1. Encode
        >>> question_hidden_states = model.question_encoder(input_ids)[0]
        >>> # 2. Retrieve
        >>> docs_dict = retriever(input_ids.numpy(), question_hidden_states.detach().numpy(), return_tensors="pt")
        >>> doc_scores = torch.bmm(
        ...     question_hidden_states.unsqueeze(1), docs_dict["retrieved_doc_embeds"].float().transpose(1, 2)
        ... ).squeeze(1)
        >>> # 3. Forward to generator
        >>> outputs = model(
        ...     context_input_ids=docs_dict["context_input_ids"],
        ...     context_attention_mask=docs_dict["context_attention_mask"],
        ...     doc_scores=doc_scores,
        ...     decoder_input_ids=labels,
        ... )

        >>> # or directly generate
        >>> generated = model.generate(
        ...     context_input_ids=docs_dict["context_input_ids"],
        ...     context_attention_mask=docs_dict["context_attention_mask"],
        ...     doc_scores=doc_scores,
        ... )
        >>> generated_string = tokenizer.batch_decode(generated, skip_special_tokens=True)
        ```NFrÑ   )rÎ   rÒ   r„   r   r   r   r   r    r!   r   r   r"   r#   r$   r%   r&   r'   r(   r)   r*   r4   )r;   r„   r<  rÎ   r<   r   rÓ   r   rÔ   r[  r   r   r    r!   r   r   r"   r#   r$   r%   r&   r'   r(   r)   r*   )ry   r{   r|   r}   r~   r   r   r    r!   r   r€   r�   r‚   rƒ   r<  rÎ   rÏ   r„   r`   rÕ   r   r   s                         r6   r´   zRagTokenForGeneration.forwardË  s+  € ð^ "Ð-��°4´;Ô3EˆØ+9Ð+E˜˜È4Ì;ÔKeˆØ%0Ð%<�k�kÀ$Ä+ÔBYˆàÐØ Ð(Ø$*Ð!ØˆIà—(’(ØØ)Ø+Ø/Ø#9Ø/Ø#9Ø!Ø+ØØ/Ø!5Ø-Øð ñ 
ô 
ˆð" ˆØ”ˆØÐØ$Ð0Ð0Ð0Ø—<’<Ø”ØÔ"ØØ'ØœÔ3Øð  ñ ô ˆDð ð 	JØ×%Ò% f¨gÔ.@À&ÑIÔIˆFå'ð 
ð 
ð 
Ø�ð
à�6ð
ð Ô)Ð)ð
ð $Ô3Ð3ð	
ð
 &Ô7Ð7ð
ð $+Ô#AÐ#Að
ð ")Ô!=Ð!=ð
ð &Ô7Ð7ð
ð 07Ô/YÐ/Yð
ð (/Ô'IÐ'Ið
ð %,Ô$CÐ$Cð
ð -4Ô,SÐ,Sð
ð )0Ô(KÐ(Kð
ð &-Ô%EÐ%Eð
ð )0Ô(KÐ(Kð
ð  &-Ô%EÐ%Eð!
ð" (/Ô'IÐ'Ið#
ð 	
r5   Úgeneration_configÚprefix_allowed_tokens_fnÚlogits_processorÚstopping_criteriac           	      ó<  ‡‡— |                       d|ddd¦  «        } | j        |fi |¤Ž\  }}|                     ¦   «         }|t          j        t          j        t          j        t          j        fvrt          d|› d�¦  «        ‚t          t          | ¦  «        t          |         ¦  «        }|                      |                     ¦   «         ¦  «         |                      |||¦  «         |                     dd¦  «        du}|                      ||¦  «         ‰�‰n| j        j        Š| j        ��<|�€9|                      ||¬¦  «        d         }|                      ||                     ¦   «                              dt.          j        ¬	¦  «                             ¦   «         t          | j        j        d
d¦  «        ‰d¬¦  «        }|d         |d         |d         }}}|                     |¦  «        }|                     |¦  «        }|                     |¦  «        }t/          j        |                     d¦  «        |                     dd¦  «        ¦  «                             d¦  «        }|j        d         ‰z  dk    sJ d‰› d|j        d         › d�¦   «         ‚|j        d         ‰z  Š| j         j         !                    ¦   «         } |||d¬¦  «        }t/          j"        ‰|j#        z  df|j$        t.          j%        tM          |  '                    ¦   «         ¦  «        j(        ¬¦  «        }|j        d         }|d         }d%ˆˆfd„	} |||j#        ¬¦  «        } |||j#        ¬¦  «        |d<   | )                    |j#        d¬¦  «        }||d<   ||d<   ||d<   ‰|d<   |j*        |d <   |  +                    |||||	|j(        ¬!¦  «        }|  ,                    ||
¬"¦  «        }|  -                    ||d|j        d         |j.        dz
  ¬#¦  «          || |f|||d$œ|¤|¤ŽS )&aÙ  
        Implements RAG token decoding.

        Args:
            input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
                The sequence used as a prompt for the generation. If `input_ids` is not passed, then
                `context_input_ids` has to be provided.
            attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
                Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:

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

                [What are attention masks?](../glossary#attention-mask)
            context_input_ids (`torch.LongTensor` of shape `(batch_size * config.n_docs, config.max_combined_length)`, *optional*, returned when *output_retrieved=True*):
                Input IDs post-processed from the retrieved documents and the question encoder `input_ids` by the
                retriever.

                If the model has is not initialized with a `retriever`, `context_input_ids` has to be provided to the
                forward pass. `context_input_ids` are returned by [`~RagRetriever.__call__`].
            context_attention_mask (`torch.LongTensor` of shape `(batch_size * config.n_docs, config.max_combined_length)`, *optional*, returned when *output_retrieved=True*):
                Attention mask post-processed from the retrieved documents and the question encoder `input_ids` by the
                retriever.

                If the model has is not initialized with a `retriever`, `context_input_ids` has to be provided to the
                forward pass. `context_input_ids` are returned by [`~RagRetriever.__call__`].
            doc_scores (`torch.FloatTensor` of shape `(batch_size, config.n_docs)`):
                Score between each retrieved document embeddings (see `retrieved_doc_embeds`) and
                `question_encoder_last_hidden_state`.

                If the model has is not initialized with a `retriever`, `context_input_ids` has to be provided to the
                forward pass. `context_input_ids` are returned by [`~RagRetriever.__call__`].
            n_docs (`int`, *optional*, defaults to `config.n_docs`)
                Number of documents to retrieve and/or number of documents for which to generate an answer.
            generation_config (`~generation.GenerationConfig`, *optional*):
                The generation configuration to be used as base parametrization for the generation call. `**kwargs`
                passed to generate matching the attributes of `generation_config` will override them. If
                `generation_config` is not provided, the default will be used, which has the following loading
                priority: 1) from the `generation_config.json` model file, if it exists; 2) from the model
                configuration. Please note that unspecified parameters will inherit [`~generation.GenerationConfig`]'s
                default values, whose documentation should be checked to parameterize generation.
            prefix_allowed_tokens_fn (`Callable[[int, torch.Tensor], list[int]]`, *optional*):
                If provided, this function constraints the beam search to allowed tokens only at each step. If not
                provided no constraint is applied. This function takes 2 arguments `inputs_ids` and the batch ID
                `batch_id`. It has to return a list with the allowed tokens for the next generation step conditioned on
                the previously generated tokens `inputs_ids` and the batch ID `batch_id`. This argument is useful for
                constrained generation conditioned on the prefix, as described in [Autoregressive Entity
                Retrieval](https://huggingface.co/papers/2010.00904).
            logits_processor (`LogitsProcessorList`, *optional*):
                Custom logits processors that complement the default logits processors built from arguments and a
                model's config. If a logit processor is passed that is already created with the arguments or a model's
                config an error is thrown.
            stopping_criteria (`StoppingCriteriaList`, *optional*):
                Custom stopping criteria that complement the default stopping criteria built from arguments and a
                model's config. If a stopping criteria is passed that is already created with the arguments or a
                model's config an error is thrown.
            kwargs (`dict[str, Any]`, *optional*):
                Ad hoc parametrization of `generate_config` and/or additional model-specific kwargs that will be
                forwarded to the `forward` function of the model.

        Return:
            `torch.LongTensor` of shape `(batch_size * num_return_sequences, sequence_length)`: The generated
            sequences. The second dimension (sequence_length) is either equal to `max_length` or shorter if all batches
            finished early due to the `eos_token_id`.
        NFz!RAG model is not compatible with z5 generation. Please check your generation parameters.r|   râ   r   r‡   rˆ   r‹   rŒ   r�   r    r!   r   r   rO   r“   r”   r•   T)r{   r|   r†   )rŠ   r‰   r’   Úlast_hidden_statec                 ó  •— | d d d d …f                               ‰d‰f| j        dd …         z   ¦  «        } |                      ‰|‰f| j        dd …         z   ¦  «        } |                       ‰|z  ‰z  f| j        dd …         z   ¦  «        S )Nr   r   )Úreshaper    Úexpand)Útensorrà   rô   r„   s     €€r6   Úextend_enc_outputz9RagTokenForGeneration.generate.<locals>.extend_enc_output  sœ   ø€ à˜D $¨¨¨˜MÔ*×2Ò2°JÀÀ6Ð3JÈVÌ\ÐZ[ÐZ\ÐZ\ÔM]Ñ3]Ñ^Ô^ˆFà—]’] J°	¸6Ð#BÀVÄ\ÐRSÐRTÐRTÔEUÑ#UÑVÔVˆFà—>’> :°	Ñ#9¸FÑ#BÐ"DÀvÄ|ÐTUÐTVÐTVÔGWÑ"WÑXÔXÐXr5   )rà   r–   r   r}   r„   r€   )r]  Úinput_ids_seq_lengthÚencoder_input_idsr^  r_  r‰   )r]  r`  )Úgeneration_moderô   Úmax_cache_length)r_  r`  r]  r˜   )/Ú_extract_generation_mode_kwargsÚ_prepare_generation_configÚget_generation_moder   ÚSAMPLEÚGREEDY_SEARCHÚBEAM_SEARCHÚBEAM_SAMPLEÚ
ValueErrorr�   ru   r   Ú_validate_model_kwargsÚcopyÚ_validate_generation_moder]   Ú_prepare_special_tokensr;   r„   r?   rV   r™   rš   r/   r›   rœ   rW   r¡   r¢   r£   r¤   r    r<   Úget_encoderÚfullrà   Údecoder_start_token_idÚlongÚnextÚ
parametersr‰   r¥   r€   Ú_get_logits_processorÚ_get_stopping_criteriaÚ_prepare_cache_for_generationÚ
max_length)ry   r{   r|   r    r!   r   r„   r]  r^  r_  r`  r`   Úgeneration_mode_kwargsrð   rj  Údecoding_methodÚkwargs_has_attention_maskrò   Úoutr   Úencoderr}   rh  rb  rg  Úprepared_logits_processorÚprepared_stopping_criteriarô   s         `                    @r6   rè   zRagTokenForGeneration.generatex  s"  øø€ ðb "&×!EÒ!EÀdÈFÐTYÐ[_ÐaeÑ!fÔ!fÐØ*I¨$Ô*IÐJ[Ð*fÐ*fÐ_eÐ*fÐ*fÑ'Ð˜<Ø+×?Ò?ÑAÔAˆØÝÔ!ÝÔ(ÝÔ&ÝÔ&ð	#
ð 
ð 
õ Øz°OÐzÐzÐzñô ð õ "¥$ t¡*¤*Õ.FÀÔ.WÑXÔXˆØ×#Ò# L×$5Ò$5Ñ$7Ô$7Ñ8Ô8Ð8Ø×&Ò& Ð8IÐKaÑbÔbÐbà$0×$4Ò$4Ð5EÀtÑ$LÔ$LÐTXÐ$XÐ!Ø×$Ò$Ð%6Ð8QÑRÔRÐRð "Ð-��°4´;Ô3Eˆð Œ>Ñ%Ð*;Ñ*CØ%)×%:Ò%:¸9ÐUcÐ%:Ñ%dÔ%dÐefÔ%gÐ"Ø—.’.ØØ&×-Ò-Ñ/Ô/×2Ò2¸%ÅuÄ}Ð2ÑUÔU×[Ò[Ñ]Ô]Ý˜tœ~Ô4°hÀÑEÔEØØ#ð !ñ ô ˆCð Ð'Ô(ØÐ,Ô-ØÐ*Ô+ð 8LÐ5Ðð $8×#:Ò#:Ð;QÑ#RÔ#RÐ Ø 1× 4Ò 4°YÑ ?Ô ?ÐØ%;×%>Ò%>¸yÑ%IÔ%IÐ"õ œÐ#9×#CÒ#CÀAÑ#FÔ#FÐH\×HfÒHfÐghÐjkÑHlÔHlÑmÔm×uÒuØñô ˆJð "Ô'¨Ô*¨VÑ3¸Ò9Ð9Ð9ð.Ð\bð .ð .Ø!Ô'¨Ô*ð.ð .ð .ñ :Ô9Ð9ð 'Ô,¨QÔ/°6Ñ9ˆ
à”(Ô$×0Ò0Ñ2Ô2ˆØ!˜'Ð,=ÐNdÐrvÐwÑwÔwˆå”JØÐ+Ô5Ñ5°qÐ9ØÔ4Ý”*Ý˜ŸšÑ)Ô)Ñ*Ô*Ô1ð	
ñ 
ô 
ˆ	ð  )œ¨rÔ2ÐØ+Ð,?Ô@Ðð	Yð 	Yð 	Yð 	Yð 	Yð 	Yð 	Yð "3Ð!2Ð3IÐUfÔUpÐ!qÑ!qÔ!qÐØ/@Ð/@ØÐ):Ô)Dð0
ñ 0
ô 0
ˆÐ+Ñ,ð  ×1Ò1Ð2CÔ2MÐSTÐ1ÑUÔUˆ
ð &0ˆ�\Ñ"Ø*9ˆÐ&Ñ'Ø)?ˆÐ%Ñ&Ø!'ˆ�XÑØ$5Ô$?ˆ�[Ñ!à$(×$>Ò$>Ø/Ø!5Ø/Ø%=Ø-ØÔ#ð %?ñ %
ô %
Ð!ð &*×%@Ò%@Ø/ÐCTð &Añ &
ô &
Ð"ð 	×*Ò*ØØØ Ø ” qÔ)Ø.Ô9¸AÑ=ð 	+ñ 	
ô 	
ð 	
ð ˆØØð
ð 7Ø8Ø/ð
ð 
ð %ð
ð ð
ð 
ð 	
r5   c                 ó2   — |                       ||¦  «        }|S r˜   )rX  )ry   r   rK  s      r6   Ú_temporary_reorder_cachez.RagTokenForGeneration._temporary_reorder_cacheE  s   € ð ×-Ò-¨o¸xÑHÔHˆØÐr5   c                 ó>   — | j         j                             ¦   «         S r˜   )r<   rW   Úget_input_embeddingsr×   s    r6   rŒ  z*RagTokenForGeneration.get_input_embeddingsL  s   € ØŒxÔ!×6Ò6Ñ8Ô8Ð8r5   c                 ó>   — | j         j                             ¦   «         S r˜   )r<   rW   Úget_output_embeddingsr×   s    r6   rŽ  z+RagTokenForGeneration.get_output_embeddingsO  s   € ØŒxÔ!×7Ò7Ñ9Ô9Ð9r5   c                 ó@   — | j         j                             |¦  «        S r˜   )r<   rW   Úset_output_embeddings)ry   Únew_embeddingss     r6   r�  z+RagTokenForGeneration.set_output_embeddingsR  s   € ØŒxÔ!×7Ò7¸ÑGÔGÐGr5   c                 óº   — |€| j         j        }|                     |j        ¦  «        }|dd…dd…f                              ¦   «         |dd…dd…f<   ||dd…df<   |S )zCShift input ids one token to the right, and pad with start_token_idNr’   r   r   )r;   rz  Ú	new_zerosr    Úclone)ry   r{   Ústart_token_idÚshifted_input_idss       r6   Úshift_tokens_rightz(RagTokenForGeneration.shift_tokens_rightU  su   € àÐ!Ø!œ[Ô?ˆNØ%×/Ò/°	´Ñ@Ô@ÐØ#,¨Q¨Q¨Q°°°¨VÔ#4×#:Ò#:Ñ#<Ô#<Ð˜!˜!˜!˜Q˜R˜R˜%Ñ Ø"0Ð˜!˜!˜!˜Q˜$ÑØ Ð r5   Frþ   c                 ó(  ‡ ‡— |�|n‰ j         j        }t          j        ‰d d …dd …f         ‰                     ‰j        d         d¦  «                             ‰ j         j        j        ¦  «        gd¦  «        Šˆ ˆfd„}‰  	                    |||¦  «        }‰ 
                    d¦  «        Š‰                     ¦   «         |                     ¦   «         k    sJ ‚|                     d‰¬¦  «        }	|                     dd¬¦  «        }
 ||	|
¦  «        \  }	}
|	                     d¦  «        }	|
                     d¦  «        }
|	 }|
 }|r(|                     ¦   «         }|                     ¦   «         }||                     d¦  «        z  }d|z
  |z  ||z  z   }|S )	Nr   r   c                 ó   •— ‰                      ‰j        j        j        ¦  «        }|                     ¦   «         r,|                      |d¦  «         |                     |d¦  «         |                      d¦  «        |                     d¦  «        fS r  r  r  s      €€r6   r  z1RagTokenForGeneration.get_nll.<locals>._mask_padse  r  r5   r’   r  Tr  r  )r;   r„   r/   r  r  r    r  rW   ræ   r[  r¢   r—   r  r  r  )ry   r  r   r
  rÎ   rÒ   r„   r  r#  r  r  r$  r%  r&  r   s   `  `           r6   rÓ   zRagTokenForGeneration.get_nll^  s¨  øø€ Ø!Ð-��°4´;Ô3Eˆå”Ø�A�A�A�q�r�r�EŒ]˜FŸJšJ v¤|°A¤¸Ñ:Ô:×@Ò@ÀÄÔAVÔAcÑdÔdÐeÐghñ
ô 
ˆð	:ð 	:ð 	:ð 	:ð 	:ð 	:ð ×'Ò'¨
°JÀÑGÔGˆà×!Ò! "Ñ%Ô%ˆØ�zŠz‰|Œ|˜|×/Ò/Ñ1Ô1Ò1Ð1Ð1Ð1à× Ò  R¨vÐ Ñ6Ô6ˆØ!×%Ò%¨"°dÐ%Ñ;Ô;ˆ
Ø#˜ B¨
Ñ3Ô3‰ˆˆJØ�VŠV�A‰YŒYˆØ—^’^ AÑ&Ô&ˆ
à�3ˆØ!�kˆàð 	,Ø—|’|‘~”~ˆHØ%Ÿ/š/Ñ+Ô+ˆKà˜,×+Ò+¨BÑ/Ô/Ñ/ˆØ�g‘ Ñ)¨E°KÑ,?Ñ?ˆØˆr5   rµ   )NNNNNNr˜   r2  )Frþ   N)+r+   r,   r-   r	   r   r   rs   rÅ   rË   r=  r3  r?   rW   rV   r5  rX  r[  r   r/   r2   r0   r3   r¶   r·   r   r¸   r¹   r   r´   r4  r   r   r
   r   rê   rè   rŠ  rŒ  rŽ  r�  r—  rÓ   rº   r»   s   @r6   r7  r7  I  sš  ø€ € € € € ð +/Ø37Ø,0Ø)-ðð à  4Ñ'ðð *¨DÑ0ðð # TÑ)ð	ð
   $Ñ&ðð ð ð ð ð ð@' |ð 'ð 'ð 'ð 'ð+¸Oð +ð +ð +ð +ð ØØØØØð
ð 
ð 
ð 
ð: ð"ð "ñ „Xð"ð ð"ð "ñ „Xð"ð ð)ð )ñ „Xð)ð ð5ð 5ñ „\ð5ð@	4ð 	4ð 	4ð 	4ð ð .2Ø37Ø=AØ59Ø:>Ø(,Ø59Ø:>Ø/3Ø!%Ø)-Ø,0Ø(,Ø&*Ø#'Ø*.Ø!ð%j
ð j
àÔ# dÑ*ðj
ð Ô)¨DÑ0ðj
ð ˜u U¤\Ô2Ô3°dÑ:ð	j
ð
 !Ô+¨dÑ2ðj
ð !&Ô 0°4Ñ 7ðj
ð  ™ðj
ð !Ô+¨dÑ2ðj
ð !&Ô 0°4Ñ 7ðj
ð Ô%¨Ñ,ðj
ð ˜$‘;ðj
ð   $™;ðj
ð # T™kðj
ð  ™+ðj
ð ˜t™ðj
ð  ˜D‘[ð!j
ð" Ô  4Ñ'ð#j
ð$ �d‘
ð%j
ð( 
"ð)j
ð j
ð j
ñ „^ðj
ðX €U„]�_„_ð .2Ø26Ø59Ø:>Ø/3Ø!Ø59ØTXØ7JÐ7JÑ7LÔ7LØ9MÐ9MÑ9OÔ9OðI
ð I
àÔ# dÑ*ðI
ð Ô(¨4Ñ/ðI
ð !Ô+¨dÑ2ð	I
ð
 !&Ô 0°4Ñ 7ðI
ð Ô%¨Ñ,ðI
ð �d‘
ðI
ð ,¨dÑ2ðI
ð #+¨C°´Ð+>ÀÀSÄ	Ð+IÔ"JÈTÑ"QðI
ð .°Ñ4ðI
ð 0°$Ñ6ðI
ð 
Ô	ðI
ð I
ð I
ñ „_ðI
ðXð ð ð9ð 9ð 9ð:ð :ð :ðHð Hð Hð!ð !ð !ð !ð"ð "ð "ð "ð "ð "ð "ð "r5   r7  )rm   r:   r½   r7  ))r.   Úcollections.abcr   Údataclassesr   r/   r   Úcache_utilsr   r   Úconfiguration_utilsr	   Ú
generationr
   r   r   r   r   Úgeneration.utilsr   Úmodeling_outputsr   Úmodeling_utilsr   Úutilsr   r   Úconfiguration_ragr   Úretrieval_ragr   Ú
get_loggerr+   Úloggerr   r8   r:   rm   r½   r7  Ú__all__r4   r5   r6   ú<module>r¨     sU  ðð  Ð à $Ð $Ð $Ð $Ð $Ð $Ø !Ð !Ð !Ð !Ð !Ð !à €€€Ø Ð Ð Ð Ð Ð à 5Ð 5Ð 5Ð 5Ð 5Ð 5Ð 5Ð 5Ø 3Ð 3Ð 3Ð 3Ð 3Ð 3Ø vÐ vÐ vÐ vÐ vÐ vÐ vÐ vÐ vÐ vÐ vÐ vÐ vÐ vØ 8Ð 8Ð 8Ð 8Ð 8Ð 8Ø +Ð +Ð +Ð +Ð +Ð +Ø -Ð -Ð -Ð -Ð -Ð -Ø ,Ð ,Ð ,Ð ,Ð ,Ð ,Ð ,Ð ,Ø (Ð (Ð (Ð (Ð (Ð (Ø 'Ð 'Ð 'Ð 'Ð 'Ð 'ð 
ˆÔ	˜HÑ	%Ô	%€ð €ððñ ô ð
 ðWLð WLð WLð WLð WL˜{ñ WLô WLñ „ñô ðWLðt Ø
ðTLð TLð TLð TLð TL˜ñ TLô TLñ „ñ „ðTLðn €ððñ ô ð ðIoð Ioð Ioð Ioð Io˜ñ Ioô Ioñ „ñô ðIoðX ð[
ð [
ð [
ð [
ð [
Ð!ñ [
ô [
ñ „ð[
ð| €ððñ ô ð
oð oð oð oð oÐ1ñ oô oñô ð
oðd €ððñ ô ð
rð rð rð rð rÐ.°ñ rô rñô ð
rðj bÐ
aÐ
a€€€r5   