Ë
    óÿæi'2  ã                   ó¤  — d dl mZmZmZmZmZ d dlZd dlmZ ddgZ	eee
   ej                  eeej                        ef   Zde_        dedee
   fd	„Zdedej                  fd
„Zdedeeej                        fd„Zdedefd„Zdedefd„Zdee   deeej                        fd„Zdeeej                        de
dej,                  deeej                        fd„Zdedefd„Zdee   dej                  de
deej                  ej                  ej                  f   fd„Zdedee   ddfd„Z G d„ dej6                  j8                  «      Zy)é    )ÚCallableÚDictÚListÚOptionalÚTupleN)ÚRNNTÚ
HypothesisÚRNNTBeamSearchz™Hypothesis generated by RNN-T beam search decoder,
    represented as tuple of (tokens, prediction network output, prediction network state, score).
    ÚhypoÚreturnc                 ó   — | d   S ©Nr   © ©r   s    ús/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchaudio/models/rnnt_decoder.pyÚ_get_hypo_tokensr      ó   € Ø�‰7€Nó    c                 ó   — | d   S ©Né   r   r   s    r   Ú_get_hypo_predictor_outr      r   r   c                 ó   — | d   S )Né   r   r   s    r   Ú_get_hypo_stater      r   r   c                 ó   — | d   S )Né   r   r   s    r   Ú_get_hypo_scorer      r   r   c                 ó   — t        | d   «      S r   )Ústrr   s    r   Ú_get_hypo_keyr!       s   € Üˆt�A‰w‹<Ðr   Úhyposc                 óV  — g }t        t        t        | d   «      «      «      D ]~  }g }t        t        t        | d   «      |   «      «      D ]C  }|j                  t	        j
                  | D �cg c]  }t        |«      |   |   ‘Œ c}«      «       ŒE |j                  |«       Œ€ |S c c}w r   )ÚrangeÚlenr   ÚappendÚtorchÚcat)r"   ÚstatesÚiÚbatched_state_componentsÚjr   s         r   Ú_batch_stater-   $   s    € Ø')€FÜ”3” u¨Q¡xÓ0Ó1Ö2ˆØ79Ð Ü”sœ?¨5°©8Ó4°QÑ7Ó8Ö9ˆAØ$×+Ñ+¬E¯I©IÑ_dÓ6eÑ_dÐW[´ÀtÓ7LÈQÑ7OÐPQÓ7RÐ_dÑ6eÓ,fÕgð :à�‰Ð.Õ/ð	 3ð
 €Mùò 7fs   Á,B&r)   ÚidxÚdevicec                 ó¨   — t        j                  |g|¬«      }| D ��cg c]"  }|D �cg c]  }|j                  d|«      ‘Œ c}‘Œ$ c}}S c c}w c c}}w )N©r/   r   )r'   ÚtensorÚindex_select)r)   r.   r/   Ú
