Ë
    óÿæiZ³  ã                   óD  — d dl Z d dlmZmZmZmZ d dlZd dlmZmZ d dl	m
Z dgZd"dededed	ed
ej                  j                   f
d„Z	 	 	 	 	 	 d#dededededeeeeee   f      deded	ed
ej                  j$                  fd„Zded
efd„Z G d„ dej*                  «      Z G d„ dej*                  «      Z G d„ dej*                  «      Z G d„ dej*                  «      Z G d„ dej*                  «      Z G d„ d ej*                  «      Z G d!„ dej*                  «      Zy)$é    N)ÚListÚOptionalÚTupleÚUnion)ÚnnÚTensor)Ú
functionalÚ	Tacotron2Úin_dimÚout_dimÚbiasÚw_init_gainÚreturnc                 ó  — t         j                  j                  | ||¬«      }t         j                  j                  j	                  |j
                  t         j                  j                  j                  |«      ¬«       |S )a  Linear layer with xavier uniform initialization.

    Args:
        in_dim (int): Size of each input sample.
        out_dim (int): Size of each output sample.
        bias (bool, optional): If set to ``False``, the layer will not learn an additive bias. (Default: ``True``)
        w_init_gain (str, optional): Parameter passed to ``torch.nn.init.calculate_gain``
            for setting the gain parameter of ``xavier_uniform_``. (Default: ``linear``)

    Returns:
        (torch.nn.Linear): The corresponding linear layer.
    ©r   ©Úgain)Útorchr   ÚLinearÚinitÚxavier_uniform_ÚweightÚcalculate_gain)r   r   r   r   Úlinears        úp/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchaudio/models/tacotron2.pyÚ_get_linear_layerr   )   sT   € ô �X‰X�_‰_˜V W°4ˆ_Ó8€FÜ	‡H�H‡M�M×!Ñ! &§-¡-´e·h±h·m±m×6RÑ6RÐS^Ó6_Ð!Ô`Ø€Mó    Úin_channelsÚout_channelsÚkernel_sizeÚstrideÚpaddingÚdilationc           	      ó\  — |€'|dz  dk7  rt        d«      ‚t        ||dz
  z  dz  «      }t        j                  j	                  | ||||||¬«      }t        j                  j
                  j                  |j                  t        j                  j
                  j                  |«      ¬«       |S )al  1D convolution with xavier uniform initialization.

    Args:
        in_channels (int): Number of channels in the input image.
        out_channels (int): Number of channels produced by the convolution.
        kernel_size (int, optional): Number of channels in the input image. (Default: ``1``)
        stride (int, optional): Number of channels in the input image. (Default: ``1``)
        padding (str, int or tuple, optional): Padding added to both sides of the input.
            (Default: dilation * (kernel_size - 1) / 2)
        dilation (int, optional): Number of channels in the input image. (Default: ``1``)
        w_init_gain (str, optional): Parameter passed to ``torch.nn.init.calculate_gain``
            for setting the gain parameter of ``xavier_uniform_``. (Default: ``linear``)

    Returns:
        (torch.nn.Conv1d): The corresponding Conv1D layer.
    é   é   zkernel_size must be odd)r    r!   r"   r#   r   r   )	Ú
ValueErrorÚintr   r   ÚConv1dr   r   r   r   )	r   r   r    r!   r"   r#   r   r   Úconv1ds	            r   Ú_get_conv1d_layerr+   ;   sŸ   € ð4 €Ø˜‰?˜aÒÜÐ6Ó7Ð7Ü�h +°¡/Ñ2°QÑ6Ó7ˆä�X‰X�_‰_ØØØØØØØð ó €Fô 
‡H�H‡M�M×!Ñ! &§-¡-´e·h±h·m±m×6RÑ6RÐS^Ó6_Ð!Ô`à€Mr   Úlengthsc                 ó  — t        j                  | «      j                  «       }t        j                  d|| j                  | j
                  ¬«      }|| j                  d«      k  j                  «       }t        j                  |d«      }|S )al  Returns a binary mask based on ``lengths``. The ``i``-th row and ``j``-th column of the mask
    is ``1`` if ``j`` is smaller than ``i``-th element of ``lengths.

    Args:
        lengths (Tensor): The length of each element in the batch, with shape (n_batch, ).

    Returns:
        mask (Tensor): The binary mask, with shape (n_batch, max of ``lengths``).
    r   )ÚdeviceÚdtyper&   )	r   ÚmaxÚitemÚaranger.   r/   Ú	unsqueezeÚbyteÚle)r,   Úmax_lenÚidsÚmasks       r   Ú_get_mask_from_lengthsr9   i   sj   € ô �i‰i˜Ó ×%Ñ%Ó'€GÜ
