Ë
    S^(hßY  ã                   ó”  — d dl Z d dlZd dlmZ d dlmZ d dlmZ ddl	m
Z
 ddlmZ  ee«      ZdZ G d„ d«      Z G d	„ d
«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Z G d„ d e«      Z G d!„ d"e«      Zy)#é    N)Úsparseé   )Úadd_start_docstrings)Ú
get_loggerad  
    Args:
        input_ids (`jnp.ndarray` of shape `(batch_size, sequence_length)`):
            Indices of input sequence tokens in the vocabulary.

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

            [What are input IDs?](../glossary#input-ids)
        scores (`jnp.ndarray` of shape `(batch_size, config.vocab_size)`):
            Prediction scores of a language modeling head. These can be logits for each vocabulary when not using beam
            search or log softmax for each vocabulary token when using beam search
        kwargs (`Dict[str, Any]`, *optional*):
            Additional logits processor specific kwargs.

    Return:
        `jnp.ndarray` of shape `(batch_size, config.vocab_size)`: The processed prediction scores.

c                   óv   — e Zd ZdZ ee«      dej                  dej                  dej                  fd„«       Zy)ÚFlaxLogitsProcessorzSAbstract base class for all logit processors that can be applied during generation.Ú	input_idsÚscoresÚreturnc                 ó2   — t        | j                  › d�«      ‚)z"Flax method for processing logits.úH is an abstract class. Only classes inheriting this class can be called.©ÚNotImplementedErrorÚ	__class__©Úselfr	   r
   s      úi/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/generation/flax_logits_process.pyÚ__call__zFlaxLogitsProcessor.__call__6   ó!   € ô "Ø�~‰~ÐÐfÐgó
ð 	
ó    N©	Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Ú!LOGITS_PROCESSOR_INPUTS_DOCSTRINGÚjnpÚndarrayr   © r   r   r   r   3   s>   „ Ù]áÐ;Ó<ð
 #§+¡+ð 
°s·{±{ð 
ÀsÇ{Á{ò 
ó =ñ
r   r   c                   óv   — e Zd ZdZ ee«      dej                  dej                  dej                  fd„«       Zy)ÚFlaxLogitsWarperzjAbstract base class for all logit warpers that can be applied during generation with multinomial sampling.r	   r
   r   c                 ó2   — t        | j                  › d�«      ‚)zFlax method for warping logits.r   r   r   s      r   r   zFlaxLogitsWarper.__call__A   r   r   Nr   r   r   r   r!   r!   >   s>   „ ÙtáÐ;Ó<ð
 #§+¡+ð 
°s·{±{ð 
ÀsÇ{Á{ò 
ó =ñ
r   r!   c            	       óz   — e Zd ZdZ ee«      dej                  dej                  dedej                  fd„«       Z	y)ÚFlaxLogitsProcessorLista.  
    This class can be used to create a list of [`FlaxLogitsProcessor`] or [`FlaxLogitsWarper`] to subsequently process
    a `scores` input tensor. This class inherits from list and adds a specific *__call__* method to apply each
    [`FlaxLogitsProcessor`] or [`FlaxLogitsWarper`] to the inputs.
    r	   r
   Úcur_lenr   c                 ór  ‡— | D ]°  }t        j                  |j                  «      j                  }t	        |«      dkD  rmt        ˆfd„t        |j                  «       «      dd  D «       «      s3t        dt        |j                  «       «      › d|j                  › d�«      ‚ ||||fi ‰¤Ž}Œ§ ||||«      }Œ² |S )Né   c              3   ó&   •K  — | ]  }|‰v –— Œ
 y ­w©Nr   )Ú.0ÚargÚkwargss     €r   ú	<genexpr>z3FlaxLogitsProcessorList.__call__.<locals>.<genexpr>U   s   øè ø€ ÒS¨S˜3 &œ=ÑSùs   ƒr   z,Make sure that all the required parameters: z for z$ are passed to the logits processor.)
ÚinspectÚ	signaturer   Ú
parametersÚlenÚallÚlistÚkeysÚ
ValueErrorr   )r   r	   r
   r%   r,   Ú	processorÚfunction_argss       `  r   r   z FlaxLogitsProcessorList.__call__P   sÃ   ø€ àò 
	?ˆIÜ#×-Ñ-¨i×.@Ñ.@ÓA×LÑLˆMÜ�=Ó! AÒ%ÜÓS´D¸×9KÑ9KÓ9MÓ4NÈqÈrÐ4RÔSÔSÜ$ØFÄtÈM×L^ÑL^ÓL`ÓGaÐFbÐbgØ$×.Ñ.Ð/Ð/SðUóð ñ # 9¨f°gÑHÀÑH‘á" 9¨f°gÓ>‘ð
	?ð ˆr   N)