idx_tensorÚstate_tupleÚstates         r   Ú_slice_stater7   .   sM   € Ü—‘˜s˜e¨FÔ3€JÙ\bÔcÑ\bÈ[¹KÓH¹K°5ˆU×Ñ  :Õ.¸KÓHÐ\bÒcÐcùÒHùÓcs   ž	A§A	Á AÁ	Ac                 óH   — t        | «      t        t        | «      «      dz   z  S r   )r   r%   r   r   s    r   Ú_default_hypo_sort_keyr9   3   s"   € Ü˜4Ó ¤CÔ(8¸Ó(>Ó$?À!Ñ$CÑDÐDr   Únext_token_probsÚ
beam_widthc                 óR  — t        j                  | D �cg c]  }t        |«      ‘Œ c}«      j                  d«      }||d d …d d…f   z   }|j	                  d«      j                  |«      \  }}|j                  |j                  d   d¬«      }||j                  d   z  }	|||	fS c c}w )Nr   éÿÿÿÿÚtrunc)Úrounding_mode)r'   r2   r   Ú	unsqueezeÚreshapeÚtopkÚdivÚshape)
r"   r:   r;   ÚhÚhypo_scoresÚnonblank_scoresÚnonblank_nbest_scoresÚnonblank_nbest_idxÚnonblank_nbest_hypo_idxÚnonblank_nbest_tokens
             r   Ú_compute_updated_scoresrL   7   s¹   € ô
 —,‘,¹EÓB¹E°q¤°Õ 2¸EÑBÓC×MÑMÈaÓP€KØ!Ð$4²Q¸¸¸°VÑ$<Ñ<€OØ0?×0GÑ0GÈÓ0K×0PÑ0PÐQ[Ó0\Ñ-ÐÐ-Ø0×4Ñ4°_×5JÑ5JÈ1Ñ5MÐ]dÐ4ÓeÐØ-°×0EÑ0EÀaÑ0HÑHÐØ Ð"9Ð;OÐOÐOùò  Cs   ”B$Ú	hypo_listc                 ób   — t        |«      D ]!  \  }}t        | «      t        |«      k(  sŒ||=  y  y ©N)Ú	enumerater!   )r   rM   r*   Úelems       r   Ú_remove_hyporR   D   s1   € Ü˜YÖ'‰ˆˆ4Ü˜Ó¤-°Ó"5Ó5Ø˜!�Ùñ (r   c                   ó6  ‡ — e Zd ZdZ	 	 	 d#dedededeee	gef      deddfˆ fd	„Z
d
ej                  dee	   fd„Zdej                  dee	   d
ej                  dej                  fd„Zdee	   dee	   dej                  deee	f   dee	   f
d„Zdee	   dee	   dej                  deded
ej                  dee	   fd„Zdee	   dee   dee   ded
ej                  dee	   fd„Zdej                  deee	      dedee	   fd„Zdej                  dej                  dedee	   fd„Zej0                  j2                  	 	 d$dej                  dej                  ded eeeej                           d!eee	      deee	   eeej                        f   fd"„«       Zˆ xZS )%r
   a)  Beam search decoder for RNN-T model.

    See Also:
        * :class:`torchaudio.pipelines.RNNTBundle`: ASR pipeline with pretrained model.

    Args:
        model (RNNT): RNN-T model to use.
        blank (int): index of blank token in vocabulary.
        temperature (float, optional): temperature to apply to joint network output.
            Larger values yield more uniform samples. (Default: 1.0)
        hypo_sort_key (Callable[[Hypothesis], float] or None, optional): callable that computes a score
            for a given hypothesis to rank hypotheses by. If ``None``, defaults to callable that returns
            hypothesis score normalized by token sequence length. (Default: None)
        step_max_tokens (int, optional): maximum number of tokens to emit per input time step. (Default: 100)
    NÚmodelÚblankÚtemperatureÚhypo_sort_keyÚstep_max_tokensr   c                 ó’   •— t         ‰| �  «        || _        || _        || _        |€t
        | _        || _        y || _        || _        y rO   )ÚsuperÚ__init__rT   rU   rV   r9   rW   rX   )ÚselfrT   rU   rV   rW   rX   Ú	__class__s         €r   r[   zRNNTBeamSearch.__init__\   sP   ø€ ô 	‰ÑÔØˆŒ
ØˆŒ
Ø&ˆÔàÐ Ü!7ˆDÔð  /ˆÕð "/ˆDÔà.ˆÕr   r/   c                 óô   — | j                   }d }t        j                  dg|¬«      }| j                  j	                  t        j                  |gg|¬«      ||«      \  }}}|g|d   j                  «       |df}|gS )Nr   r1   r   g        )rU   r'   r2   rT   ÚpredictÚdetach)	r\   r/   Útokenr6   Ú
one_tensorÚpred_outÚ_Ú
pred_stateÚ	init_hypos	            r   Ú_init_b_hyposzRNNTBeamSearch._init_b_hyposp   s   € Ø—
‘
ˆØˆä—\‘\ 1 #¨fÔ5ˆ
Ø"&§*¡*×"4Ñ"4´U·\±\ÀEÀ7À)ÐTZÔ5[Ð]gÐinÓ"oÑˆ�!�ZàˆGØ�Q‰K×ÑÓ ØØð	
ˆ	ð ˆ{Ðr   Úenc_outr"   c                 ó¤  — t        j                  dg|¬«      }t        j                  |D �cg c]  }t        |«      ‘Œ c}d¬«      }| j                  j                  |||t        j                  dgt        |«      z  |¬«      «      \  }}}t         j                  j                  j                  || j                  z  d¬«      }|d d …ddf   S c c}w )Nr   r1   r   )Údimr   )r'   r2   Ústackr   rT   Újoinr%   ÚnnÚ
functionalÚlog_softmaxrV   )	r\   rh   r"   r/   rb   rE   Úpredictor_outÚ
joined_outrd   s	            r   Ú_gen_next_token_probsz$RNNTBeamSearch._gen_next_token_probs~   s»   € ô —\‘\ 1 #¨fÔ5ˆ
ÜŸ™ÉÓ$OÉÀAÔ%<¸QÕ%?ÈÑ$OÐUVÔWˆØŸ:™:Ÿ?™?ØØØÜ�L‰L˜!˜œs 5›zÑ)°&Ô9ó	
Ñˆ
�A�qô —X‘X×(Ñ(×4Ñ4°ZÀ$×BRÑBRÑ5RÐXYÐ4ÓZˆ
Øš!˜Q ˜'Ñ"Ð"ùò %Ps   ¬CÚb_hyposÚa_hyposr:   Úkey_to_b_hypoc                 ój  — t        t        |«      «      D ]Ã  }||   }t        |«      ||df   z   }t        |«      |v rQ|t        |«         }t	        ||«       t        t        j                  t        |«      «      j                  |«      «      }	nt        |«      }	t        |«      t        |«      t        |«      |	f}|j                  |«       ||t        |«      <   ŒÅ t        j                  |D �
cg c]  }
t        |
«      ‘Œ c}
«      j                  «       \  }}|D �cg c]  }||   ‘Œ	 c}S c c}
w c c}w )Nr=   )r$   r%   r   r!   rR   Úfloatr'   r2   Ú	logaddexpr   r   r   r&   Úsort)r\   rs   rt   r:   ru   r*   Úh_aÚappend_blank_scoreÚh_bÚscorer   rd   Ú
sorted_idxr.   s                 r   Ú_gen_b_hyposzRNNTBeamSearch._gen_b_hyposŒ   s!  € ô ”s˜7“|Ö$ˆAØ˜!‘*ˆCÜ!0°Ó!5Ð8HÈÈBÈÑ8OÑ!OÐÜ˜SÓ! ]Ñ2Ø#¤M°#Ó$6Ñ7�Ü˜S 'Ô*ÜœeŸl™l¬?¸3Ó+?Ó@×JÑJÐK]Ó^Ó_‘äÐ0Ó1�ä  Ó%Ü'¨Ó,Ü Ó$Øð	ˆCð �N‰N˜3ÔØ03ˆMœ-¨Ó,Ò-ð! %ô" Ÿ™ÉÓ%PÉÀ¤o°dÕ&;ÈÑ%PÓQ×VÑVÓX‰ˆˆ:Ù(2Ó3©
 �˜“¨