�,‰,�q˜'¨'¯.©.ÀÇÁÔ
N€CØ�'×#Ñ# AÓ&Ñ&×,Ñ,Ó.€DÜ�8‰8�D˜!Ó€DØ€Kr   c                   ó@   ‡ — e Zd ZdZdededefˆ fd„Zdedefd„Zˆ xZS )	Ú_LocationLayera  Location layer used in the Attention model.

    Args:
        attention_n_filter (int): Number of filters for attention model.
        attention_kernel_size (int): Kernel size for attention model.
        attention_hidden_dim (int): Dimension of attention hidden representation.
    Úattention_n_filterÚattention_kernel_sizeÚattention_hidden_dimc           	      óš   •— t         ‰| �  «        t        |dz
  dz  «      }t        d|||ddd¬«      | _        t        ||dd¬«      | _        y )Nr&   r%   F)r    r"   r   r!   r#   Útanh©r   r   )ÚsuperÚ__init__r(   r+   Úlocation_convr   Úlocation_dense)Úselfr<   r=   r>   r"   Ú	__class__s        €r   rC   z_LocationLayer.__init__ƒ   s`   ø€ ô 	‰ÑÔÜÐ,¨qÑ0°AÑ5Ó6ˆÜ.ØØØ-ØØØØô
ˆÔô 0ØÐ 4¸5Èfô
ˆÕr   Úattention_weights_catr   c                 ón   — | j                  |«      }|j                  dd«      }| j                  |«      }|S )a�  Location layer used in the Attention model.

        Args:
            attention_weights_cat (Tensor): Cumulative and previous attention weights
                with shape (n_batch, 2, max of ``text_lengths``).

        Returns:
            processed_attention (Tensor): Cumulative and previous attention weights
                with shape (n_batch, ``attention_hidden_dim``).
        r&   r%   )rD   Ú	transposerE   )rF   rH   Úprocessed_attentions      r   Úforwardz_LocationLayer.forward˜   sA   € ð #×0Ñ0Ð1FÓGÐØ1×;Ñ;¸A¸qÓAÐà"×1Ñ1Ð2EÓFÐØ"Ð"r   ©	Ú__name__Ú
__module__Ú__qualname__Ú__doc__r(   rC   r   rL   Ú__classcell__©rG   s   @r   r;   r;   z   s<   ø„ ñð
àð
ð  #ð
ð "õ	
ð*#¨Vð #¸÷ #r   r;   c                   ó~   ‡ — e Zd ZdZdedededededdfˆ fd	„Zd
edededefd„Zdedededededeeef   fd„Z	ˆ xZ
S )Ú
_Attentionaº  Locally sensitive attention model.

    Args:
        attention_rnn_dim (int): Number of hidden units for RNN.
        encoder_embedding_dim (int): Number of embedding dimensions in the Encoder.
        attention_hidden_dim (int): Dimension of attention hidden representation.
        attention_location_n_filter (int): Number of filters for Attention model.
        attention_location_kernel_size (int): Kernel size for Attention model.
    Úattention_rnn_dimÚencoder_embedding_dimr>   Úattention_location_n_filterÚattention_location_kernel_sizer   Nc                 óÞ   •— t         ‰| �  «        t        ||dd¬«      | _        t        ||dd¬«      | _        t        |dd¬«      | _        t        |||«      | _        t        d«       | _	        y )NFr@   rA   r&   r   Úinf)
rB   rC   r   Úquery_layerÚmemory_layerÚvr;   Úlocation_layerÚfloatÚscore_mask_value)rF   rV   rW   r>   rX   rY   rG   s         €r   rC   z_Attention.__init__¶   sx   ø€ ô 	‰ÑÔÜ,Ð->Ð@TÐ[`ÐntÔuˆÔÜ-Ø!Ð#7¸eÐQWô
ˆÔô #Ð#7¸ÀÔGˆŒÜ,Ø'Ø*Ø ó
ˆÔô
 "' u£ ˆÕr   ÚqueryÚprocessed_memoryrH   c                 óÞ   — | j                  |j                  d«      «      }| j                  |«      }| j                  t	        j
                  ||z   |z   «      «      }|j                  d«      }|S )a=  Get the alignment vector.

        Args:
            query (Tensor): Decoder output with shape (n_batch, n_mels * n_frames_per_step).
            processed_memory (Tensor): Processed Encoder outputs
                with shape (n_batch, max of ``text_lengths``, attention_hidden_dim).
            attention_weights_cat (Tensor): Cumulative and previous attention weights
                with shape (n_batch, 2, max of ``text_lengths``).

        Returns:
            alignment (Tensor): attention weights, it is a tensor with shape (batch, max of ``text_lengths``).
        r&   r%   )r\   r3   r_   r^   r   r@   Úsqueeze)rF   rb   rc   rH   Úprocessed_queryÚprocessed_attention_weightsÚenergiesÚ	alignments           r   Ú_get_alignment_energiesz"_Attention._get_alignment_energiesË   sh   € ð ×*Ñ*¨5¯?©?¸1Ó+=Ó>ˆØ&*×&9Ñ&9Ð:OÓ&PÐ#Ø—6‘6œ%Ÿ*™* _Ð7RÑ%RÐUeÑ%eÓfÓgˆà×$Ñ$ QÓ'ˆ	ØÐr   Úattention_hidden_stateÚmemoryr8   c                 ó  — | j                  |||«      }|j                  || j                  «      }t        j                  |d¬«      }t        j                  |j                  d«      |«      }|j                  d«      }||fS )a¹  Pass the input through the Attention model.

        Args:
            attention_hidden_state (Tensor): Attention rnn last output with shape (n_batch, ``attention_rnn_dim``).
            memory (Tensor): Encoder outputs with shape (n_batch, max of ``text_lengths``, ``encoder_embedding_dim``).
            processed_memory (Tensor): Processed Encoder outputs
                with shape (n_batch, max of ``text_lengths``, ``attention_hidden_dim``).
            attention_weights_cat (Tensor): Previous and cumulative attention weights
                with shape (n_batch, current_num_frames * 2, max of ``text_lengths``).
            mask (Tensor): Binary mask for padded data with shape (n_batch, current_num_frames).

        Returns:
            attention_context (Tensor): Context vector with shape (n_batch, ``encoder_embedding_dim``).
            attention_weights (Tensor): Attention weights with shape (n_batch, max of ``text_lengths``).
        r&   ©Údim)	rj   Úmasked_fillra   ÚFÚsoftmaxr   Úbmmr3   re   )	rF   rk   rl   rc   rH   r8   ri   Úattention_weightsÚattention_contexts	            r   rL   z_Attention.forwardà   s~   € ð. ×0Ñ0Ð1GÐIYÐ[pÓqˆ	à×)Ñ)¨$°×0EÑ0EÓFˆ	äŸI™I i°QÔ7ÐÜ!ŸI™IÐ&7×&AÑ&AÀ!Ó&DÀfÓMÐØ-×5Ñ5°aÓ8Ðà Ð"3Ð3Ð3r   )rN   rO   rP   rQ   r(   rC   r   rj   r   rL   rR   rS   s   @r   rU   rU   «   s²   ø„ ñð.àð.ð  #ð.ð "ð	.ð
 &)ð.ð ),ð.ð 
õ.ð*¨Vð Àvð Ðflð Ðqwó ð*4à &ð4ð ð4ð !ð	4ð
  &ð4ð ð4ð 
ˆv�vˆ~Ñ	÷4r   rU   c                   óF   ‡ — e Zd ZdZdedee   ddfˆ fd„Zdedefd„Zˆ xZ	S )	Ú_PrenetzÒPrenet Module. It is consists of ``len(output_size)`` linear layers.

    Args:
        in_dim (int): The size of each input sample.
        output_sizes (list): The output dimension of each linear layers.
    r   Ú	out_sizesr   Nc                 óÆ   •— t         ‰| �  «        |g|d d z   }t        j                  t	        ||«      D ��cg c]  \  }}t        ||d¬«      ‘Œ c}}«      | _        y c c}}w )NéÿÿÿÿFr   )rB   rC   r   Ú
ModuleListÚzipr   Úlayers)rF   r   rx   Úin_sizesÚin_sizeÚout_sizerG   s         €r   rC   z_Prenet.__init__
  s_   ø€ Ü‰ÑÔØ�8˜i¨¨˜nÑ,ˆÜ—m‘mÜY\Ð]eÐgpÔYqÔrÑYqÑBUÀ7ÈHÔ˜w¨°uÖ=ÐYqÒró
ˆ�ùÛrs   ·A
Úxc                 óŠ   — | j                   D ]3  }t        j                  t        j                   ||«      «      dd¬«      }Œ5 |S )zÙPass the input through Prenet.

        Args:
            x (Tensor): The input sequence to Prenet with shape (n_batch, in_dim).

        Return:
            x (Tensor): Tensor with shape (n_batch, sizes[-1])
        ç      à?T)ÚpÚtraining)r}   rq   ÚdropoutÚrelu)rF   r�   r   s      r   rL   z_Prenet.forward  s6   € ð —k”kˆFÜ—	‘	œ!Ÿ&™&¡¨£Ó+¨s¸TÔB‰Að "àˆr   )