r   r   r   r   r   r   r   r   Úintr   r   r   r   r$   r$   I   sL   „ ññ Ð;Ó<ð #§+¡+ð °s·{±{ð ÈSð Ð_b×_jÑ_jò ó =ñr   r$   c                   óp   — e Zd ZdZdefd„Zdej                  dej                  dedej                  fd„Z	y	)
ÚFlaxTemperatureLogitsWarperzÍ
    [`FlaxLogitsWarper`] for temperature (exponential scaling output probability distribution).

    Args:
        temperature (`float`):
            The value used to module the logits distribution.
    Útemperaturec                 óX   — t        |t        «      r|dkD  st        d|› �«      ‚|| _        y )Nr   z:`temperature` has to be a strictly positive float, but is )Ú
isinstanceÚfloatr5   r;   )r   r;   s     r   Ú__init__z$FlaxTemperatureLogitsWarper.__init__i   s/   € Ü˜+¤uÔ-°kÀA²oÜÐYÐZeÐYfÐgÓhÐhà&ˆÕr   r	   r
   r%   r   c                 ó$   — || j                   z  }|S r)   )r;   ©r   r	   r
   r%   s       r   r   z$FlaxTemperatureLogitsWarper.__call__o   s   € Ø˜$×*Ñ*Ñ*ˆØˆr   N)