Ñ3Ð3ùò &QùÚ3s   Ã.D+ÄD0Útr;   c                 ó¦  — t        |||«      \  }}}	t        |«      |k  rt        d«       }
nt        ||    «      }
g }g }g }t	        |«      D ]f  }t        ||   «      }||
kD  sŒt        ||   «      }|j                  ||   «       |j                  t        |	|   «      «       |j                  |«       Œh |r| j                  |||||«      }|S g }|S )NÚinf)rL   r%   rw   r   r$   Úintr&   Ú_gen_new_hypos)r\   rt   rs   r:   r€   r;   r/   rH   rJ   rK   Úb_nbest_scoreÚ
base_hyposÚ
new_tokensÚ
new_scoresr*   r}   Ú
a_hypo_idxÚ	new_hyposs                     r   Ú_gen_a_hyposzRNNTBeamSearch._gen_a_hypos§   sý   € ô $ GÐ-=¸zÓJñ		
Ø!Ø#Ø ô ˆw‹<˜*Ò$Ü" 5›\˜M‰Mä+¨G°Z°KÑ,@ÓAˆMà')ˆ
Ø "ˆ
Ø"$ˆ
Ü�zÖ"ˆAÜÐ/°Ñ2Ó3ˆEØ�}Ó$Ü Ð!8¸Ñ!;Ó<�
Ø×!Ñ! '¨*Ñ"5Ô6Ø×!Ñ!¤#Ð&:¸1Ñ&=Ó">Ô?Ø×!Ñ! %Õ(ð #ñ Ø×+Ñ+¨J¸
ÀJÐPQÐSYÓZˆIð Ðð +-ˆIàÐr   r†   ÚtokensÚscoresc           
      ó®  — t        j                  |D �cg c]  }|g‘Œ c}|¬«      }t        |«      }| j                  j	                  |t        j                  dgt        |«      z  |¬«      |«      \  }	}
}g }t        |«      D ]K  \  }}t        |«      ||   gz   }|j                  ||	|   j                  «       t        |||«      ||   f«       ŒM |S c c}w )Nr1   r   )r'   r2   r-   rT   r_   r%   rP   r   r&   r`   r7   )r\   r†   rŒ   r�   r€   r/   ra   Ú
tgt_tokensr)   rc   rd   Úpred_statesrŠ   r*   rz   r‡   s                   r   r„   zRNNTBeamSearch._gen_new_hyposÍ   sÜ   € ô —\‘\¹Ó"?¹¨u E¢7¸Ñ"?ÈÔOˆ
Ü˜jÓ)ˆØ#'§:¡:×#5Ñ#5ØÜ�L‰L˜!˜œs :›Ñ.°vÔ>Øó$
Ñ ˆ�!�[ð
 ')ˆ	Ü 
Ö+‰FˆAˆsÜ)¨#Ó.°&¸±)°Ñ<ˆJØ×Ñ˜j¨(°1©+×*<Ñ*<Ó*>ÄÈ[ÐZ[Ð]cÓ@dÐflÐmnÑfoÐpÕqð ,ð Ðùò #@s   ”
Cr   c           	      ó˜  — |j                   d   }|j                  }g }|€| j                  |«      n|}t        |«      D ]ÿ  }|}t        j
                  j                  t        t           g «      }i }	d}