rN   rO   rP   rQ   r(   r   rC   r   rL   rR   rS   s   @r   rw   rw     s9   ø„ ñð
˜sð 
¨t°C©yð 
¸Tõ 
ð˜ð  F÷ r   rw   c                   óD   ‡ — e Zd ZdZdedededefˆ fd„Zdedefd	„Zˆ xZS )
Ú_Postneta  Postnet Module.

    Args:
        n_mels (int): Number of mel bins.
        postnet_embedding_dim (int): Postnet embedding dimension.
        postnet_kernel_size (int): Postnet kernel size.
        postnet_n_convolution (int): Number of postnet convolutions.
    Ún_melsÚpostnet_embedding_dimÚpostnet_kernel_sizeÚpostnet_n_convolutionc                 óÄ  •— t         ‰
| �  «        t        j                  «       | _        t        |«      D ]�  }|dk(  r|n|}||dz
  k(  r|n|}||dz
  k(  rdnd}||dz
  k(  r|n|}	| j                  j                  t        j                  t        |||dt        |dz
  dz  «      d|¬«      t        j                  |	«      «      «       Œ’ t        | j                  «      | _        y )Nr   r&   r   r@   r%   ©r    r!   r"   r#   r   )rB   rC   r   r{   ÚconvolutionsÚrangeÚappendÚ
Sequentialr+   r(   ÚBatchNorm1dÚlenÚn_convs)rF   rŠ   r‹   rŒ   r�   Úir   r   Ú	init_gainÚnum_featuresrG   s             €r   rC   z_Postnet.__init__*  sé   ø€ ô 	‰ÑÔÜŸM™M›OˆÔäÐ,Ö-ˆAØ$%¨¢F™&Ð0EˆKØ%&Ð+@À1Ñ+DÒ%E™6ÐK`ˆLØ$%Ð*?À!Ñ*CÒ$D™È&ˆIØ%&Ð+@À1Ñ+DÒ%E™6ÐK`ˆLØ×Ñ×$Ñ$Ü—‘Ü%Ø#Ø$Ø$7Ø Ü #Ð%8¸1Ñ%<ÀÑ$AÓ BØ!"Ø$-ôô —N‘N <Ó0óõð .ô( ˜4×,Ñ,Ó-ˆ�r   r�   r   c                 ó,  — t        | j                  «      D ]{  \  }}|| j                  dz
  k  r<t        j                  t        j                   ||«      «      d| j                  ¬«      }ŒTt        j                   ||«      d| j                  ¬«      }Œ} |S )a  Pass the input through Postnet.

        Args:
            x (Tensor): The input sequence with shape (n_batch, ``n_mels``, max of ``mel_specgram_lengths``).

        Return:
            x (Tensor): Tensor with shape (n_batch, ``n_mels``, max of ``mel_specgram_lengths``).
        r&   rƒ   )r…   )Ú	enumerater�   r–   rq   r†   r   r@   r…   )rF   r�   r—   Úconvs       r   rL   z_Postnet.forwardJ  sr   € ô ! ×!2Ñ!2Ö3‰GˆAˆtØ�4—<‘< !Ñ#Ò#Ü—I‘IœeŸj™j©¨a«Ó1°3ÀÇÁÔO‘ä—I‘I™d 1›g s°T·]±]ÔC‘ð	 4ð ˆr   rM   rS   s   @r   r‰   r‰      sG   ø„ ñð.àð.ð  #ð.ð !ð	.ð
  #õ.ð@˜ð  F÷ r   r‰   c                   óH   ‡ — e Zd ZdZdedededdfˆ fd„Zded	edefd
„Zˆ xZS )Ú_Encodera§  Encoder Module.

    Args:
        encoder_embedding_dim (int): Number of embedding dimensions in the encoder.
        encoder_n_convolution (int): Number of convolution layers in the encoder.
        encoder_kernel_size (int): The kernel size in the encoder.

    Examples
        >>> encoder = _Encoder(3, 512, 5)
        >>> input = torch.rand(10, 20, 30)
        >>> output = encoder(input)  # shape: (10, 30, 512)
    rW   Úencoder_n_convolutionÚencoder_kernel_sizer   Nc                 óÄ  •— t         ‰| �  «        t        j                  «       | _        t        |«      D ]e  }t        j                  t        |||dt        |dz
  dz  «      dd¬«      t        j                  |«      «      }| j                  j                  |«       Œg t        j                  |t        |dz  «      ddd¬«      | _        | j                  j                  «        y )Nr&   r%   r‡   r�   T)Úbatch_firstÚbidirectional)rB   rC   r   r{   r�   r‘   r“   r+   r(   r”   r’   ÚLSTMÚlstmÚflatten_parameters)rF   rW   rŸ   r    Ú_Ú
conv_layerrG   s         €r   rC   z_Encoder.__init__k  sÌ   ø€ ô 	‰ÑÔäŸM™M›OˆÔÜÐ,Ö-ˆAÜŸ™Ü!Ø)Ø)Ø 3ØÜÐ!4°qÑ!8¸AÑ =Ó>ØØ &ôô —‘Ð4Ó5óˆJð ×Ñ×$Ñ$ ZÕ0ð .ô —G‘GØ!ÜÐ%¨Ñ)Ó*ØØØô
ˆŒ	ð 	�	‰	×$Ñ$Õ&r   r�   Úinput_lengthsc                 ó¼  — | j                   D ]<  }t        j                  t        j                   ||«      «      d| j                  «      }Œ> |j                  dd«      }|j                  «       }t        j                  j                  j                  ||d¬«      }| j                  |«      \  }}t        j                  j                  j                  |d¬«      \  }}|S )a_  Pass the input through the Encoder.

        Args:
            x (Tensor): The input sequences with shape (n_batch, encoder_embedding_dim, n_seq).
            input_lengths (Tensor): The length of each input sequence with shape (n_batch, ).

        Return:
            x (Tensor): A tensor with shape (n_batch, n_seq, encoder_embedding_dim).
        rƒ   r&   r%   T)r¢   )r�   rq   r†   r‡   r…   rJ   Úcpur   ÚutilsÚrnnÚpack_padded_sequencer¥   Úpad_packed_sequence)rF   r�   r©   rœ   Úoutputsr§   s         r   rL   z_Encoder.forwardŒ  s¬   € ð ×%Ô%ˆDÜ—	‘	œ!Ÿ&™&¡ a£›/¨3°·±Ó>‰Að &ð �K‰K˜˜1Óˆà%×)Ñ)Ó+ˆÜ�H‰H�L‰L×-Ñ-¨a°ÈDÐ-ÓQˆà—Y‘Y˜q“\‰
ˆ�Ü—X‘X—\‘\×5Ñ5°gÈ4Ð5ÓP‰
ˆ�àˆr   rM   rS   s   @r   rž   rž   ]  sN   ø„ ñð'à"ð'ð  #ð'ð !ð	'ð
 