r   r   r   r   r>   r?   r   r   r8   r   r   r   r   r:   r:   `   sC   „ ñð' Eó 'ð #§+¡+ð °s·{±{ð ÈSð ÐUX×U`ÑU`ô r   r:   c                   óŒ   — e Zd ZdZ ed«       dfdededefd„Zdej                  d	ej                  d
edej                  fd„Z	y)ÚFlaxTopPLogitsWarpera=  
    [`FlaxLogitsWarper`] that performs top-p, i.e. restricting to top tokens summing to prob_cut_off <= prob_cut_off.

    Args:
        top_p (`float`):
            If set to < 1, only the smallest set of most probable tokens with probabilities that add up to `top_p` or
            higher are kept for generation.
        filter_value (`float`, *optional*, defaults to -inf):
            All filtered values will be set to this float value.
        min_tokens_to_keep (`int`, *optional*, defaults to 1):
            Minimum number of tokens that cannot be filtered.
    ÚInfé   Útop_pÚfilter_valueÚmin_tokens_to_keepc                 óÄ   — t        |t        «      r
|dk  s|dkD  rt        d|› �«      ‚t        |t        «      r|dk  rt        d|› �«      ‚|| _        || _        || _        y )Nr   g      ð?z.`top_p` has to be a float > 0 and < 1, but is rE   z:`min_tokens_to_keep` has to be a positive integer, but is )r=   r>   r5   r8   rF   rG   rH   )r   rF   rG   rH   s       r   r?   zFlaxTopPLogitsWarper.__init__‚   sj   € Ü˜%¤Ô'¨E°AªI¸ÀºÜÐMÈeÈWÐUÓVÐVÜÐ,¬cÔ2Ð7IÈAÒ7MÜÐYÐZlÐYmÐnÓoÐoàˆŒ
Ø(ˆÔØ"4ˆÕr   r	   r
   r%   r   c                 óX  — t        j                  ||j                  d   «      \  }}t        j                  || j
                  «      }t        j                  j                  |d¬«      j                  d¬«      }|| j                  k  }t        j                  |d«      }||j                  d d …df   j                  d«      z  }|j                  d d …d | j                  …f   j                  d«      }t        j                  |||«      }	t        j                   j!                  ||	«      d   }
|
S )Néÿÿÿÿ©ÚaxisrE   r   T)ÚlaxÚtop_kÚshaper   Ú	full_likerG   ÚjaxÚnnÚsoftmaxÚcumsumrF   ÚrollÚatÚsetrH   ÚwhereÚsort_key_val)r   r	   r
   r%   Útopk_scoresÚtopk_indicesÚmask_scoresÚcumulative_probsÚ
score_maskÚtopk_next_scoresÚnext_scoress              r   r   zFlaxTopPLogitsWarper.__call__Œ   sþ   € Ü$'§I¡I¨f°f·l±lÀ2Ñ6FÓ$GÑ!ˆ�\ä—m‘m F¨D×,=Ñ,=Ó>ˆÜŸ6™6Ÿ>™>¨+¸B˜>Ó?×FÑFÈBÐFÓOÐØ%¨¯
©
Ñ2ˆ
ô —X‘X˜j¨!Ó,ˆ
Ø�j—m‘m¢A q DÑ)×-Ñ-¨dÓ3Ñ3ˆ
ð  —]‘]¢1Ð&?¨×(?Ñ(?Ð&?Ð#?Ñ@×DÑDÀTÓJˆ
äŸ9™9 Z°¸kÓJÐÜ—g‘g×*Ñ*¨<Ð9IÓJÈ2ÑNˆàÐr   N©
r   r   r   r   r>   r8   r?   r   r   r   r   r   r   rC   rC   t   sa   „ ññ =BÀ%»L¸=Ðdeñ 5˜eð 5°5ð 5Ð^aó 5ð #§+¡+ð °s·{±{ð ÈSð ÐUX×U`ÑU`ô r   rC   c                   óŒ   — e Zd ZdZ ed«       dfdededefd„Zdej                  d	ej                  d
edej                  fd„Z	y)ÚFlaxTopKLogitsWarperaæ  
    [`FlaxLogitsWarper`] that performs top-k, i.e. restricting to the k highest probability elements.

    Args:
        top_k (`int`):
            The number of highest probability vocabulary tokens to keep for top-k-filtering.
        filter_value (`float`, *optional*, defaults to -inf):
            All filtered values will be set to this float value.
        min_tokens_to_keep (`int`, *optional*, defaults to 1):
            Minimum number of tokens that cannot be filtered.
    rD   rE   rO   rG   rH   c                 óz   — t        |t        «      r|dk  rt        d|› �«      ‚t        ||«      | _        || _        y )Nr   z6`top_k` has to be a strictly positive integer, but is )r=   r8   r5   ÚmaxrO   rG   )r   rO   rG   rH   s       r   r?   zFlaxTopKLogitsWarper.__init__­   s>   € Ü˜%¤Ô%¨°!ªÜÐUÐV[ÐU\Ð]Ó^Ð^ä˜Ð 2Ó3ˆŒ
Ø(ˆÕr   r	   r
   r%   r   c                 ó  — |j                   \  }}t        j                  ||z  | j                  «      }t	        | j
                  |j                   d   «      }t        j
                  ||«      \  }}	t        j                  t        j                  |«      |z  d d …d f   ||f«      j                  «       }
|j                  «       }|	j                  «       |
z   }|j                  |   j                  |«      }|j                  ||«      }|S )NrK   )rP   r   ÚfullrG   ÚminrO   rN   Úbroadcast_toÚarangeÚflattenrW   rX   Úreshape)r   r	   r
   r%   Ú
batch_sizeÚ
vocab_sizeÚnext_scores_flatÚtopkr[   r\   ÚshiftÚtopk_scores_flatÚtopk_indices_flatra   s                 r   r   zFlaxTopKLogitsWarper.__call__´   sì   € Ø!'§¡Ñˆ
�JÜŸ8™8 J°Ñ$;¸T×=NÑ=NÓOÐä�4—:‘:˜vŸ|™|¨BÑ/Ó0ˆÜ$'§I¡I¨f°dÓ$;Ñ!ˆ�\Ü× Ñ ¤#§*¡*¨ZÓ"8¸:Ñ"EÂqÈ$ÀwÑ!OÐR\Ð^bÐQcÓd×lÑlÓnˆØ&×.Ñ.Ó0ÐØ(×0Ñ0Ó2°UÑ:Ðà+×.Ñ.Ð/@ÑA×EÑEÐFVÓWÐØ&×.Ñ.¨z¸:ÓFˆØÐr   Nrb   r   r   r   rd   rd       sa   „ ñ
ñ ;@À»,¸Ðbcñ )˜cð )°ð )Ð\_ó )ð #§+¡+ð °s·{±{ð ÈSð ÐUX×U`ÑU`ô r   rd   c                   óp   — e Zd ZdZdefd„Zdej                  dej                  dedej                  fd„Zy	)
Ú!FlaxForcedBOSTokenLogitsProcessorzÑ
    [`FlaxLogitsProcessor`] that enforces the specified token as the first generated token.

    Args:
        bos_token_id (`int`):
            The id of the token to force as the first generated token.
    Úbos_token_idc                 ó   — || _         y r)   )rw   )r   rw   s     r   r?   z*FlaxForcedBOSTokenLogitsProcessor.__init__Ì   s
   € Ø(ˆÕr   r	   r
   r%   r   c                 ó  — t        j                  |j                  t        d«       «      }dt        j                  |dz
  «      z
  }t        j
                  ||j                  d d …| j                  f   j                  d«      |«      }|S ©NÚinfrE   r   )	r   rh   rP   r>   Úbool_rY   rW   rw   rX   ©r   r	   r
   r%   Ú
new_scoresÚapply_penaltys         r   r   z*FlaxForcedBOSTokenLogitsProcessor.__call__Ï   sk   € Ü—X‘X˜fŸl™l¬U°5«\¨MÓ:ˆ
àœCŸI™I g°¡kÓ2Ñ2ˆä—‘˜=¨*¯-©-º¸4×;LÑ;LÐ8LÑ*M×*QÑ*QÐRSÓ*TÐV\Ó]ˆàˆr   N©	r   r   r   r   r8   r?   r   r   r   r   r   r   rv   rv   Ã   sC   „ ñð) Só )ð #§+¡+ð °s·{±{ð ÈSð ÐUX×U`ÑU`ô r   rv   c                   ót   — e Zd ZdZdedefd„Zdej                  dej                  dedej                  fd	„Zy
)Ú!FlaxForcedEOSTokenLogitsProcessorae  
    [`FlaxLogitsProcessor`] that enforces the specified token as the last generated token when `max_length` is reached.

    Args:
        max_length (`int`):
            The maximum length of the sequence to be generated.
        eos_token_id (`int`):
            The id of the token to force as the last generated token when `max_length` is reached.
    Ú
max_lengthÚeos_token_idc                 ó    — || _         || _        y r)   )rƒ   r„   )r   rƒ   r„   s      r   r?   z*FlaxForcedEOSTokenLogitsProcessor.__init__ä   s   € Ø$ˆŒØ(ˆÕr   r	   r
   r%   r   c                 ó,  — t        j                  |j                  t        d«       «      }dt        j                  || j
                  z
  dz   «      z
  }t        j                  ||j                  d d …| j                  f   j                  d«      |«      }|S rz   )
r   rh   rP   r>   r|   rƒ   rY   rW   r„   rX   r}   s         r   r   z*FlaxForcedEOSTokenLogitsProcessor.__call__è   su   € Ü—X‘X˜fŸl™l¬U°5«\¨MÓ:ˆ
àœCŸI™I g°·±Ñ&?À!Ñ&CÓDÑDˆä—‘˜=¨*¯-©-º¸4×;LÑ;LÐ8LÑ*M×*QÑ*QÐRSÓ*TÐV\Ó]ˆàˆr   Nr€   r   r   r   r‚   r‚   Ù   sJ   „ ñð) 3ð )°có )ð #§+¡+ð °s·{±{ð ÈSð ÐUX×U`ÑU`ô r   r‚   c                   ót   — e Zd ZdZdedefd„Zdej                  dej                  dedej                  fd	„Zy
)ÚFlaxMinLengthLogitsProcessora3  
    [`FlaxLogitsProcessor`] enforcing a min-length by setting EOS probability to 0.

    Args:
        min_length (`int`):
            The minimum length below which the score of `eos_token_id` is set to `-float("Inf")`.
        eos_token_id (`int`):
            The id of the *end-of-sequence* token.
    Ú
min_lengthr„   c                 ó¬   — t        |t        «      r|dk  rt        d|› �«      ‚t        |t        «      r|dk  rt        d|› �«      ‚|| _        || _        y )Nr   z2`min_length` has to be a positive integer, but is z4`eos_token_id` has to be a positive integer, but is )r=   r8   r5   r‰   r„   )r   r‰   r„   s      r   r?   z%FlaxMinLengthLogitsProcessor.__init__ý   s\   € Ü˜*¤cÔ*¨j¸1ªnÜÐQÐR\ÐQ]Ð^Ó_Ð_ä˜,¬Ô,°¸qÒ0@ÜÐSÐT`ÐSaÐbÓcÐcà$ˆŒØ(ˆÕr   r	   r
   r%   r   c                 óê   — dt        j                  || j                  z
  dd«      z
  }t        j                  ||j                  d d …| j
                  f   j                  t        d«       «      |«      }|S )NrE   r   r{   )r   Úclipr‰   rY   rW   r„   rX   r>   ©r   r	   r
   r%   r   s        r   r   z%FlaxMinLengthLogitsProcessor.__call__  s`   € àœCŸH™H W¨t¯©Ñ%>ÀÀ1ÓEÑEˆä—‘˜=¨&¯)©)²A°t×7HÑ7HÐ4HÑ*I×*MÑ*MÌuÐUZË|ÈmÓ*\Ð^dÓeˆàˆr   Nr€   r   r   r   rˆ   rˆ   ò   sJ   „ ñð) 3ð )°có )ð #§+¡+ð °s·{±{ð ÈSð ÐUX×U`ÑU`ô r   rˆ   c                   ó"   — e Zd ZdZd„ Zdefd„Zy)Ú(FlaxSuppressTokensAtBeginLogitsProcessoraº  
    [`FlaxLogitsProcessor`] supressing a list of tokens as soon as the `generate` function starts generating using
    `begin_index` tokens. This should ensure that the tokens defined by `begin_suppress_tokens` are not sampled at the
    beginning of the generation.

    Args:
        begin_suppress_tokens (`List[int]`):
            Tokens to not sample.
        begin_index (`int`):
            Index where the tokens are suppressed.
    c                 ó2   — t        |«      | _        || _        y r)   )r3   Úbegin_suppress_tokensÚbegin_index)r   r‘   r’   s      r   r?   z1FlaxSuppressTokensAtBeginLogitsProcessor.__init__  s   € Ü%)Ð*?Ó%@ˆÔ"Ø&ˆÕr   r%   c                 óæ   — dt        j                  || j                  z
  «      z
  }t        j                  ||j                  d d …| j
                  f   j                  t        d«       «      |«      }|S )NrE   r{   )r   r|   r’   rY   rW   r‘   rX   r>   r�   s        r   r   z1FlaxSuppressTokensAtBeginLogitsProcessor.__call__!  sa   € ØœCŸI™I g°×0@Ñ0@Ñ&@ÓAÑAˆä—‘˜=¨&¯)©)²A°t×7QÑ7QÐ4QÑ*R×*VÑ*VÔX]Ð^cÓXdÐWdÓ*eÐgmÓnˆàˆr   N)r   r   r   r   r?   r8   r   r   r   r   r�   r�     s   „ ñ
ò'ð°3ô r   r�   c                   óp   — e Zd ZdZdefd„Zdej                  dej                  dedej                  fd„Z	y	)
Ú!FlaxSuppressTokensLogitsProcessorzõ
    [`FlaxLogitsProcessor`] suppressing a list of tokens at each decoding step. The processor will set their log probs
    to be `-inf` so they are not sampled.

    Args:
        suppress_tokens (`list`):
            Tokens to not sample.
    Úsuppress_tokensc                 ó$   — t        |«      | _        y r)   )r3   r–   )r   r–   s     r   r?   z*FlaxSuppressTokensLogitsProcessor.__init__3  s   € Ü# OÓ4ˆÕr   r	   r
   r%   r   c                 ón   — |j                   d| j                  f   j                  t        d«       «      }|S )N.r{   )rW   r–   rX   r>   rA   s       r   r   z*FlaxSuppressTokensLogitsProcessor.__call__6  s1   € Ø—‘˜3 × 4Ñ 4Ð4Ñ5×9Ñ9¼5À»<¸-ÓHˆàˆr   N)
r   r   r   r   r3   r?   r   r   r8   r   r   r   r   r•   r•   )  sC   „ ñð5¨ó 5ð #§+¡+ð °s·{±{ð ÈSð ÐUX×U`ÑU`ô r   r•   c                   ój   — e Zd ZdZd„ Zdej                  dej                  dedej                  fd„Zy)	ÚFlaxForceTokensLogitsProcessora½  
    [`FlaxLogitsProcessor`] that takes a list of pairs of integers which indicates a mapping from generation indices to
    token indices that will be forced before sampling. The processor will set their log probs to 0 and all other tokens
    to `-inf` so that they are sampled at their corresponding index.

    Args:
        force_token_map (`list`):
            Map giving token ids and indices where they will be forced to be sampled.
    c                 óD  — t        |«      }t        j                  t        |j	                  «       «      dz   t        j
                  ¬«      dz  }|j                  «       D ]&  \  }}|€Œ	|j                  |   j                  |«      }Œ( t        j
                  |«      | _	        y )NrE   ©ÚdtyperK   )
Údictr   Úonesrf   r4   Úint32ÚitemsrW   rX   Úforce_token_array)r   Úforce_token_mapr¢   ÚindexÚtokens        r   r?   z'FlaxForceTokensLogitsProcessor.__init__G  s�   € Ü˜Ó/ˆô  ŸH™H¤c¨/×*>Ñ*>Ó*@Ó&AÀAÑ&EÌcÏiÉiÔXÐ[]Ñ]ÐØ+×1Ñ1Ó3ò 	K‰LˆE�5ØÑ Ø$5×$8Ñ$8¸Ñ$?×$CÑ$CÀEÓ$JÑ!ð	Kô "%§¡Ð+<Ó!=ˆÕr   r	   r
   r%   r   c                 óŽ   ‡ ‡‡‡— ˆˆ fd„Št        j                  ‰‰ j                  j                  d   k\  ˆfd„ˆˆˆˆ fd„«      Š‰S )Nc                 ó  •— ‰j                   d   }‰j                  |    }t        j                  ‰‰j                  ¬«      t        d«       z  }t        j                  |df‰j                  ¬«      }t        j                  ||d|f«      }|S )Nr   rœ   r{   rE   )	rP   r¢   r   Ú	ones_liker�   r>   ÚzerosrN   Údynamic_update_slice)Úgeneration_idxrn   Úcurrent_tokenr~   Úupdatesr
   r   s        €€r   Ú_force_tokenz=FlaxForceTokensLogitsProcessor.__call__.<locals>._force_tokenS  sv   ø€ ØŸ™ a™ˆJØ ×2Ñ2°>ÑBˆMäŸ™ v°V·\±\ÔBÄeÈEÃlÀ]ÑRˆJÜ—i‘i ¨Q °v·|±|ÔDˆGÜ×1Ñ1°*¸gÈÈ=ÐGYÓZˆJØÐr   r   c                  ó   •— ‰ S r)   r   ©r
   s   €r   ú<lambda>z9FlaxForceTokensLogitsProcessor.__call__.<locals>.<lambda>_  s   ø€ �F€ r   c                  ó`   •— t        j                  ‰j                  ‰   dk\  ˆ ˆfd„ˆfd„«      S )Nr   c                  ó   •—  ‰ ‰«      S r)   r   )r®   r%   s   €€r   r±   zKFlaxForceTokensLogitsProcessor.__call__.<locals>.<lambda>.<locals>.<lambda>d  s   ø€ ™ WÓ-€ r   c                  ó   •— ‰ S r)   r   r°   s   €r   r±   zKFlaxForceTokensLogitsProcessor.__call__.<locals>.<lambda>.<locals>.<lambda>f  s   ø€ ˜€ r   )rN   Úcondr¢   )r®   r%   r
   r   s   €€€€r   r±   z9FlaxForceTokensLogitsProcessor.__call__.<locals>.<lambda>a  s)   ø€ ”C—H‘HØ×&Ñ& wÑ/°1Ñ4ä-ãó€ r   )rN   rµ   r¢   rP   )r   r	   r
   r%   r®   s   ` ``@r   r   z'FlaxForceTokensLogitsProcessor.__call__R  s@   û€ õ	ô —‘Ø�t×-Ñ-×3Ñ3°AÑ6Ñ6ãöó
ˆð ˆr   N)	r   r   r   r   r?   r   r   r8   r   r   r   r   rš   rš   <  s<   „ ñò	>ð #§+¡+ð °s·{±{ð ÈSð ÐUX×U`ÑU`ô r   rš   c                   ó   — e Zd ZdZd„ Zd„ Zy)Ú#FlaxWhisperTimeStampLogitsProcessora{  
    Whisper specific Processor. This processor can be used to force a list of tokens. The processor will set their log
    probs to `inf` so that they are sampled at their corresponding index.

    Args:
        generate_config (`GenerateConfig`):
            The generate config used to generate the output. The following parameters are required:
                eos_token_id (`int`, *optional*, defaults to 50257):
                    The id of the *end-of-sequence* token.
                no_timestamps_token_id (`int`, *optional*, defaults to 50363):
                    The id of the `"<|notimestamps|>"` token.
                max_initial_timestamp_index (`int`, *optional*, defaults to 1):
                    Used to set the maximum value of the initial timestamp. This is used to prevent the model from
                    predicting timestamps that are too far in the future.
    c                 ó`  — |j                   | _         |j                  | _        |j                  dz   | _        |dz   | _        |j                  r| xj                  dz  c_        t        |d«      r|j                  | _        n|j                  | _        | j                  €|j                  | _        y y )NrE   r   Úmax_initial_timestamp_index)r„   Úno_timestamps_token_idÚtimestamp_beginr’   Úis_multilingualÚhasattrr¹   ro   )r   Úgenerate_configÚmodel_configÚdecoder_input_lengths       r   r?   z,FlaxWhisperTimeStampLogitsProcessor.__init__}  sž   € Ø+×8Ñ8ˆÔØ&5×&LÑ&LˆÔ#Ø.×EÑEÈÑIˆÔà/°!Ñ3ˆÔà×*Ò*à×Ò Ñ!ÕÜ�?Ð$AÔBØ/>×/ZÑ/ZˆDÕ,à/;×/FÑ/FˆDÔ,Ø×+Ñ+Ð3Ø/;×/FÑ/FˆDÕ,ð 4r   c                 óŠ  ‡ ‡— |j                   d d …‰ j                  f   j                  t        d«       «      }ˆˆ fd„} t	        j
                  |«      ||«      }t        j                  ‰‰ j                  k(  dd«      }t        j                  ‰ j                  d u|d«      }‰ j                  ‰ j                  z   }t        j                  ||j                   d d …|dz   d …f   j                  t        d«       «      |«      }t        j                  j                  |d¬«      }ˆ fd„} t	        j
                  |«      ||«      }|S )	Nr{   c                 óf  •— t        j                  ‰‰j                  z
  dk\  dd«      }t        j                  | ‰dz
     ‰j                  k\  |d«      }t        j                  ‰‰j                  z
  dk  dd«      }t        j                  | ‰dz
     ‰j                  k\  d|«      }t        j                  |t        j                  |dkD  |j                  ‰j                  d  j                  t        d«       «      |j                  d ‰j                   j                  t        d«       «      «      |«      S )NrE   TFr   r   r{   )r   rY   r’   r»   rW   rX   r>   r„   )Úinput_ids_kÚscores_kÚlast_was_timestampÚpenultimate_was_timestampr%   r   s       €€r   Úhandle_pairszBFlaxWhisperTimeStampLogitsProcessor.__call__.<locals>.handle_pairs’  s"  ø€ Ü!$§¡¨G°d×6FÑ6FÑ,FÈ1Ñ+LÈdÐTYÓ!ZÐÜ!$§¡Ø˜G a™KÑ(¨D×,@Ñ,@Ñ@Ø+Øó"Ðô ),¯	©	°7¸T×=MÑ=MÑ3MÐQRÑ2RÐTXÐZ_Ó(`Ð%Ü(+¯	©	Ø˜G a™KÑ(¨D×,@Ñ,@Ñ@ØØ)ó)Ð%ô —9‘9Ø"Ü—	‘	Ø-°Ñ1Ø—K‘K × 4Ñ 4Ð 6Ð7×;Ñ;¼UÀ5»\¸MÓJØ—K‘KÐ 3 $×"3Ñ"3Ð4×8Ñ8¼%À»,¸ÓGóð
 óð r   TFrE   rK   rL   c                 ó8  •— t         j                  j                  | ‰j                  d  d¬«      }t	        j
                  | d ‰j                   «      }t	        j                  ||kD  |j                  d ‰j                   j                  t        d«       «      |«      S )NrK   rL   r{   )
rR   rS   Ú	logsumexpr»   r   rf   rY   rW   rX   r>   )Ú
logprobs_krÄ   Útimestamp_logprobÚmax_text_token_logprobr   s       €r   Úhandle_cumulative_probszMFlaxWhisperTimeStampLogitsProcessor.__call__.<locals>.handle_cumulative_probs¿  sŒ   ø€ Ü #§¡× 0Ñ 0°¸D×<PÑ<PÐ<RÐ1SÐZ\Ð 0Ó ]ÐÜ%(§W¡W¨ZÐ8N¸$×:NÑ:NÐ-OÓ%PÐ"Ü—9‘9Ø!Ð$:Ñ:Ø—‘Ð2˜d×2Ñ2Ð3×7Ñ7¼¸u»¸ÓFØóð r   )rW   rº   rX   r>   rR   Úvmapr   rY   r’   r¹   r»   rS   Úlog_softmax)	r   r	   r
   r%   rÇ   Úapply_max_initial_timestampÚlast_allowedÚlogprobsrÍ   s	   `  `     r   r   z,FlaxWhisperTimeStampLogitsProcessor.__call__Ž  s"  ù€ à—‘š1˜d×9Ñ9Ð9Ñ:×>Ñ>ÄÀeÃ¸}ÓMˆõ	ð2 (”—‘˜,Ó'¨	°6Ó:ˆä&)§i¡i°¸4×;KÑ;KÑ0KÈTÐSXÓ&YÐ#Ü&)§i¡iØ×,Ñ,°DÐ8Ø0Øó'
Ð#ð ×+Ñ+¨d×.NÑ.NÑNˆä—‘Ø'Ø�I‰I’a˜¨Ñ)Ñ+Ð+Ñ,×0Ñ0´%¸³,°Ó?Øó
ˆô —6‘6×%Ñ% f°2Ð%Ó6ˆô	ð 3”—‘Ð1Ó2°8¸VÓDˆàˆr   N)r   r   r   r   r?   r   r   r   r   r·   r·   l  s   „ ñò Gó"<r   r·   c                   óÐ   — e Zd ZdZdefd„Zdej                  dedefd„Zdej                  d	ej                  fd
„Z	dej                  dej                  ded	ej                  fd„Z
y)Ú FlaxNoRepeatNGramLogitsProcessora9  
    [`FlaxLogitsProcessor`] that enforces no repetition of n-grams. See
    [Fairseq](https://github.com/pytorch/fairseq/blob/a07cb6f40480928c9e0548b737aadd36ee66ac76/fairseq/sequence_generator.py#L345).

    Args:
        ngram_size (`int`):
            All ngrams of size `ngram_size` can only occur once.
    Ú
ngram_sizec                 óX   — t        |t        «      r|dk  rt        d|› �«      ‚|| _        y )Nr   z;`ngram_size` has to be a strictly positive integer, but is )r=   r8   r5   rÕ   )r   rÕ   s     r   r?   z)FlaxNoRepeatNGramLogitsProcessor.__init__×  s.   € Ü˜*¤cÔ*¨j¸AªoÜÐZÐ[eÐZfÐgÓhÐhØ$ˆ�r   r	   ro   r%   c           	      óÜ  ‡ ‡‡— ‰j                   \  Š}|‰ j                  dz
  z
  }|‰ j                  dz
  z
  }ˆˆˆ fd„}‰|z  ‰ j                  dz   f}t        j                  j	                  d‰|z  |t        j                  |‰j                  ¬«      «      }	t        j                  ‰|z  «      ‰|z  k  j                  d«      }
t        j                  |
|	f‰f|f‰ j                  z  z   ¬«      S )a  
        get a matrix of size (batch_size,) + (vocab_size,)*n (for n-grams) that
        represent the n-grams that occurred previously.
        The BCOO representation allow to store only the few non-zero entries, instead of the full (huge) matrix
        rE   c                 ó  •— | ‰z  }| ‰z  }|j                   |    j                  t        j                  |gt	        ‰j
                  «      D �cg c]  }t        j                  ‰«      |||z   f   ‘Œ! c}z   «      «      S c c}w r)   )rW   rX   r   ÚarrayÚrangerÕ   )ÚiÚvalÚbÚposÚjrn   r	   r   s        €€€r   Úbody_funzFFlaxNoRepeatNGramLogitsProcessor.get_previous_ngrams.<locals>.body_funè  s   ø€ Ø�J‘ˆAØ�z‘/ˆCØ—6‘6˜!‘9—=‘=Ü—	‘	àðô BGÀtÇÁÓAWÖX¸A”s—y‘y Ó+¨A¨s°Q©w¨JÓ7ÒXñYóóð ùò
 Ys   Á$A=r   rœ   Úfloat32)rP   )rP   rÕ   rR   rN   Ú	fori_loopr   r©   r�   rk   Úastyper   ÚBCOO)r   r	   ro   r%   Úseq_lenÚ
seq_ngramsÚ
cur_ngramsrà   rP   Úall_update_indicesÚdatarn   s   ``         @r   Úget_previous_ngramsz4FlaxNoRepeatNGramLogitsProcessor.get_previous_ngramsÜ  sá   ú€ ð (Ÿo™oÑˆ
�Gà §¡°!Ñ 3Ñ4ˆ
à §¡°!Ñ 3Ñ4ˆ
ö
	ð ˜jÑ(¨$¯/©/¸AÑ*=Ð>ˆÜ ŸW™W×.Ñ.Øˆz˜JÑ&¨´#·)±)¸EÈÏÉÔ2Yó
Ðô
 —
‘
˜:¨
Ñ2Ó3°jÀ:Ñ6MÑM×UÑUÐV_Ó`ˆä�{‰{˜DÐ"4Ð5¸j¸]ÈjÈ]Ð]a×]lÑ]lÑMlÑ=lÔmÐmr   Úlatest_tokensr   c                 óŒ   — t         j                  t        j                  d„ «       «       }t        j                   |||«      «      S )zt
        Determines which tokens must be banned given latest tokens and the previously seen
        ngrams.
        c                 ó   — |t        | «         S r)   )Útuple)rë   Úprevious_ngramss     r   Úinner_fnzIFlaxNoRepeatNGramLogitsProcessor.get_banned_tokens_mask.<locals>.inner_fn  s   € ð #¤5¨Ó#7Ñ8Ð8r   )r   ÚsparsifyrR   rÎ   Úbcoo_todense)r   rë   rï   rð   s       r   Úget_banned_tokens_maskz7FlaxNoRepeatNGramLogitsProcessor.get_banned_tokens_maskþ  s@   € ô 
