§
    ‚Štj²  ã                   ó¾   — d dl Z ddlmZmZ  ej        e¦  «        Z e¦   «         rd dlZ	 dde j        de j        de j        de j        d	e	d
e
de j        fd„Z	 dd„ZdS )é    Né   )Úis_torchaudio_availableÚloggingÚmean_volumeÚlogitsÚtargetsÚlogit_lengthsÚtarget_lengthsÚblank_token_idÚ	reductionÚreturnc           	      óŒ  — t          ¦   «         st          d¦  «        ‚d}||vr3t          d|› dd                     d„ |D ¦   «         ¦  «        › d�¦  «        ‚|                     | j        ¦  «        }t          j                             |  	                    ¦   «          
                    ¦   «         |                     | j        ¦  «                             ¦   «         |                     | j        ¦  «                             ¦   «         |                     ¦   «         |d¬	¦  «        }|d
k    r;|                     ¦   «         | 	                    ¦   «                              ¦   «         z  S |dk    r|                     ¦   «         S |dk    r)|| 	                    ¦   «         z                       ¦   «         S |dk    r|                     ¦   «         S |S )a  
    Compute standard RNN-T (RNN Transducer) loss (https://huggingface.co/papers/1211.3711).

    Thin wrapper around [`torchaudio.functional.rnnt_loss`]. torchaudio is queried with `reduction="none"` to get
    the per-sample negative log-likelihoods, and the requested reduction is applied here. The reduction names and
    formulas mirror NeMo's `RNNTLoss` (the reference implementation used to train/finetune Parakeet), so that loss
    magnitudes and gradient scaling match when finetuning other RNNT models like Parakeet:

    - `"mean_volume"`: sum of per-sample losses divided by the sum of target lengths (per-token average over the
      whole batch). This is what `nvidia/parakeet-rnnt-0.6b` is trained with (`rnnt_reduction: mean_volume`).
    - `"mean_batch"`: plain average of per-sample losses over the batch (NeMo's default).
    - `"mean"`: per-sample loss divided by its own target length, then averaged over the batch.
    - `"sum"`: sum of per-sample losses.
    - `"none"`: per-sample losses, unreduced.

    Args:
        logits: Joint token logits of shape `(batch, T, U+1, vocab_size)`.
        targets: Target labels of shape `(batch, U)`.
        logit_lengths: Encoder output lengths of shape `(batch,)`.
        target_lengths: Target lengths of shape `(batch,)`.
        blank_token_id: Blank token id.
        reduction: Loss reduction method. One of `"mean_volume"`, `"mean_batch"`, `"mean"`, `"sum"`, or `"none"`.

    Returns:
        Scalar loss tensor (or per-example losses if `reduction="none"`).

    zWComputing the RNN-T loss requires torchaudio. Install it with `pip install torchaudio`.)r   Ú
mean_batchÚmeanÚsumÚnonezInvalid reduction mode "z". Expected one of z, c              3   ó4   K  — | ]}t          |¦  «        V — Œd S )N)Úrepr)Ú.0Úrs     úY/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/loss/loss_rnnt.pyú	<genexpr>zrnnt_loss.<locals>.<genexpr>D   s*   è è € ÐNqÐNqÐ[\ÍtÐTUÉwÌwÐNqÐNqÐNqÐNqÐNqÐNqó    ú.r   )r   r   r	   r
   Úblankr   r   r   r   r   )r   ÚImportErrorÚ
ValueErrorÚjoinÚtoÚdeviceÚ
torchaudioÚ
functionalÚ	rnnt_lossÚfloatÚ
contiguousÚintr   r   )r   r   r	   r
   r   r   Úvalid_reductionsÚlossess           r   r#   r#      s²  € õH #Ñ$Ô$ð uÝÐsÑtÔtÐtàKÐØÐ(Ð(Ð(ÝØt yÐtÐtÀTÇYÂYÐNqÐNqÐ`pÐNqÑNqÔNqÑEqÔEqÐtÐtÐtñ
ô 
ð 	
ð $×&Ò& v¤}Ñ5Ô5€NÝÔ"×,Ò,Ø�|Š|‰~Œ~×(Ò(Ñ*Ô*Ø—
’
˜6œ=Ñ)Ô)×-Ò-Ñ/Ô/Ø#×&Ò& v¤}Ñ5Ô5×9Ò9Ñ;Ô;Ø%×)Ò)Ñ+Ô+ØØð -ñ ô €Fð �MÒ!Ð!Ø�zŠz‰|Œ|˜n×2Ò2Ñ4Ô4×8Ò8Ñ:Ô:Ñ:Ð:Ø	�lÒ	"Ð	"Ø�{Š{‰}Œ}ÐØ	�fÒ	Ð	Ø˜×-Ò-Ñ/Ô/Ñ/×5Ò5Ñ7Ô7Ð7Ø	�eÒ	Ð	Ø�zŠz‰|Œ|ÐØ€Mr   c                 ó,   — t          | |||||¬¦  «        S )N)r   r   r	   r
   r   r   )r#   )r   Úlabelsr	   Úlabel_lengthsr   r   Úkwargss          r   ÚParakeetForRNNTLossr-   \   s-   € õ ØØØ#Ø$Ø%Øðñ ô ð r   )r   )ÚtorchÚutilsr   r   Ú
get_loggerÚ__name__Úloggerr!   ÚTensorr&   Ústrr#   r-   © r   r   ú<module>r6      sç   ðð €€€à 4Ð 4Ð 4Ð 4Ð 4Ð 4Ð 4Ð 4ð 
ˆÔ	˜HÑ	%Ô	%€àÐÑÔð ØÐÐÐð #ð?ð ?ØŒLð?àŒ\ð?ð ”<ð?ð ”Lð	?ð
 ð?ð ð?ð „\ð?ð ?ð ?ð ?ðP ðð ð ð ð ð r   