|rs| j                  |d d …||dz   …f   ||«      }|j                  «       }| j                  ||||	«      }|
| j                  k(  rn | j                  ||||||«      }|r|
dz  }
|rŒst	        j                  |D �cg c]  }| j                  |«      ‘Œ c}«      j!                  |«      \  }}|D �cg c]  }||   ‘Œ	 }}�Œ |S c c}w c c}w )Nr   r   )rD   r/   rg   r$   r'   ÚjitÚannotater   r	   rr   Úcpur   rX   r‹   r2   rW   rB   )r\   rh   r   r;   Ún_time_stepsr/   rt   rs   r€   ru   Úsymbols_current_tr:   Úhyprd   r~   r.   s                   r   Ú_searchzRNNTBeamSearch._searchâ   sm  € ð —}‘} QÑ'ˆØ—‘ˆà$&ˆØ04°�$×$Ñ$ VÔ,À$ˆÜ�|Ö$ˆAØˆGÜ—i‘i×(Ñ(¬¬jÑ)9¸2Ó>ˆGØ35ˆMØ !ÐáØ#'×#=Ñ#=¸gÂaÈÈQÐQRÉUÈÀlÑ>SÐU\Ð^dÓ#eÐ Ø#3×#7Ñ#7Ó#9Ð Ø×+Ñ+¨G°WÐ>NÐP]Ó^�à$¨×(<Ñ(<Ò<Øà×+Ñ+ØØØ$ØØØó�ñ Ø%¨Ñ*Ð%ò# ô& "ŸL™LÉWÓ)UÉWÀc¨$×*<Ñ*<¸SÕ*AÈWÑ)UÓV×[Ñ[Ð\fÓg‰MˆAˆzÙ/9Ó:©z¨�w˜s“|¨zˆGÒ:ð5 %ð8 ˆùò *VùÚ:s   Ã:E
Ä/EÚinputÚlengthc                 óÎ  — |j                  «       dk7  r0|j                  «       dk(  r|j                  d   dk(  st        d«      ‚|j                  «       dk(  r|j                  d«      }|j                  dk7  r|j                  dk7  rt        d«      ‚|j                  «       dk(  r|j                  d«      }| j                  j                  ||«      \  }}| j                  |d	|«      S )
a  Performs beam search for the given input sequence.

        T: number of frames;
        D: feature dimension of each frame.

        Args:
            input (torch.Tensor): sequence of input frames, with shape (T, D) or (1, T, D).
            length (torch.Tensor): number of valid frames in input
                sequence, with shape () or (1,).
            beam_width (int): beam size to use during search.

        Returns:
            List[Hypothesis]: top-``beam_width`` hypotheses found by beam search.
        r   r   r   r   ú*input must be of shape (T, D) or (1, T, D)r   ©r   ú"length must be of shape () or (1,)N)rj   rD   Ú