õ'ðB˜ð °ð ¸6÷ r   rž   c            !       ó¢  ‡ — e Zd ZdZdededededededed	ed
ededededededdfˆ fd„Zdedefd„Z	dede
eeeeeeeef   fd„Zdedefd„Zdededede
eeef   fd„Zdedededed ed!ed"ed#eded$ed%ede
eeeeeeeeef	   fd&„Zded'ed(ede
eeef   fd)„Zdedefd*„Zej$                  j&                  ded(ede
eeeef   fd+„«       Zˆ xZS ),Ú_Decodera,  Decoder with Attention model.

    Args:
        n_mels (int): number of mel bins
        n_frames_per_step (int): number of frames processed per step, only 1 is supported
        encoder_embedding_dim (int): the number of embedding dimensions in the encoder.
        decoder_rnn_dim (int): number of units in decoder LSTM
        decoder_max_step (int): maximum number of output mel spectrograms
        decoder_dropout (float): dropout probability for decoder LSTM
        decoder_early_stopping (bool): stop decoding when all samples are finished
        attention_rnn_dim (int): number of units in attention LSTM
        attention_hidden_dim (int): dimension of attention hidden representation
        attention_location_n_filter (int): number of filters for attention model
        attention_location_kernel_size (int): kernel size for attention model
        attention_dropout (float): dropout probability for attention LSTM
        prenet_dim (int): number of ReLU units in prenet layers
        gate_threshold (float): probability threshold for stop token
    rŠ   Ún_frames_per_steprW   Údecoder_rnn_dimÚdecoder_max_stepÚdecoder_dropoutÚdecoder_early_stoppingrV   r>   rX   rY   Úattention_dropoutÚ
prenet_dimÚgate_thresholdr   Nc                 óæ  •— t         ‰| �  «        || _        || _        || _        || _        || _        || _        || _        || _	        || _
        || _        || _        t        ||z  ||g«      | _        t        j                   ||z   |«      | _        t%        |||	|
|«      | _        t        j                   ||z   |d«      | _        t+        ||z   ||z  «      | _        t+        ||z   ddd¬«      | _        y )NTr&   ÚsigmoidrA   )rB   rC   rŠ   r³   rW   rV   r´   r¹   rµ   rº   r¸   r¶   r·   rw   Úprenetr   ÚLSTMCellÚattention_rnnrU   Úattention_layerÚdecoder_rnnr   Úlinear_projectionÚ
gate_layer)rF   rŠ   r³   rW   r´   rµ   r¶   r·   rV   r>   rX   rY   r¸   r¹   rº   rG   s                  €r   rC   z_Decoder.__init__¹  s  ø€ ô$ 	‰ÑÔØˆŒØ!2ˆÔØ%:ˆÔ"Ø!2ˆÔØ.ˆÔØ$ˆŒØ 0ˆÔØ,ˆÔØ!2ˆÔØ.ˆÔØ&<ˆÔ#ä˜fÐ'8Ñ8¸:ÀzÐ:RÓSˆŒäŸ[™[¨Ð6KÑ)KÐM^Ó_ˆÔä)ØØ!Ø Ø'Ø*ó 
ˆÔô Ÿ;™;Ð'8Ð;PÑ'PÐRaÐcgÓhˆÔä!2°?ÐEZÑ3ZÐ\bÐevÑ\vÓ!wˆÔä+ØÐ3Ñ3°Q¸TÈyô
ˆ�r   rl   c                 ó¸   — |j                  d«      }|j                  }|j                  }t        j                  || j
                  | j                  z  ||¬«      }|S )am  Gets all zeros frames to use as the first decoder input.

        Args:
            memory (Tensor): Encoder outputs with shape (n_batch, max of ``text_lengths``, ``encoder_embedding_dim``).

        Returns:
            decoder_input (Tensor): all zeros frames with shape
                (n_batch, max of ``text_lengths``, ``n_mels * n_frames_per_step``).
        r   ©r/   r.   ©Úsizer/   r.   r   ÚzerosrŠ   r³   ©rF   rl   Ún_batchr/   r.   Údecoder_inputs         r   Ú_get_initial_framez_Decoder._get_initial_frameì  óN   € ð —+‘+˜a“.ˆØ—‘ˆØ—‘ˆÜŸ™ G¨T¯[©[¸4×;QÑ;QÑ-QÐY^ÐgmÔnˆØÐr   c                 ó‚  — |j                  d«      }|j                  d«      }|j                  }|j                  }t        j                  || j
                  ||¬«      }t        j                  || j
                  ||¬«      }t        j                  || j                  ||¬«      }t        j                  || j                  ||¬«      }	t        j                  ||||¬«      }
t        j                  ||||¬«      }t        j                  || j                  ||¬«      }| j                  j                  |«      }||||	|
|||fS )a  Initializes attention rnn states, decoder rnn states, attention
        weights, attention cumulative weights, attention context, stores memory
        and stores processed memory.

        Args:
            memory (Tensor): Encoder outputs with shape (n_batch, max of ``text_lengths``, ``encoder_embedding_dim``).

        Returns:
            attention_hidden (Tensor): Hidden state of the attention LSTM with shape (n_batch, ``attention_rnn_dim``).
            attention_cell (Tensor): Hidden state of the attention LSTM with shape (n_batch, ``attention_rnn_dim``).
            decoder_hidden (Tensor): Hidden state of the decoder LSTM with shape (n_batch, ``decoder_rnn_dim``).
            decoder_cell (Tensor): Hidden state of the decoder LSTM with shape (n_batch, ``decoder_rnn_dim``).
            attention_weights (Tensor): Attention weights with shape (n_batch, max of ``text_lengths``).
            attention_weights_cum (Tensor): Cumulated attention weights with shape (n_batch, max of ``text_lengths``).
            attention_context (Tensor): Context vector with shape (n_batch, ``encoder_embedding_dim``).
            processed_memory (Tensor): Processed encoder outputs
                with shape (n_batch, max of ``text_lengths``, ``attention_hidden_dim``).
        r   r&   rÅ   )
