
    ^j3                         d dl mZ d dlZddlmZmZ ddlmZ  G d d      Z G d d	e      Z	e G d
 de             Z
 G d de      Z G d de      Z G d de      Zy)    )	dataclassN   )GenerationMixinStoppingCriteria)ModelOutputc                   "    e Zd Zd Zd Z	 ddZy)ParakeetRNNTDecoderCachec                 J    || _         d | _        d | _        d | _        d| _        y )NF)configcachehidden_state
cell_stateis_initialized)selfr   s     {/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/transformers/models/parakeet/generation_parakeet.py__init__z!ParakeetRNNTDecoderCache.__init__   s'    *.
15/3$)    c                 P   t        j                  |j                  d   d| j                  j                  |j
                  |j                        | _        t        j                  | j                  j                  |j                  d   | j                  j                  |j
                  |j                        | _	        t        j                  | j                  j                  |j                  d   | j                  j                  |j
                  |j                        | _
        t         j                  j                  | j                         t         j                  j                  | j                         t         j                  j                  | j                         d| _        y )Nr      devicedtypeT)torchzerosshaper   decoder_hidden_sizer   r   r   num_decoder_layersr   r   _dynamomark_static_addressr   )r   hidden_statess     r   lazy_initializationz,ParakeetRNNTDecoderCache.lazy_initialization   s)   [["KK++ ''%%

 "KKKK**"KK++ ''%%
  ++KK**"KK++ ''%%
 	))$**5))$*;*;<))$//:"r   Nc                 j   | j                   s| j                  |       |R| j                  j                  |       | j                  j                  |       | j
                  j                  |       y |j                  |j                        }|j                  d   }|j                  d|d      }|j                  |dd      }t        j                  ||| j
                        | _        t        j                  ||| j                        | _        t        j                  ||| j                        | _        y )Nr   r   )r   r!   r   copy_r   r   tor   r   viewr   where)r   decoder_outputr   r   mask