�‰Ü	�‰ñ	9ó 
ó 
ð	9ô ×"Ñ"¡8¨M¸?Ó#KÓLÐLr   r
   c                 ó†   ‡ ‡‡‡— ˆˆˆˆ fd„}t         j                  j                  ‰‰ j                  dz
  k\  |ˆfd„«      }|S )Nc            
      ó"  •— ‰j                   \  } }‰j                  ‰|‰«      }t        j                  ‰j                   d   ‰j                  dz
  f‰j
                  ¬«      }t        j                  j                  |t        j                  j                  ‰d‰‰j                  dz
  z
  f‰j                   d   ‰j                  dz
  f«      d«      }‰j                  ||«      j                  d«      }t        j                  |t        d«       ‰«      S )Nr   rE   rœ   )r   r   Úboolr{   )rP   rê   r   r©   rÕ   r�   rR   rN   rª   Údynamic_sliceró   rã   rY   r>   )	Ú_ro   rï   rë   Úbanned_tokens_indices_maskr%   r	   r
   r   s	        €€€€r   Útrue_fnz:FlaxNoRepeatNGramLogitsProcessor.__call__.<locals>.true_fn  sô   ø€ Ø"ŸL™L‰MˆAˆzà"×6Ñ6°yÀ*ÈgÓVˆOô  ŸI™I y§¡°qÑ'9¸4¿?¹?ÈQÑ;NÐ&OÐW`×WfÑWfÔgˆMÜŸG™G×8Ñ8ØÜ—‘×%Ñ%Ø  7¨d¯o©oÀÑ.AÑ#BÐCÀiÇoÁoÐVWÑFXÐ[_×[jÑ[jÐmnÑ[nÐEpóð óˆMð *.×)DÑ)DÀ]ÐTcÓ)d×)kÑ)kÐlrÓ)sÐ&Ü—9‘9Ð7¼%À»,¸ÈÓOÐOr   rE   c                  ó   •— ‰ S r)   r   r°   s   €r   r±   z;FlaxNoRepeatNGramLogitsProcessor.__call__.<locals>.<lambda>  s   ø€ ÐQW€ r   )rR   rN   rµ   rÕ   )r   r	   r
   r%   rú   Úoutputs   ````  r   r   z)FlaxNoRepeatNGramLogitsProcessor.__call__  s4   û€ ÷	Pô& —‘—‘˜w¨$¯/©/¸AÑ*=Ñ=ÀËÓXˆØˆr   N)r   r   r   r   r8   r?   r   r   rê   ró   r   r   r   r   rÔ   rÔ   Í  sˆ   „ ñð% 3ó %ð
 n¨S¯[©[ð  nÀcð  nÐTWó  nðDM°C·K±Kð MÐUX×U`ÑU`ó Mð #§+¡+ð °s·{±{ð ÈSð ÐUX×U`ÑU`ô r   rÔ   )r.   rR   Újax.laxrN   Ú	jax.numpyÚnumpyr   Újax.experimentalr   Úutilsr   Úutils.loggingr   r   Úloggerr   r   r!   r3   r$   r:   rC   rd   rv   r‚   rˆ   r�   r•   rš   r·   rÔ   r   r   r   ú<module>r     sç   ðó  ã 
Ý Ý Ý #å (Ý &ñ 
�HÓ	€ð%Ð !÷*
ñ 
÷
ñ 
ô˜dô ô.Ð"2ô ô()Ð+ô )ôX Ð+ô  ôFÐ(;ô ô,Ð(;ô ô2Ð#6ô ô<Ð/Bô ô2Ð(;ô ô&-Ð%8ô -ô`^Ð*=ô ^ôBSÐ':õ Sr   