Ë
    T^(hé  ã                   óò   — d 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d	Zd
ej                   dededej                   fd„Z G d„ de	«      Z G d„ de«      Z G d„ de«      Zg d¢Zy)zFlax mT5 model.é    Né   )Úloggingé   )ÚFlaxT5EncoderModelÚFlaxT5ForConditionalGenerationÚFlaxT5Modelé   )Ú	MT5ConfigÚT5ConfigÚ	input_idsÚpad_token_idÚdecoder_start_token_idÚreturnc                 ó  — t        j                  | «      }|j                  dd…dd…f   j                  | dd…dd…f   «      }|j                  dd…df   j                  |«      }t        j                  |dk(  ||«      }|S )z1
    Shift input ids one token to the right.
    Nr	   éÿÿÿÿr   iœÿÿÿ)ÚjnpÚ
zeros_likeÚatÚsetÚwhere)r   r   r   Úshifted_input_idss       úg/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/mt5/modeling_flax_mt5.pyÚshift_tokens_rightr      sƒ   € ô Ÿ™ yÓ1ÐØ)×,Ñ,ªQ°±¨UÑ3×7Ñ7¸	Â!ÀSÀbÀSÀ&Ñ8IÓJÐØ)×,Ñ,ªQ°¨TÑ2×6Ñ6Ð7MÓNÐäŸ	™	Ð"3°tÑ";¸\ÐK\Ó]ÐØÐó    c                   ó   — e Zd ZdZdZeZy)ÚFlaxMT5Modela  
    This class overrides [`FlaxT5Model`]. Please check the superclass for the appropriate documentation alongside usage
    examples.

    Examples:

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

    >>> model = FlaxMT5Model.from_pretrained("google/mt5-small")
    >>> tokenizer = AutoTokenizer.from_pretrained("google/mt5-small")

    >>> article = "UN Offizier sagt, dass weiter verhandelt werden muss in Syrien."
    >>> summary = "Weiter Verhandlung in Syrien."
    >>> inputs = tokenizer(article, return_tensors="np")

    >>> decoder_input_ids = tokenizer(text_target=summary, return_tensors="np").input_ids

    >>> outputs = model(input_ids=inputs["input_ids"], decoder_input_ids=decoder_input_ids)
    >>> hidden_states = outputs.last_hidden_state
    ```Úmt5N©Ú__name__Ú
__module__Ú__qualname__Ú__doc__Ú
model_typer
   Úconfig_class© r   r   r   r   *   ó   „ ñð, €JØ�Lr   r   c                   ó   — e Zd ZdZdZeZy)ÚFlaxMT5EncoderModela	  
    This class overrides [`FlaxT5EncoderModel`]. Please check the superclass for the appropriate documentation
    alongside usage examples.

    Examples:

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

    >>> model = FlaxT5EncoderModel.from_pretrained("google/mt5-small")
    >>> tokenizer = AutoTokenizer.from_pretrained("google/mt5-small")

    >>> article = "UN Offizier sagt, dass weiter verhandelt werden muss in Syrien."
    >>> summary = "Weiter Verhandlung in Syrien."
    >>> inputs = tokenizer(article, return_tensors="np")

    >>> decoder_input_ids = tokenizer(text_target=summary, return_tensors="np").input_ids

    >>> outputs = model(input_ids=inputs["input_ids"])
    >>> hidden_states = outputs.last_hidden_state
    ```r   Nr   r%   r   r   r(   r(   E   r&   r   r(   c                   ó   — e Zd ZdZdZeZy)ÚFlaxMT5ForConditionalGenerationa-  
    This class overrides [`FlaxT5ForConditionalGeneration`]. Please check the superclass for the appropriate
    documentation alongside usage examples.

    Examples:

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

    >>> model = FlaxMT5ForConditionalGeneration.from_pretrained("google/mt5-small")
    >>> tokenizer = AutoTokenizer.from_pretrained("google/mt5-small")

    >>> article = "UN Offizier sagt, dass weiter verhandelt werden muss in Syrien."
    >>> summary = "Weiter Verhandlung in Syrien."
    >>> inputs = tokenizer(article, return_tensors="np")

    >>> decoder_input_ids = tokenizer(text_target=summary, return_tensors="np").input_ids

    >>> outputs = model(**inputs, decoder_input_ids=decoder_input_ids)
    >>> logits = outputs.logits
    ```r   Nr   r%   r   r   r*   r*   `   r&   r   r*   )r(   r*   r   )r"   Ú	jax.numpyÚnumpyr   Úutilsr   Út5.modeling_flax_t5r   r   r   Úconfiguration_mt5r
   Ú
get_loggerr   ÚloggerÚ_CONFIG_FOR_DOCÚndarrayÚintr   r   r(   r*   Ú__all__r%   r   r   ú<module>r6      s�   ðñ å å ß aÑ aÝ (ð 
ˆ×	Ñ	˜HÓ	%€à€ð	 #§+¡+ð 	¸Sð 	ÐZ]ð 	Ðbe×bmÑbmó 	ô�;ô ô6Ð,ô ô6Ð&Dô ò6 U�r   