batch_sizemask_hmask_ds           r   updatezParakeetRNNTDecoderCache.update<   s     ""$$^4<##L1OO!!*-JJ^, 77>001D'--a0JYYq*a0FYYz1a0FV^TZZHDJ %FL$BSBS TD#kk&*dooNDOr   N)__name__
__module____qualname__r   r!   r,    r   r   r	   r	      s    *#D Or   r	   c                       e Zd Zy)ParakeetTDTDecoderCacheN)r.   r/   r0   r1   r   r   r3   r3   V   s    r   r3   c                       e Zd ZU dZej
                  ed<   dZej
                  dz  ed<   dZe	e	ej                        dz  ed<   dZe	e	ej                        dz  ed<   y)ParakeetRNNTGenerateOutputa  
    Outputs of Parakeet transducer (RNN-T / TDT) generation.

    Args:
        sequences (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
            Generated token sequences (including blank tokens).
        durations (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Per-step durations in frames. Combined with `sequences`, this is sufficient
            to reconstruct full timestamp information (frame indices are the cumulative sum
            of durations).
        attentions (`tuple(tuple(torch.FloatTensor))`, *optional*):
            Encoder attention weights per layer.
        hidden_states (`tuple(tuple(torch.FloatTensor))`, *optional*):
            Encoder hidden states per layer.
    	sequencesN	durations
attentionsr    )r.   r/   r0   __doc__r   
LongTensor__annotations__r7   r8   tupleFloatTensorr    r1   r   r   r5   r5   Y   sh      )-Iu$&-9=JeE--./$6=<@M5u0012T9@r   r5   c                       e Zd ZdZd Zd Zy)EncoderExhaustedCriteriazVStops generation when all batch elements have walked past their encoder output length.c                     || _         y r-   )model)r   rA   s     r   r   z!EncoderExhaustedCriteria.__init__t   s	    
r   c                     | j                   j                  =t        j                  |j                  d   t        j
                  |j                        S | j                   j                  S )Nr   r   r   )rA   _encoder_finishedr   r   r   boolr   )r   	input_idsscoreskwargss       r   __call__z!EncoderExhaustedCriteria.__call__w   sH    ::''/;;yq1IL\L\]]zz+++r   N)r.   r/   r0   r9   r   rI   r1   r   r   r?   r?   q   s    `,r   r?   c                   \     e Zd ZdZ fdZ fdZ fdZ fdZd Z fdZ	d	 fd	Z
 xZS )
ParakeetRNNTGenerationMixina  Generation mixin for Parakeet RNN-T models, and the base for all Parakeet transducer generation.

    Handles the transducer machinery shared by RNN-T and TDT: encoder frame tracking, decoder cache
    preparation, encoder-exhaustion stopping, and output-buffer sizing. For RNN-T greedy decoding the encoder
    frame pointer advances by one frame on every blank emission and stays put on every non-blank emission; a
    ``max_symbols_per_step`` guard forces an advance after too many consecutive non-blank emissions at the same
    frame, mirroring NeMo's greedy RNN-T decoding. The duration-aware [`ParakeetTDTGenerationMixin`] extends this
    by advancing the frame pointer by a predicted duration instead.
    c                 Z    t        |   |i |}|j                  t        |              |S r-   )super_get_stopping_criteriaappendr?   )r   argsrH   criteria	__class__s       r   rN   z2ParakeetRNNTGenerationMixin._get_stopping_criteria   s.    714B6B067r   c                    t        |   |g|i |}|j                  d d dd d f   }|j                  d      }|| j                  j
                  k(  }| j                  t        j                  |      | _        t        j                  |t        j                  | j                        | j                  dz         }|| j                  k\  }	t        j                  ||	z  t        j                  |      |      | _        ||	z  j                         }
|d   |
z   |d<   | j                  j                  |
       |d   |d   k\  | _        |S )Ndimr   encoder_frame_idxsencoder_valid_lengths)rM   #_update_model_kwargs_for_generationlogitsargmaxr   blank_token_id_symbols_at_framer   
zeros_liker&   max_symbols_per_steplong_step_durationsrO   rD   )r   outputsrP   rH   model_kwargsrZ   tokens
blank_masksymbolsforce_advanceadvancerR   s              r   rY   z?ParakeetRNNTGenerationMixin._update_model_kwargs_for_generation   s>   wB7\T\U[\2q)2&t{{999
 !!)%*%5%5f%=D"++j%*:*:4;Q;Q*RTXTjTjmnTno4#<#<<!&Z--GIYIYZaIbdk!l -335-9:N-ORY-Y)* 	##G,!-.B!C|TkGl!lr   c                     |rx|j                   l| j                  j                  t        j                  |j
                  d   g|j                              j                         }| j                  |z  |_	        d}t        | -  ||||||      S )Nr   )r   F)max_new_tokensencoder_get_subsampling_output_lengthr   tensorr   r   itemr_   
max_lengthrM   _prepare_generated_length)	r   generation_confighas_default_max_lengthhas_default_min_lengthmodel_input_nameinput_ids_lengthinputs_tensorencoder_seq_lenrR   s	           r   rp   z5ParakeetRNNTGenerationMixin._prepare_generated_length   s     "&7&F&F&N"llIIm11!45m>R>RSdf  ,0+D+D+V(%*"w0""
 	
r   c                    t        |   |i |\  }}}h d}d}|j                         D 	ci c]  \  }}	||vr|j                  |      s||	 }
}}	 | j                  d||j                  dd       dd|
}||d<   |j                  |j                  j                  d      }nl|j                  j                  d   }t        j                  |f|j                  j                  d	   t        j                  |j                  j                  
      }||d<   t        j                  |j                  d   |j                  t        j                        |d<   |||fS c c}	}w )N>   attention_maskinput_featuresoutput_attention_mask)decoder_
cross_attn	use_cachepast_key_valuescache_paramsry   T)rz   ry   r{   encoder_outputsrT   r   r   rC   rX   r   rW   r1   )rM   _prepare_model_inputsitems
startswithget_audio_featuresgetry   sumlast_hidden_stater   r   fullr`   r   r   )r   rP   rH   inputs
input_namerc   explicitirrelevant_prefixkeyvalueencoder_kwargsr   rX   r)   rR   s                 r   r   z1ParakeetRNNTGenerationMixin._prepare_model_inputs   s|   +07+H$+YRX+Y(
LPf +002
U("3>>:K+L J
 
 2$11 
!'++,<dC"&
 	
 +:&'))5$3$B$B$F$Fr$J!(::@@CJ$)JJ1177:jj&88??	%! 1F,--2[[LLO==**.
)* z<//A
s   "Ec                 4    t        | j                        |d<   y )Ndecoder_cache)r	   r   )r   rq   rc   rP   rH   s        r   _prepare_cache_for_generationz9ParakeetRNNTGenerationMixin._prepare_cache_for_generation   s    (@(M_%r   c                 v   ddl m} t        
|   |g|i |}|j	                  d      j                  |d   j                  j                        }|d   j                  }|j                  d   |j                  d   }	}|j                  |	dz
        } ||t        j                  |      |d f         |d<   |S )Nr   )ParakeetEncoderModelOutputrW   r   r   )max)pooler_output)modeling_parakeetr   rM   prepare_inputs_for_generationpopr$   r   r   r   clampr   arange)r   rF   rP   rH   r   model_inputsrW   r   r)   max_encoder_lenrR   s             r   r   z9ParakeetRNNTGenerationMixin.prepare_inputs_for_generation   s    Aw<YXXQWX)--.BCFF*+99@@
 %%67EE&3&9&9!&<m>Q>QRS>TO