rÇ   r/   r.   r   rÈ   rV   r´   rW   rÀ   r]   )rF   rl   rÊ   Úmax_timer/   r.   Úattention_hiddenÚattention_cellÚdecoder_hiddenÚdecoder_cellrt   Úattention_weights_cumru   rc   s                 r   Ú_initialize_decoder_statesz#_Decoder._initialize_decoder_statesý  s   € ð* —+‘+˜a“.ˆØ—;‘;˜q“>ˆØ—‘ˆØ—‘ˆä Ÿ;™; w°×0FÑ0FÈeÐ\bÔcÐÜŸ™ W¨d×.DÑ.DÈEÐZ`ÔaˆäŸ™ W¨d×.BÑ.BÈ%ÐX^Ô_ˆÜ—{‘{ 7¨D×,@Ñ,@ÈÐV\Ô]ˆä!ŸK™K¨°ÀÈvÔVÐÜ %§¡¨G°XÀUÐSYÔ ZÐÜ!ŸK™K¨°×1KÑ1KÐSXÐagÔhÐà×/Ñ/×<Ñ<¸VÓDÐð ØØØØØ!ØØð	
ð 		
r   Údecoder_inputsc                 óÜ   — |j                  dd«      }|j                  |j                  d«      t        |j                  d«      | j                  z  «      d«      }|j                  dd«      }|S )ak  Prepares decoder inputs.

        Args:
            decoder_inputs (Tensor): Inputs used for teacher-forced training, i.e. mel-specs,
                with shape (n_batch, ``n_mels``, max of ``mel_specgram_lengths``)

        Returns:
            inputs (Tensor): Processed decoder inputs with shape (max of ``mel_specgram_lengths``, n_batch, ``n_mels``).
        r&   r%   r   rz   )rJ   ÚviewrÇ   r(   r³   )rF   rÖ   s     r   Ú_parse_decoder_inputsz_Decoder._parse_decoder_inputs.  so   € ð (×1Ñ1°!°QÓ7ˆØ'×,Ñ,Ø×Ñ Ó"Ü�×#Ñ# AÓ&¨×)?Ñ)?Ñ?Ó@Øó
ˆð (×1Ñ1°!°QÓ7ˆØÐr   Úmel_specgramÚgate_outputsÚ
alignmentsc                 óF  — |j                  dd«      j                  «       }|j                  dd«      j                  «       }|j                  dd«      j                  «       }|j                  d   d| j                  f} |j                  |Ž }|j                  dd«      }|||fS )aq  Prepares decoder outputs for output

        Args:
            mel_specgram (Tensor): mel spectrogram with shape (max of ``mel_specgram_lengths``, n_batch, ``n_mels``)
            gate_outputs (Tensor): predicted stop token with shape (max of ``mel_specgram_lengths``, n_batch)
            alignments (Tensor): sequence of attention weights from the decoder
                with shape (max of ``mel_specgram_lengths``, n_batch, max of ``text_lengths``)

        Returns:
            mel_specgram (Tensor): mel spectrogram with shape (n_batch, ``n_mels``, max of ``mel_specgram_lengths``)
            gate_outputs (Tensor): predicted stop token with shape (n_batch, max of ``mel_specgram_lengths``)
            alignments (Tensor): sequence of attention weights from the decoder
                with shape (n_batch, max of ``mel_specgram_lengths``, max of ``text_lengths``)
        r   r&   rz   r%   )rJ   Ú
contiguousÚshaperŠ   rØ   )rF   rÚ   rÛ   rÜ   rß   s        r   Ú_parse_decoder_outputsz_Decoder._parse_decoder_outputsC  s¡   € ð&  ×)Ñ)¨!¨QÓ/×:Ñ:Ó<ˆ
à#×-Ñ-¨a°Ó3×>Ñ>Ó@ˆà#×-Ñ-¨a°Ó3×>Ñ>Ó@ˆà×#Ñ# AÑ&¨¨D¯K©KÐ8ˆØ(�|×(Ñ(¨%Ð0ˆà#×-Ñ-¨a°Ó3ˆà˜\¨:Ð5Ð5r   rË   rÐ   rÑ   rÒ   rÓ   rt   rÔ   ru   rc   r8   c           	      óž  — t        j                  ||fd«      }| j                  |||f«      \  }}t        j                  || j
                  | j                  «      }t        j                  |j                  d«      |j                  d«      fd¬«      }| j                  ||	|
||«      \  }}||z  }t        j                  ||fd«      }| j                  |||f«      \  }}t        j                  || j                  | j                  «      }t        j                  ||fd¬«      }| j                  |«      }| j                  |«      }|||||||||f	S )a&	  Decoder step using stored states, attention and memory

        Args:
            decoder_input (Tensor): Output of the Prenet with shape (n_batch, ``prenet_dim``).
            attention_hidden (Tensor): Hidden state of the attention LSTM with shape (n_batch, ``attention_rnn_dim``).
            attention_cell (Tensor): Hidden state of the attention LSTM with shape (n_batch, ``attention_rnn_dim``).
            decoder_hidden (Tensor): Hidden state of the decoder LSTM with shape (n_batch, ``decoder_rnn_dim``).
            decoder_cell (Tensor): Hidden state of the decoder LSTM with shape (n_batch, ``decoder_rnn_dim``).
            attention_weights (Tensor): Attention weights with shape (n_batch, max of ``text_lengths``).
            attention_weights_cum (Tensor): Cumulated attention weights with shape (n_batch, max of ``text_lengths``).
            attention_context (Tensor): Context vector with shape (n_batch, ``encoder_embedding_dim``).
            memory (Tensor): Encoder output with shape (n_batch, max of ``text_lengths``, ``encoder_embedding_dim``).
            processed_memory (Tensor): Processed Encoder outputs
                with shape (n_batch, max of ``text_lengths``, ``attention_hidden_dim``).
            mask (Tensor): Binary mask for padded data with shape (n_batch, current_num_frames).

        Returns:
            decoder_output: Predicted mel spectrogram for the current frame with shape (n_batch, ``n_mels``).
            gate_prediction (Tensor): Prediction of the stop token with shape (n_batch, ``1``).
            attention_hidden (Tensor): Hidden state of the attention LSTM with shape (n_batch, ``attention_rnn_dim``).
            attention_cell (Tensor): Hidden state of the attention LSTM with shape (n_batch, ``attention_rnn_dim``).
            decoder_hidden (Tensor): Hidden state of the decoder LSTM with shape (n_batch, ``decoder_rnn_dim``).
            decoder_cell (Tensor): Hidden state of the decoder LSTM with shape (n_batch, ``decoder_rnn_dim``).
            attention_weights (Tensor): Attention weights with shape (n_batch, max of ``text_lengths``).
            attention_weights_cum (Tensor): Cumulated attention weights with shape (n_batch, max of ``text_lengths``).
            attention_context (Tensor): Context vector with shape (n_batch, ``encoder_embedding_dim``).
        rz   r&   rn   )r   Úcatr¿   rq   r†   r¸   r…   r3   rÀ   rÁ   r¶   rÂ   rÃ   )rF   rË   rÐ   rÑ   rÒ   rÓ   rt   rÔ   ru   rl   rc   r8   Ú
cell_inputrH   Ú decoder_hidden_attention_contextÚdecoder_outputÚgate_predictions                    r   Údecodez_Decoder.decodec  si  € ôR —Y‘Y Ð/@ÐAÀ2ÓFˆ
à+/×+=Ñ+=¸jÐK[Ð]kÐJlÓ+mÑ(Ð˜.ÜŸ9™9Ð%5°t×7MÑ7MÈtÏ}É}Ó]Ðä %§	¡	Ð+<×+FÑ+FÀqÓ+IÐK`×KjÑKjÐklÓKmÐ*nÐtuÔ vÐØ/3×/CÑ/CØ˜fÐ&6Ð8MÈtó0
Ñ,ÐÐ,ð 	Ð!2Ñ2ÐÜŸ	™	Ð#3Ð5FÐ"GÈÓLˆà'+×'7Ñ'7¸ÈÐXdÐGeÓ'fÑ$ˆ˜ÜŸ™ >°4×3GÑ3GÈÏÉÓWˆä+0¯9©9°nÐFWÐ5XÐ^_Ô+`Ð(Ø×/Ñ/Ð0PÓQˆàŸ/™/Ð*JÓKˆð ØØØØØØØ!Øð

ð 
	
r   Úmel_specgram_truthÚmemory_lengthsc                 ó   — | j                  |«      j                  d«      }| j                  |«      }t        j                  ||fd¬«      }| j                  |«      }t        |«      }| j                  |«      \  }}}	}
}}}}g g g }}}t        |«      |j                  d«      dz
  k  r„|t        |«         }| j                  ||||	|