ValueErrorr@   rT   Ú
transcriber˜   )r\   r™   rš   r;   rh   rd   s         r   ÚforwardzRNNTBeamSearch.forward  sÄ   € ð �9‰9‹;˜!Ò U§Y¡Y£[°AÒ%5¸%¿+¹+Àa¹.ÈAÒ:MÜÐIÓJÐJØ�9‰9‹;˜!ÒØ—O‘O AÓ&ˆEà�<‰<˜2Ò &§,¡,°$Ò"6ÜÐAÓBÐBØ�:‰:‹<˜1ÒØ×%Ñ% aÓ(ˆFà—Z‘Z×*Ñ*¨5°&Ó9‰
ˆ�Ø�|‰|˜G T¨:Ó6Ð6r   r6   Ú
hypothesisc                 óÖ  — |j                  «       dk7  r0|j                  «       dk(  r|j                  d   dk(  st        d«      ‚|j                  «       dk(  r|j                  d«      }|j                  dk7  r|j                  dk7  rt        d«      ‚|j                  «       dk(  r|j                  d«      }| j                  j                  |||«      \  }}}| j                  |||«      |fS )	a™  Performs beam search for the given input sequence in streaming mode.

        T: number of frames;
        D: feature dimension of each frame.

        Args:
            input (torch.Tensor): sequence of input frames, with shape (T, D) or (1, T, D).
            length (torch.Tensor): number of valid frames in input
                sequence, with shape () or (1,).
            beam_width (int): beam size to use during search.
            state (List[List[torch.Tensor]] or None, optional): list of lists of tensors
                representing transcription network internal state generated in preceding
                invocation. (Default: ``None``)
            hypothesis (List[Hypothesis] or None): hypotheses from preceding invocation to seed
                search with. (Default: ``None``)

        Returns:
            (List[Hypothesis], List[List[torch.Tensor]]):
                List[Hypothesis]
                    top-``beam_width`` hypotheses found by beam search.
                List[List[torch.Tensor]]
                    list of lists of tensors representing transcription network
                    internal state generated in current invocation.
        r   r   r   r   rœ   r   r�   rž   )rj   rD   rŸ   r@   rT   Útranscribe_streamingr˜   )r\   r™   rš   r;   r6   r¢   rh   rd   s           r   ÚinferzRNNTBeamSearch.infer'  sÏ   € ðB �9‰9‹;˜!Ò U§Y¡Y£[°AÒ%5¸%¿+¹+Àa¹.ÈAÒ:MÜÐIÓJÐJØ�9‰9‹;˜!ÒØ—O‘O AÓ&ˆEà�<‰<˜2Ò &§,¡,°$Ò"6ÜÐAÓBÐBØ�:‰:‹<˜1ÒØ×%Ñ% aÓ(ˆFà ŸJ™J×;Ñ;¸EÀ6È5ÓQÑˆ��EØ�|‰|˜G Z°Ó<¸eÐCÐCr   )g      ð?Néd   )NN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   rƒ   rw   r   r   r	   r[   r'   r/   r   rg   ÚTensorrr   r   r    r   r‹   r„   r˜   r¡   r’   Úexportr   r¥   Ú__classcell__)r]   s   @r   r
   r
   K   sÕ  ø„ ñð( !ØAEØ"ñ/àð/ð ð/ð ð	/ð
   ¨*¨°uÐ)<Ñ =Ñ>ð/ð ð/ð 