/55/A:M5N*D'Z(@BTVZ(Z[+
&' r   c                    d | _         d | _        g | _        t        |   d||d|}t        j                  | j                  d      }t        j                  t        j                  |j                  d   d|j                  |j                        |gd      }| ` | `| `t        t        |t              r|j                  |      S ||      S )N)r   rq   r   rU   r   rC   )r6   r7   r1   )rD   r]   ra   rM   generater   stackcatr   r   r   r   r5   
isinstancer   r6   )r   r   rq   rH   rb   r7   rR   s         r   r   z$ParakeetRNNTGenerationMixin.generate   s    !%!%!'"`&DU`Y_`KK 4 4!<	II[[+QiooiN^N^_ajkqr
	 "D$:D<P)+5g{+Kg''
 	
QX
 	
r   )NN)r.   r/   r0   r9   rN   rY   rp   r   r   r   r   __classcell__)rR   s   @r   rK   rK   }   s2    
0
6$0LN"
 
r   rK   c                       e Zd ZdZd Zy)ParakeetTDTGenerationMixina  Generation mixin for Parakeet TDT models.

    Extends [`ParakeetRNNTGenerationMixin`] with duration-aware decoding: instead of advancing the encoder frame
    pointer by one on each blank emission, the joint network predicts a per-step duration and the pointer advances
    by that amount. The shared setup (encoder frame tracking, decoder cache, stopping criteria, output buffer
    sizing) is inherited unchanged.
    c                     t        j                  | |g|i |}|j                  d d dd d f   }|d d d | j                  j                  f   j                  d      }|d d | j                  j                  d f   j                  d      }|| j                  j                  k(  }t        j                  ||dk(  z  t        j                  |      |      }|d   |z   |d<   | j                  j                  |       |d   |d   k\  | _        |S )NrT   rU   r   rW   rX   )r   rY   rZ   r   
vocab_sizer[   r\   r   r&   	ones_likera   rO   rD   )	r   rb   rP   rH   rc   rZ   rd   r7   re   s	            r   rY   z>ParakeetTDTGenerationMixin._update_model_kwargs_for_generation  s    'JJ4QXj[_jcij 2q)3T[[33334;;;C1dkk44667>>2>F	 t{{999
KK
i1n =uy?Y[de	-9:N-OR[-[)*##I. "..B!C|TkGl!lr   N)r.   r/   r0   r9   rY   r1   r   r   r   r     s    r   r   )dataclassesr   r   
generationr   r   utilsr   r	   r3   r5   r?   rK   r   r1   r   r   <module>r      sr    "  ;  ;O ;O~ =6 < A A A.	,/ 	,O
/ O
d!< r   