||||||«      \	  }}}}}	}
}}}||j                  d«      gz  }||j                  d«      gz  }||gz  }t        |«      |j                  d«      dz
  k  rŒ„| j                  t        j                  |«      t        j                  |«      t        j                  |«      «      \  }}}|||fS )aî  Decoder forward pass for training.

        Args:
            memory (Tensor): Encoder outputs
                with shape (n_batch, max of ``text_lengths``, ``encoder_embedding_dim``).
            mel_specgram_truth (Tensor): Decoder ground-truth mel-specs for teacher forcing
                with shape (n_batch, ``n_mels``, max of ``mel_specgram_lengths``).
            memory_lengths (Tensor): Encoder output lengths for attention masking
                (the same as ``text_lengths``) with shape (n_batch, ).

        Returns:
            mel_specgram (Tensor): Predicted mel spectrogram
                with shape (n_batch, ``n_mels``, max of ``mel_specgram_lengths``).
            gate_outputs (Tensor): Predicted stop token for each timestep
                with shape (n_batch,  max of ``mel_specgram_lengths``).
            alignments (Tensor): Sequence of attention weights from the decoder
                with shape (n_batch,  max of ``mel_specgram_lengths``, max of ``text_lengths``).
        r   rn   r&   )rÌ   r3   rÙ   r   râ   r½   r9   rÕ   r•   rÇ   rç   re   rà   Ústack)rF   rl   rè   ré   rË   rÖ   r8   rÐ   rÑ   rÒ   rÓ   rt   rÔ   ru   rc   Úmel_outputsrÛ   rÜ   Ú
mel_outputÚgate_outputrÚ   s                        r   rL   z_Decoder.forward­  s½  € ð, ×/Ñ/°Ó7×AÑAÀ!ÓDˆØ×3Ñ3Ð4FÓGˆÜŸ™ M°>Ð#BÈÔJˆØŸ™ ^Ó4ˆä% nÓ5ˆð ×+Ñ+¨FÓ3ñ		
ØØØØØØ!ØØð 13°B¸ :�\ˆÜ�+Ó ×!4Ñ!4°QÓ!7¸!Ñ!;Ò;Ø*¬3¨{Ó+;Ñ<ˆMð —‘ØØ ØØØØ!Ø%Ø!ØØ Øóñ
ØØØ ØØØØ!Ø%Ø!ð ˜J×.Ñ.¨qÓ1Ð2Ñ2ˆKØ˜[×0Ñ0°Ó3Ð4Ñ4ˆLØÐ,Ð-Ñ-ˆJô9 �+Ó ×!4Ñ!4°QÓ!7¸!Ñ!;Ó;ð< 26×1LÑ1LÜ�K‰K˜Ó$¤e§k¡k°,Ó&?ÄÇÁÈZÓAXó2
Ñ.ˆ�l Jð ˜\¨:Ð5Ð5r   c                 ó¸   — |j                  d«      }|j                  }|j                  }t        j                  || j
                  | j                  z  ||¬«      }|S )aU  Gets all zeros frames to use as the first decoder input

        args:
            memory (Tensor): Encoder outputs
                with shape (n_batch, max of ``text_lengths``, ``encoder_embedding_dim``).

        returns:
            decoder_input (Tensor): All zeros frames with shape(n_batch, ``n_mels`` * ``n_frame_per_step``).
        r   rÅ   rÆ   rÉ   s         r   Ú_get_go_framez_Decoder._get_go_frameù  rÍ   r   c                 ó€  — |j                  d«      |j                  }}| j                  |«      }t        |«      }| j	                  |«      \  }}}	}
}}}}t        j                  |gt
        j                  |¬«      }t        j                  |gt
        j                  |¬«      }g }g }g }t        | j                  «      D ]ñ  }| j                  |«      }| j                  ||||	|
||||||«      \	  }}}}}	}
}}}|j                  |j                  d«      «       |j                  |j                  dd«      «       |j                  |«       || xx   dz  cc<   |t        j                   |j#                  d«      «      | j$                  kD  z  }| j&                  rt        j(                  |«      r n|}Œó t+        |«      | j                  k(  rt-        j.                  d«       t        j0                  |d¬«      }t        j0                  |d¬«      }t        j0                  |d¬«      }| j3                  |||«      \  }}}||||fS )a’  Decoder inference

        Args:
            memory (Tensor): Encoder outputs
                with shape (n_batch, max of ``text_lengths``, ``encoder_embedding_dim``).
            memory_lengths (Tensor): Encoder output lengths for attention masking
                (the same as ``text_lengths``) with shape (n_batch, ).

        Returns:
            mel_specgram (Tensor): Predicted mel spectrogram
                with shape (n_batch, ``n_mels``, max of ``mel_specgram_lengths``).
            mel_specgram_lengths (Tensor): the length of the predicted mel spectrogram (n_batch, ))
            gate_outputs (Tensor): Predicted stop token for each timestep
                with shape (n_batch,  max of ``mel_specgram_lengths``).
            alignments (Tensor): Sequence of attention weights from the decoder
                with shape (n_batch,  max of ``mel_specgram_lengths``, max of ``text_lengths``).
        r   rÅ   r&   zZReached max decoder steps. The generated spectrogram might not cover the whole transcript.rn   )rÇ   r.   rð   r9   rÕ   r   rÈ   Úint32Úboolr‘   rµ   r½   rç   r’   r3   rJ   r¼   re   rº   r·   Úallr•   ÚwarningsÚwarnrâ   rà   )rF   rl   ré   Ú
batch_sizer.   rË   r8   rÐ   rÑ   rÒ   rÓ   rt   rÔ   ru   rc   Úmel_specgram_lengthsÚfinishedÚmel_specgramsrÛ   rÜ   r§   rÚ   rî   s                          r   Úinferz_Decoder.infer
  s6  € ð& $Ÿ[™[¨›^¨V¯]©]�Fˆ