õ/ð( E§L¡Lð °T¸*Ñ5Eó ð#Ø—|‘|ð#Ø,0°Ñ,<ð#ØFKÇlÁlð#à	�‰ó#ð4à�jÑ!ð4ð �jÑ!ð4ð  Ÿ,™,ð	4ð
 ˜C ˜OÑ,ð4ð 
ˆjÑ	ó4ð6$à�jÑ!ð$ð �jÑ!ð$ð  Ÿ,™,ð	$ð
 ð$ð ð$ð —‘ð$ð 
ˆjÑ	ó$ðLà˜Ñ$ðð �S‘	ðð �U‘ð	ð
 ðð —‘ðð 
ˆjÑ	óð*'à—‘ð'ð �t˜JÑ'Ñ(ð'ð ð	'ð
 
ˆjÑ	ó'ðR7˜UŸ\™\ð 7°5·<±<ð 7ÈSð 7ÐUYÐZdÑUeó 7ð8 ‡Y�Y×Ñð 59Ø15ñ+Dà�|‰|ð+Dð —‘ð+Dð ð	+Dð
 ˜˜T %§,¡,Ñ/Ñ0Ñ1ð+Dð ˜T *Ñ-Ñ.ð+Dð 
ˆt�JÑ  d¨5¯<©<Ñ&8Ñ!9Ð9Ñ	:ò+Dó ô+Dr   )Útypingr   r   r   r   r   r'   Útorchaudio.modelsr   Ú__all__rƒ   r«   rw   r	   rª   r   r   r   r   r    r!   r-   r/   r7   r9   rL   rR   rm   ÚModuler
   r   r   r   Ú<module>r²      sÕ  ðß 8Õ 8ã Ý "ð Ð)Ð
*€ð �4˜‘9˜eŸl™l¨D°°e·l±lÑ1CÑ,DÀeÐKÑL€
ð€
Ô ð
˜:ð ¨$¨s©)ó ð *ð °·±ó ð˜*ð ¨¨d°5·<±<Ñ.@Ñ)Aó ð˜*ð ¨ó ð˜
ð  só ð˜˜ZÑ(ð ¨T°$°u·|±|Ñ2DÑ-Eó ðd˜˜d 5§<¡<Ñ0Ñ1ð d¸ð dÀUÇ\Á\ð dÐVZÐ[_Ð`e×`lÑ`lÑ[mÑVnó dð
E ð E°ó Eð
PØ�
Ñð
Pà—l‘lð
Pð ð
Pð ˆ5�<‰<˜Ÿ™ u§|¡|Ð3Ñ4ó	
Pð�zð ¨d°:Ñ.>ð À4ó ôHD�U—X‘X—_‘_õ HDr   