à×*Ñ*¨6Ó2ˆä% nÓ5ˆð ×+Ñ+¨FÓ3ñ		
ØØØØØØ!ØØô  %Ÿ{™{¨J¨<¼u¿{¹{ÐSYÔZÐÜ—;‘; 
˜|´5·:±:ÀfÔMˆØ&(ˆØ%'ˆØ#%ˆ
Ü�t×,Ñ,Ö-ˆAØ ŸK™K¨Ó6ˆMð —‘ØØ ØØØØ!Ø%Ø!ØØ Øóñ
ØØØ ØØØØ!Ø%Ø!ð × Ñ  ×!7Ñ!7¸Ó!:Ô;Ø×Ñ × 5Ñ 5°a¸Ó ;Ô<Ø×ÑÐ/Ô0Ø  ( Ó+¨qÑ0Ó+àœŸ™ k×&9Ñ&9¸!Ó&<Ó=À×@SÑ@SÑSÑSˆHØ×*Ò*¬u¯y©y¸Ô/BÙà(‰MðG .ôJ ˆ}Ó ×!6Ñ!6Ò6Ü�M‰MØoôô Ÿ	™	 -°QÔ7ˆÜ—y‘y °1Ô5ˆÜ—Y‘Y˜z¨qÔ1ˆ
à26×2MÑ2MÈmÐ]iÐkuÓ2vÑ/ˆ�| ZàÐ2°LÀ*ÐLÐLr   )rN   rO   rP   rQ   r(   r`   ró   rC   r   rÌ   r   rÕ   rÙ   rà   rç   rL   rð   r   ÚjitÚexportrû   rR   rS   s   @r   r²   r²   ¥  s^  ø„ ñð&1
àð1
ð ð1
ð  #ð	1
ð
 ð1
ð ð1
ð ð1
ð !%ð1
ð ð1
ð "ð1
ð &)ð1
ð ),ð1
ð !ð1
ð ð1
ð ð1
ð  
õ!1
ðf¨ð °Fó ð"/
Øð/
à	ˆv�v˜v v¨v°v¸vÀvÐMÑ	Nó/
ðb°Fð ¸vó ð*6Ø"ð6Ø28ð6ØFLð6à	ˆv�v˜vÐ%Ñ	&ó6ð@H
àðH
ð !ðH
ð ð	H
ð
 ðH
ð ðH
ð "ðH
ð  &ðH
ð "ðH
ð ðH
ð !ðH
ð ðH
ð 
ˆv�v˜v v¨v°v¸vÀvÈvÐUÑ	VóH
ðTJ6ØðJ6Ø28ðJ6ØJPðJ6à	ˆv�v˜vÐ%Ñ	&óJ6ðX Fð ¨vó ð" ‡Y�Y×ÑðWM˜Fð WM°Fð WM¸uÀVÈVÐU[Ð]cÐEcÑ?dò WMó ôWMr   r²   c            /       ó2  ‡ — e Zd ZdZ	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 	 d"dededededededed	ed
ededededededededededededededdf.ˆ fd„Zdedededede	eeeef   f
d„Z
ej                  j                  d#ded ee   de	eeef   fd!„«       Zˆ xZS )$r
   aÑ	  Tacotron2 model from *Natural TTS Synthesis by Conditioning WaveNet on Mel Spectrogram Predictions*
    :cite:`shen2018natural` based on the implementation from
    `Nvidia Deep Learning Examples <https://github.com/NVIDIA/DeepLearningExamples/>`_.

    See Also:
        * :class:`torchaudio.pipelines.Tacotron2TTSBundle`: TTS pipeline with pretrained model.

    Args:
        mask_padding (bool, optional): Use mask padding (Default: ``False``).
        n_mels (int, optional): Number of mel bins (Default: ``80``).
        n_symbol (int, optional): Number of symbols for the input text (Default: ``148``).
        n_frames_per_step (int, optional): Number of frames processed per step, only 1 is supported (Default: ``1``).
        symbol_embedding_dim (int, optional): Input embedding dimension (Default: ``512``).
        encoder_n_convolution (int, optional): Number of encoder convolutions (Default: ``3``).
        encoder_kernel_size (int, optional): Encoder kernel size (Default: ``5``).
        encoder_embedding_dim (int, optional): Encoder embedding dimension (Default: ``512``).
        decoder_rnn_dim (int, optional): Number of units in decoder LSTM (Default: ``1024``).
        decoder_max_step (int, optional): Maximum number of output mel spectrograms (Default: ``2000``).
        decoder_dropout (float, optional): Dropout probability for decoder LSTM (Default: ``0.1``).
        decoder_early_stopping (bool, optional): Continue decoding after all samples are finished (Default: ``True``).
        attention_rnn_dim (int, optional): Number of units in attention LSTM (Default: ``1024``).
        attention_hidden_dim (int, optional): Dimension of attention hidden representation (Default: ``128``).
        attention_location_n_filter (int, optional): Number of filters for attention model (Default: ``32``).
        attention_location_kernel_size (int, optional): Kernel size for attention model (Default: ``31``).
        attention_dropout (float, optional): Dropout probability for attention LSTM (Default: ``0.1``).
        prenet_dim (int, optional): Number of ReLU units in prenet layers (Default: ``256``).
        postnet_n_convolution (int, optional): Number of postnet convolutions (Default: ``5``).
        postnet_kernel_size (int, optional): Postnet kernel size (Default: ``5``).
        postnet_embedding_dim (int, optional): Postnet embedding dimension (Default: ``512``).
        gate_threshold (float, optional): Probability threshold for stop token (Default: ``0.5``).
    Úmask_paddingrŠ   Ún_symbolr³   Úsymbol_embedding_dimrW   rŸ   r    r´   rµ   r¶   r·   rV   r>   rX   rY   r¸   r¹   r�   rŒ   r‹   rº   r   Nc                 ó€  •— t         ‰| �  «        || _        || _        || _        t        j                  ||«      | _        t        j
                  j                  j                  | j                  j                  «       t        |||«      | _        t        ||||	|
|||||||||«      | _        t!        ||||«      | _        y ©N)rB   rC   rÿ   rŠ   r³   r   Ú	EmbeddingÚ	embeddingr   r   r   r   rž   Úencoderr²   Údecoderr‰   Úpostnet)rF   rÿ   rŠ   r   r³   r  rW   rŸ   r    r´   rµ   r¶   r·   rV   r>   rX   rY   r¸   r¹   r�   rŒ   r‹   rº   rG   s                          €r   rC   zTacotron2.__init__†  sº   ø€ ô2 	‰ÑÔà(ˆÔØˆŒØ!2ˆÔÜŸ™ hÐ0DÓEˆŒÜ�‰�‰×%Ñ% d§n¡n×&;Ñ&;Ô<ÜÐ 5Ð7LÐNaÓbˆŒÜØØØ!ØØØØ"ØØ Ø'Ø*ØØØó
ˆŒô    Ð(=Ð?RÐTiÓjˆ�r   ÚtokensÚtoken_lengthsrÚ   rø   c                 ó  — | j                  |«      j                  dd«      }| j                  ||«      }| j                  |||¬«      \  }}}| j	                  |«      }	||	z   }	| j
                  r™t        |«      }
|
j                  | j                  |
j                  d«      |
j                  d«      «      }
|
j                  ddd«      }
|j                  |
d«       |	j                  |
d«       |j                  |
dd…ddd…f   d«       ||	||fS )a´  Pass the input through the Tacotron2 model. This is in teacher
        forcing mode, which is generally used for training.

        The input ``tokens`` should be padded with zeros to length max of ``token_lengths``.
        The input ``mel_specgram`` should be padded with zeros to length max of ``mel_specgram_lengths``.

        Args:
            tokens (Tensor): The input tokens to Tacotron2 with shape `(n_batch, max of token_lengths)`.
            token_lengths (Tensor): The valid length of each sample in ``tokens`` with shape `(n_batch, )`.
            mel_specgram (Tensor): The target mel spectrogram
                with shape `(n_batch, n_mels, max of mel_specgram_lengths)`.
            mel_specgram_lengths (Tensor): The length of each mel spectrogram with shape `(n_batch, )`.

        Returns:
            [Tensor, Tensor, Tensor, Tensor]:
                Tensor
                    Mel spectrogram before Postnet with shape `(n_batch, n_mels, max of mel_specgram_lengths)`.
                Tensor
                    Mel spectrogram after Postnet with shape `(n_batch, n_mels, max of mel_specgram_lengths)`.
                Tensor
                    The output for stop token at each time step with shape `(n_batch, max of mel_specgram_lengths)`.
                Tensor
                    Sequence of attention weights from the decoder with
                    shape `(n_batch, max of mel_specgram_lengths, max of token_lengths)`.
        r&   r%   )ré   r   g        Ng     @�@)r  rJ   r  r  r  rÿ   r9   ÚexpandrŠ   rÇ   ÚpermuteÚmasked_fill_)rF   r	  r
  rÚ   rø   Úembedded_inputsÚencoder_outputsrÛ   rÜ   Úmel_specgram_postnetr8   s              r   rL   zTacotron2.forward¹  s  € ðB Ÿ.™.¨Ó0×:Ñ:¸1¸aÓ@ˆàŸ,™, ¸ÓFˆØ15·±Ø˜\¸-ð 2>ó 2
Ñ.ˆ�l Jð  $Ÿ|™|¨LÓ9ÐØ+Ð.BÑBÐà×ÒÜ)Ð*>Ó?ˆDØ—;‘;˜tŸ{™{¨D¯I©I°a«L¸$¿)¹)ÀA»,ÓGˆDØ—<‘<  1 aÓ(ˆDà×%Ñ% d¨CÔ0Ø ×-Ñ-¨d°CÔ8Ø×%Ñ% dª1¨a²¨7¡m°SÔ9àÐ1°<ÀÐKÐKr   r,   c                 óâ  — |j                   \  }}|€It        j                  |g«      j                  |«      j	                  |j
                  |j                  «      }|€J ‚| j                  |«      j                  dd«      }| j                  ||«      }| j                  j                  ||«      \  }}}	}
| j                  |«      }||z   }|
j                  d||«      j                  dd«      }
|||
fS )aª  Using Tacotron2 for inference. The input is a batch of encoded
        sentences (``tokens``) and its corresponding lengths (``lengths``). The
        output is the generated mel spectrograms, its corresponding lengths, and
        the attention weights from the decoder.

        The input `tokens` should be padded with zeros to length max of ``lengths``.

        Args:
            tokens (Tensor): The input tokens to Tacotron2 with shape `(n_batch, max of lengths)`.
            lengths (Tensor or None, optional):
                The valid length of each sample in ``tokens`` with shape `(n_batch, )`.
                If ``None``, it is assumed that the all the tokens are valid. Default: ``None``

        Returns:
            (Tensor, Tensor, Tensor):
                Tensor
                    The predicted mel spectrogram with shape `(n_batch, n_mels, max of mel_specgram_lengths)`.
                Tensor
                    The length of the predicted mel spectrogram with shape `(n_batch, )`.
                Tensor
                    Sequence of attention weights from the decoder with shape
                    `(n_batch, max of mel_specgram_lengths, max of lengths)`.
        r&   r%   r   )rß   r   Útensorr  Útor.   r/   r  rJ   r  r  rû   r  Úunfold)rF   r	  r,   rÊ   Ú
max_lengthr  r  rÚ   rø   r§   rÜ   Úmel_outputs_postnets               r   rû   zTacotron2.inferï  sì   € ð2 %Ÿl™lÑˆ�Øˆ?Ü—l‘l J <Ó0×7Ñ7¸Ó@×CÑCÀFÇMÁMÐSY×S_ÑS_Ó`ˆGàÐ"Ð"Ð"ØŸ.™.¨Ó0×:Ñ:¸1¸aÓ@ˆØŸ,™, ¸Ó@ˆØ<@¿L¹L×<NÑ<NÈÐ`gÓ<hÑ9ˆÐ*¨A¨zà"Ÿl™l¨<Ó8ÐØ*Ð-@Ñ@Ðà×&Ñ& q¨'°7Ó;×EÑEÀaÈÓKˆ
à"Ð$8¸*ÐDÐDr   )FéP   é”   r&   é   r  é   é   é   iÐ  çš™™™™™¹?Tr  é€   é    é   r  é   r  r  r  rƒ   r  )rN   rO   rP   rQ   ró   r(   r`   rC   r   r   rL   r   rü   rý   r   rû   rR   rS   s   @r   r
   r
   e  sä  ø„ ñðD #ØØØ!"Ø$'Ø%(Ø%&Ø#$Ø#Ø $Ø!$Ø'+Ø!%Ø$'Ø+-Ø.0Ø#&ØØ%&Ø#$Ø%(Ø #ñ/1kàð1kð ð1kð ð	1kð
 ð1kð "ð1kð  #ð1kð  #ð1kð !ð1kð ð1kð ð1kð ð1kð !%ð1kð ð1kð "ð1kð  &)ð!1kð" ),ð#1kð$ !ð%1kð& ð'1kð(  #ð)1kð* !ð+1kð,  #ð-1kð. ð/1kð0 
õ11kðf4Làð4Lð ð4Lð ð	4Lð
 %ð4Lð 
ˆv�v˜v vÐ-Ñ	.ó4Lðl ‡Y�Y×Ññ&E˜Fð &E¨X°fÑ-=ð &EÈÈvÐW]Ð_eÐOeÑIfò &Eó ô&Er   )Tr   )r&   r&   Nr&   Tr   )rõ   Útypingr   r   r   r   r   r   r   Útorch.nnr	   rq   Ú__all__r(   ró   Ústrr   r   r)   r+   r9   ÚModuler;   rU   rw   r‰   rž   r²   r
   © r   r   Ú<module>r)     s|  ðó8 ß /Ó /ã ß Ý $ð ð€ñ
˜cð ¨Cð °tð ÐQTð Ðdi×dlÑdl×dsÑdsó ð* ØØ59ØØØñ+Øð+àð+ð ð+ð ð	+ð
 �e˜C  e¨C¡jÐ0Ñ1Ñ2ð+ð ð+ð ð+ð ð+ð ‡X�X‡_�_ó+ð\ Fð ¨vó ô".#�R—Y‘Yô .#ôbT4�—‘ô T4ônˆb�i‰iô ô<:ˆr�y‰yô :ôzEˆr�y‰yô EôP}Mˆr�y‰yô }Mô@qE�—	‘	õ qEr   