§
    ‚Štjš¿  ã                   óô  — d dl Z d dlmZ d dlmZ d dlmZ d dlZd dl	Z	d dl
mZ d dlmc m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mZmZ ddlmZ ddlmZmZ ddl m!Z!m"Z" ddl#m$Z$ ddl%m&Z&m'Z'm(Z(m)Z) ddl*m+Z+ ddl,m-Z- ddl.m/Z/ e(e G d„ de&¦  «        ¦   «         ¦   «         Z0e(e G d„ de&¦  «        ¦   «         ¦   «         Z1e(e G d„ de&¦  «        ¦   «         ¦   «         Z2 G d„ dej3        ¦  «        Z4 G d„ dej3        ¦  «        Z5d„ Z6 ed ¦  «        dXd!„¦   «         Z7d"e	j8        d#e9d$e	j8        fd%„Z:	 dYd'ej3        d(e	j8        d)e	j8        d*e	j8        d+e	j8        dz  d,e;d-e;d.e$e'         fd/„Z< ee7¦  «         G d0„ d1ej3        ¦  «        ¦   «         Z= ed2¦  «         G d3„ d4ej3        ¦  «        ¦   «         Z> G d5„ d6e¦  «        Z? G d7„ d8ej3        ¦  «        Z@d9„ ZA G d:„ d;ej3        ¦  «        ZB G d<„ d=ej3        ¦  «        ZC G d>„ d?ej3        ¦  «        ZD G d@„ dAej3        ¦  «        ZE G dB„ dCej3        ¦  «        ZF G dD„ dEej3        ¦  «        ZG G dF„ dGej3        ¦  «        ZH G dH„ dIej3        ¦  «        ZI G dJ„ dKej3        ¦  «        ZJ G dL„ dMej3        ¦  «        ZK G dN„ dOej3        ¦  «        ZL G dP„ dQej3        ¦  «        ZMe( G dR„ dSe"¦  «        ¦   «         ZN e(dT¬U¦  «         G dV„ dWeN¦  «        ¦   «         ZOdWdSgZPdS )Zé    N)ÚCallable)Ú	dataclass)ÚOptional)Ú	Parameteré   )Úinitialization)ÚACT2FN)ÚCache)Úuse_kernel_forward_from_hubÚuse_kernel_func_from_hubÚuse_kernelized_func)ÚGradientCheckpointingLayer)ÚROPE_INIT_FUNCTIONSÚdynamic_rope_update)ÚALL_ATTENTION_FUNCTIONSÚPreTrainedModel)ÚUnpack)ÚModelOutputÚTransformersKwargsÚauto_docstringÚcan_return_tuple)Úmaybe_autocasté   )Ú	AutoModelé   )ÚXcodec2Configc                   óŒ   — e Zd ZU dZdZej        dz  ed<   dZej	        dz  ed<   dZ
ej        dz  ed<   dZej        dz  ed<   dS )ÚXcodec2OutputaL  
    audio_values (`torch.FloatTensor` of shape `(batch_size, 1, sequence_length)`, *optional*):
        Decoded audio waveform values in the time domain, obtained using the decoder
        part of Xcodec2. These represent the reconstructed audio signal.
    audio_codes (`torch.LongTensor` of shape `(batch_size, 1, codes_length)`, *optional*):
        Discrete code embeddings computed using `model.encode`. These are the quantized
        representations of the input audio used for further processing or generation.
    latents (`torch.Tensor` of shape `(batch_size, dimension, time_steps)`):
        Quantized continuous representation of input's embedding.
    audio_codes_mask (`torch.int32` of shape `(batch_size, 1, codes_length)`, *optional*):
        Downsampled `padding_mask` for indicating valid audio codes in `audio_codes`.
    NÚaudio_valuesÚaudio_codesÚlatentsÚaudio_codes_mask)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚtorchÚFloatTensorÚ__annotations__r    Ú
LongTensorr!   ÚTensorr"   © ó    új/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/xcodec2/modeling_xcodec2.pyr   r   -   s}   € € € € € € ðð ð .2€L�%Ô# dÑ*Ð1Ð1Ñ1Ø+/€K�Ô! DÑ(Ð/Ð/Ñ/Ø#'€GˆUŒ\˜DÑ Ð'Ð'Ñ'Ø,0Ð�e”l TÑ)Ð0Ð0Ñ0Ð0Ð0r-   r   c                   ón   — e Zd ZU dZdZej        dz  ed<   dZej	        dz  ed<   dZ
ej	        dz  ed<   dS )ÚXcodec2EncoderOutputat  
    audio_codes (`torch.LongTensor` of shape `(batch_size, 1, codes_length)`, *optional*):
        Discrete code embeddings computed using `model.encode`. These represent
        the compressed, quantized form of the input audio signal that can be
        used for storage, transmission, or generation.
    latents (`torch.Tensor` of shape `(batch_size, dimension, time_steps)`):
        Quantized continuous representation of input's embedding.
    audio_codes_mask (`torch.int32` of shape `(batch_size, 1, codes_length)`, *optional*):
        Downsampled `padding_mask` for indicating valid audio codes in `audio_codes`.
    Nr    r!   r"   )r#   r$   r%   r&   r    r'   r*   r)   r!   r+   r"   r,   r-   r.   r0   r0   C   se   € € € € € € ð	ð 	ð ,0€K�Ô! DÑ(Ð/Ð/Ñ/Ø#'€GˆUŒ\˜DÑ Ð'Ð'Ñ'Ø,0Ð�e”l TÑ)Ð0Ð0Ñ0Ð0Ð0r-   r0   c                   ó2   — e Zd ZU dZdZej        dz  ed<   dS )ÚXcodec2DecoderOutputa=  
    audio_values (`torch.FloatTensor` of shape `(batch_size, 1, segment_length)`, *optional*):
        Decoded audio waveform values in the time domain, obtained by converting
        the discrete codes back into continuous audio signals. This represents
        the reconstructed audio that can be played back.
    Nr   )r#   r$   r%   r&   r   r'   r(   r)   r,   r-   r.   r2   r2   V   s6   € € € € € € ðð ð .2€L�%Ô# dÑ*Ð1Ð1Ñ1Ð1Ð1r-   r2   c                   óÔ   ‡ — e Zd ZU ej        ed<   ddefˆ fd„Ze	 	 	 ddedz  de	d         de
dz  ded	ef         fd
„¦   «         Z ej        ¦   «         ed„ ¦   «         ¦   «         Zˆ xZS )ÚXcodec2RotaryEmbeddingÚinv_freqNÚconfigc                 ó²  •— t          ¦   «                              ¦   «          |j        | _        |j        | _        || _        | j        j        d         | _        | j        }| j        dk    rt          | j                 } || j        |¦  «        \  }| _
        |                      d|d¬¦  «         |                      d|                     ¦   «         d¬¦  «         d S )NÚ	rope_typeÚdefaultr5   F©Ú
persistentÚoriginal_inv_freq)ÚsuperÚ__init__Úmax_position_embeddingsÚmax_seq_len_cachedÚoriginal_max_seq_lenr6   Úrope_parametersr8   Úcompute_default_rope_parametersr   Úattention_scalingÚregister_bufferÚclone)Úselfr6   ÚdeviceÚrope_init_fnr5   Ú	__class__s        €r.   r>   zXcodec2RotaryEmbedding.__init__f   sÊ   ø€ Ý‰Œ×ÒÑÔÐØ"(Ô"@ˆÔØ$*Ô$BˆÔ!àˆŒàœÔ4°[ÔAˆŒØ!%Ô!EˆØŒ>˜YÒ&Ð&Ý.¨t¬~Ô>ˆLØ+7¨<¸¼ÀVÑ+LÔ+LÑ(ˆ�$Ô(à×Ò˜Z¨¸eÐÑDÔDÐDØ×ÒÐ0°(·.².Ñ2BÔ2BÈuÐÑUÔUÐUÐUÐUr-   rH   ztorch.deviceÚseq_lenÚreturnztorch.Tensorc                 óü   — | j         d         }t          | dd¦  «        p| j        | j        z  }d}d|t	          j        d|dt          j        ¬¦  «                             |t          j        ¬¦  «        |z  z  z  }||fS )	a¨  
        Computes the inverse frequencies according to the original RoPE implementation
        Args:
            config ([`~transformers.PreTrainedConfig`]):
                The model configuration.
            device (`torch.device`):
                The device to use for initialization of the inverse frequencies.
            seq_len (`int`, *optional*):
                The current sequence length. Unused for this type of RoPE.
        Returns:
            Tuple of (`torch.Tensor`, `float`), containing the inverse frequencies for the RoPE embeddings and the
            post-processing scaling factor applied to the computed cos/sin (unused in this type of RoPE).
        Ú
rope_thetaÚhead_dimNç      ð?r   r   ©Údtype)rH   rR   )	rB   ÚgetattrÚhidden_sizeÚnum_attention_headsr'   ÚarangeÚint64ÚtoÚfloat)r6   rH   rK   ÚbaseÚdimÚattention_factorr5   s          r.   rC   z6Xcodec2RotaryEmbedding.compute_default_rope_parametersv   sŒ   € ð& Ô% lÔ3ˆÝ�f˜j¨$Ñ/Ô/Ðc°6Ô3EÈÔIcÑ3cˆàÐð Ø•U”\ ! S¨!µ5´;Ð?Ñ?Ô?×BÒBÈ&ÕX]ÔXcÐBÑdÔdÐgjÑjÑkñ
ˆð Ð)Ð)Ð)r-   c                 óN  — | j         d d d …d f                              ¦   «                              |j        d         dd¦  «                             |j        ¦  «        }|d d …d d d …f                              ¦   «         }t          |j        j        t          ¦  «        r|j        j        dk    r|j        j        nd}t          |d¬¦  «        5  |                     ¦   «         |                     ¦   «         z   
                    dd¦  «        }t          j        ||fd¬	¦  «        }|                     ¦   «         | j        z  }|                     ¦   «         | j        z  }	d d d ¦  «         n# 1 swxY w Y   |                     |j        ¬
¦  «        |	                     |j        ¬
¦  «        fS )Nr   éÿÿÿÿr   ÚmpsÚcpuF©Údevice_typeÚenabledr   ©r[   rQ   )r5   rY   ÚexpandÚshaperX   rH   Ú
isinstanceÚtypeÚstrr   Ú	transposer'   ÚcatÚcosrD   ÚsinrR   )
rG   ÚxÚposition_idsÚinv_freq_expandedÚposition_ids_expandedrb   ÚfreqsÚembrl   rm   s
             r.   ÚforwardzXcodec2RotaryEmbedding.forward”   s·  € ð !œM¨$°°°°4¨-Ô8×>Ò>Ñ@Ô@×GÒGÈÔHZÐ[\ÔH]Ð_aÐcdÑeÔe×hÒhÐijÔiqÑrÔrÐØ ,¨Q¨Q¨Q°°a°a°a¨ZÔ 8× >Ò >Ñ @Ô @Ðå'1°!´(´-ÅÑ'EÔ'EÐkÈ!Ì(Ì-Ð[`ÒJ`ÐJ`�a”h”m�mÐfkˆÝ¨¸UÐCÑCÔCð 	5ð 	5Ø&×,Ò,Ñ.Ô.Ð1F×1LÒ1LÑ1NÔ1NÑN×YÒYÐZ[Ð]^Ñ_Ô_ˆEÝ”)˜U E˜N°Ð3Ñ3Ô3ˆCØ—'’'‘)”)˜dÔ4Ñ4ˆCØ—'’'‘)”)˜dÔ4Ñ4ˆCð		5ð 	5ð 	5ñ 	5ô 	5ð 	5ð 	5ð 	5ð 	5ð 	5ð 	5øøøð 	5ð 	5ð 	5ð 	5ð �vŠv˜AœGˆvÑ$Ô$ c§f¢f°1´7 fÑ&;Ô&;Ð;Ð;s   ÃBE&Å&E*Å-E*©N©NNN)r#   r$   r%   r'   r+   r)   r   r>   Ústaticmethodr   ÚintÚtuplerY   rC   Úno_gradr   rt   Ú__classcell__©rJ   s   @r.   r4   r4   c   sù   ø€ € € € € € ØŒlÐÐÑðVð V˜}ð Vð Vð Vð Vð Vð Vð  à'+Ø+/Ø"ð*ð *Ø Ñ$ð*à˜Ô(ð*ð �t‘ð*ð 
ˆ~˜uÐ$Ô	%ð	*ð *ð *ñ „\ð*ð: €U„]�_„_Øð<ð <ñ Ôñ „_ð<ð <ð <ð <ð <r-   r4   c                   óH   ‡ — e Zd Zdefˆ fd„Zdej        dej        fd„Zˆ xZS )Ú
Xcodec2MLPr6   c                 ó  •— t          ¦   «                              ¦   «          || _        t          |j                 | _        t          j        |j        |j	        d¬¦  «        | _
        t          j        |j	        |j        d¬¦  «        | _        d S )NF©Úbias)r=   r>   r6   r	   Ú
hidden_actÚactivation_fnÚnnÚLinearrT   Úintermediate_sizeÚfc1Úfc2©rG   r6   rJ   s     €r.   r>   zXcodec2MLP.__init__¥   sr   ø€ Ý‰Œ×ÒÑÔÐØˆŒÝ# FÔ$5Ô6ˆÔÝ”9˜VÔ/°Ô1IÐPUÐVÑVÔVˆŒÝ”9˜VÔ5°vÔ7IÐPUÐVÑVÔVˆŒˆˆr-   Úhidden_statesrL   c                 ó„   — |                       |¦  «        }|                      |¦  «        }|                      |¦  «        }|S ru   )r‡   rƒ   rˆ   ©rG   rŠ   s     r.   rt   zXcodec2MLP.forward¬   s=   € ØŸš Ñ/Ô/ˆØ×*Ò*¨=Ñ9Ô9ˆØŸš Ñ/Ô/ˆØÐr-   ©	r#   r$   r%   r   r>   r'   r+   rt   r{   r|   s   @r.   r~   r~   ¤   sq   ø€ € € € € ðW˜}ð Wð Wð Wð Wð Wð Wð U¤\ð °e´lð ð ð ð ð ð ð ð r-   r~   c                 óœ   — | dd| j         d         dz  …f         }| d| j         d         dz  d…f         }t          j        | |fd¬¦  «        S )z*Rotates half the hidden dims of the input..Nr^   r   rd   )rf   r'   rk   )rn   Úx1Úx2s      r.   Úrotate_halfr‘   ³   s]   € à	
ˆ3Ð"�!”'˜"”+ Ñ"Ð"Ð"Ô	#€BØ	
ˆ3�”˜”˜qÑ Ð"Ð"Ð"Ô	#€BÝŒ9�r�c˜2�Y BÐ'Ñ'Ô'Ð'r-   Úrotary_pos_embc                 ó¾   — |                      |¦  «        }|                      |¦  «        }| |z  t          | ¦  «        |z  z   }||z  t          |¦  «        |z  z   }||fS )a…  Applies Rotary Position Embedding to the query and key tensors.

    Args:
        q (`torch.Tensor`): The query tensor.
        k (`torch.Tensor`): The key tensor.
        cos (`torch.Tensor`): The cosine part of the rotary embedding.
        sin (`torch.Tensor`): The sine part of the rotary embedding.
        unsqueeze_dim (`int`, *optional*, defaults to 1):
            The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and
            sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note
            that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and
            k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes
            cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have
            the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2.
    Returns:
        `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding.
    )Ú	unsqueezer‘   )ÚqÚkrl   rm   Úunsqueeze_dimÚq_embedÚk_embeds          r.   Úapply_rotary_pos_embrš   º   sc   € ð& �-Š-˜Ñ
&Ô
&€CØ
�-Š-˜Ñ
&Ô
&€CØ�3‰w�; q™>œ>¨CÑ/Ñ0€GØ�3‰w�; q™>œ>¨CÑ/Ñ0€GØ�GÐÐr-   rŠ   Ún_reprL   c                 ó¸   — | j         \  }}}}|dk    r| S | dd…dd…ddd…dd…f                              |||||¦  «        } |                      |||z  ||¦  «        S )zÔ
    This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
    num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
    r   N)rf   re   Úreshape)rŠ   r›   ÚbatchÚnum_key_value_headsÚslenrO   s         r.   Ú	repeat_kvr¡   Ô   s„   € ð
 2?Ô1DÑ.€EÐ  hØ�‚z€zØÐØ! ! ! ! Q Q Q¨¨a¨a¨a°°°Ð"2Ô3×:Ò:¸5ÐBUÐW\Ð^bÐdlÑmÔm€MØ× Ò  Ð(;¸eÑ(CÀTÈ8ÑTÔTÐTr-   ç        ÚmoduleÚqueryÚkeyÚvalueÚattention_maskÚscalingÚdropoutÚkwargsc                 ó  — t          || j        ¦  «        }t          || j        ¦  «        }	t          j        ||                     dd¦  «        ¦  «        |z  }
|�|
|z   }
t
          j                             |
dt          j        ¬¦  «         	                    |j
        ¦  «        }
t
          j                             |
|| j        ¬¦  «        }
t          j        |
|	¦  «        }|                     dd¦  «                             ¦   «         }||
fS )Nr   r   r^   ©r[   rR   ©ÚpÚtrainingr   )r¡   Únum_key_value_groupsr'   Úmatmulrj   r„   Ú
functionalÚsoftmaxÚfloat32rX   rR   r©   r¯   Ú
contiguous)r£   r¤   r¥   r¦   r§   r¨   r©   rª   Ú
key_statesÚvalue_statesÚattn_weightsÚattn_outputs               r.   Úeager_attention_forwardrº   à   sé   € õ ˜3 Ô ;Ñ<Ô<€JÝ˜U FÔ$?Ñ@Ô@€Lå”<  z×';Ò';¸A¸qÑ'AÔ'AÑBÔBÀWÑL€LØÐ!Ø# nÑ4ˆå”=×(Ò(¨¸2ÅUÄ]Ð(ÑSÔS×VÒVÐW\ÔWbÑcÔc€LÝ”=×(Ò(¨¸È6Ì?Ð(Ñ[Ô[€LÝ”,˜|¨\Ñ:Ô:€KØ×'Ò'¨¨1Ñ-Ô-×8Ò8Ñ:Ô:€Kà˜Ð$Ð$r-   c                   óÎ   ‡ — e Zd ZdZdedefˆ fd„Z	 	 	 ddej        de	ej        ej        f         dz  dej        dz  d	e
dz  d
ee         de	ej        ej        f         fd„Zˆ xZS )ÚXcodec2Attentionz=Multi-headed attention from 'Attention Is All You Need' paperr6   Ú	layer_idxc                 ó®  •— t          ¦   «                              ¦   «          || _        || _        t	          |d|j        |j        z  ¦  «        | _        |j        |j        z  | _	        | j        dz  | _
        |j        | _        d| _        t          j        |j        |j        | j        z  |j        ¬¦  «        | _        t          j        |j        |j        | j        z  |j        ¬¦  «        | _        t          j        |j        |j        | j        z  |j        ¬¦  «        | _        t          j        |j        | j        z  |j        |j        ¬¦  «        | _        d S )NrO   g      à¿Fr€   )r=   r>   r6   r½   rS   rT   rU   rO   rŸ   r°   r¨   Úattention_dropoutÚ	is_causalr„   r…   Úattention_biasÚq_projÚk_projÚv_projÚo_proj©rG   r6   r½   rJ   s      €r.   r>   zXcodec2Attention.__init__ý   sB  ø€ Ý‰Œ×ÒÑÔÐØˆŒØ"ˆŒÝ ¨
°FÔ4FÈ&ÔJdÑ4dÑeÔeˆŒØ$*Ô$>À&ÔB\Ñ$\ˆÔ!Ø”} dÑ*ˆŒØ!'Ô!9ˆÔØˆŒå”iØÔ Ô :¸T¼]Ñ JÐQWÔQfð
ñ 
ô 
ˆŒõ ”iØÔ Ô :¸T¼]Ñ JÐQWÔQfð
ñ 
ô 
ˆŒõ ”iØÔ Ô :¸T¼]Ñ JÐQWÔQfð
ñ 
ô 
ˆŒõ ”iØÔ&¨¬Ñ6¸Ô8JÐQWÔQfð
ñ 
ô 
ˆŒˆˆr-   NrŠ   Úposition_embeddingsr§   Úpast_key_valuesrª   rL   c                 ó&  — |j         d d…         }g |¢d‘| j        ‘R }|                      |¦  «                             |¦  «                             dd¦  «        }|                      |¦  «                             |¦  «                             dd¦  «        }	|                      |¦  «                             |¦  «                             dd¦  «        }
|\  }}t          ||	||d¬¦  «        \  }}	|�|                     |	|
| j	        ¦  «        \  }	}
t          j        | j        j        t          ¦  «        } || ||	|
|f| j        sdn| j        | j        dœ|¤Ž\  }} |j        g |¢d‘R Ž                      ¦   «         }|                      |¦  «        }||fS )Nr^   r   r   )r—   r¢   )r©   r¨   )rf   rO   rÂ   Úviewrj   rÃ   rÄ   rš   Úupdater½   r   Úget_interfacer6   Ú_attn_implementationrº   r¯   r¿   r¨   r�   rµ   rÅ   )rG   rŠ   rÇ   r§   rÈ   rª   Úinput_shapeÚhidden_shapeÚquery_statesr¶   r·   rl   rm   Úattention_interfacer¹   r¸   s                   r.   rt   zXcodec2Attention.forward  sÈ  € ð $Ô)¨#¨2¨#Ô.ˆØ8˜Ð8 bÐ8¨$¬-Ð8Ð8ˆà—{’{ =Ñ1Ô1×6Ò6°|ÑDÔD×NÒNÈqÐRSÑTÔTˆØ—[’[ Ñ/Ô/×4Ò4°\ÑBÔB×LÒLÈQÐPQÑRÔRˆ
Ø—{’{ =Ñ1Ô1×6Ò6°|ÑDÔD×NÒNÈqÐRSÑTÔTˆà&‰ˆˆSõ $8¸ÀjÐRUÐWZÐjkÐ#lÑ#lÔ#lÑ ˆ�jàÐ&Ø'6×'=Ò'=¸jÈ,ÐX\ÔXfÑ'gÔ'gÑ$ˆJ˜å(?Ô(MØŒKÔ,Õ.Eñ)
ô )
Ðð %8Ð$7ØØØØØð	%
ð  $œ}ÐH�C�C°$Ô2HØ”Lð	%
ð 	%
ð ð	%
ð 	%
Ñ!ˆ�\ð *�kÔ)Ð;¨;Ð;¸Ð;Ð;Ð;×FÒFÑHÔHˆØ—k’k +Ñ.Ô.ˆØ˜LÐ(Ð(r-   rv   )r#   r$   r%   r&   r   rx   r>   r'   r+   ry   r
   r   r   rt   r{   r|   s   @r.   r¼   r¼   ù   så   ø€ € € € € àGÐGð
˜}ð 
¸ð 
ð 
ð 
ð 
ð 
ð 
ð4 IMØ.2Ø(,ð))ð ))à”|ð))ð # 5¤<°´Ð#=Ô>ÀÑEð))ð œ tÑ+ð	))ð
  ™ð))ð Ð+Ô,ð))ð 
ˆuŒ|˜Uœ\Ð)Ô	*ð))ð ))ð ))ð ))ð ))ð ))ð ))ð ))r-   r¼   ÚRMSNormc                   óT   ‡ — e Zd Zd	deddfˆ fd„Zdej        dej        fd„Zd„ Zˆ xZ	S )
ÚXcodec2RMSNormç�íµ ÷Æ°>ÚepsrL   Nc                 ó¬   •— t          ¦   «                              ¦   «          t          j        t	          j        |¦  «        ¦  «        | _        || _        dS )z=
        Xcodec2RMSNorm is equivalent to T5LayerNorm
        N)r=   r>   r„   r   r'   ÚonesÚweightÚvariance_epsilon)rG   rT   rÖ   rJ   s      €r.   r>   zXcodec2RMSNorm.__init__B  sD   ø€ õ 	‰Œ×ÒÑÔÐÝ”l¥5¤:¨kÑ#:Ô#:Ñ;Ô;ˆŒØ #ˆÔÐÐr-   rŠ   c                 ó  — |j         }|                     t          j        ¦  «        }|                     d¦  «                             dd¬¦  «        }|t          j        || j        z   ¦  «        z  }| j        |                     |¦  «        z  S )Nr   r^   T)Úkeepdim)	rR   rX   r'   r´   ÚpowÚmeanÚrsqrtrÚ   rÙ   )rG   rŠ   Úinput_dtypeÚvariances       r.   rt   zXcodec2RMSNorm.forwardJ  s|   € Ø#Ô)ˆØ%×(Ò(­¬Ñ7Ô7ˆØ ×$Ò$ QÑ'Ô'×,Ò,¨R¸Ð,Ñ>Ô>ˆØ%­¬°H¸tÔ?TÑ4TÑ(UÔ(UÑUˆØŒ{˜]×-Ò-¨kÑ:Ô:Ñ:Ð:r-   c                 óH   — t          | j        j        ¦  «        › d| j        › �S )Nz, eps=)ry   rÙ   rf   rÚ   )rG   s    r.   Ú
extra_reprzXcodec2RMSNorm.extra_reprQ  s&   € Ý˜œÔ)Ñ*Ô*ÐIÐI°$Ô2GÐIÐIÐIr-   )rÕ   )
r#   r$   r%   rY   r>   r'   r+   rt   rã   r{   r|   s   @r.   rÔ   rÔ   @  sŒ   ø€ € € € € ð$ð $¨ð $¸$ð $ð $ð $ð $ð $ð $ð; U¤\ð ;°e´lð ;ð ;ð ;ð ;ðJð Jð Jð Jð Jð Jð Jr-   rÔ   c                   óÒ   ‡ — e Zd Zdedefˆ fd„Z	 	 	 	 	 ddej        dej        dz  dej        dz  d	e	dz  d
e
dz  deej        ej        f         dz  dee         dej        fd„Zˆ xZS )ÚXcodec2DecoderLayerr6   r½   c                 ó4  •— t          ¦   «                              ¦   «          |j        | _        t          ||¬¦  «        | _        t          |¦  «        | _        t          |j        |j        ¬¦  «        | _	        t          |j        |j        ¬¦  «        | _
        d S )N)r6   r½   ©rÖ   )r=   r>   rT   r¼   Ú	self_attnr~   ÚmlprÔ   Úrms_norm_epsÚinput_layernormÚpost_attention_layernormrÆ   s      €r.   r>   zXcodec2DecoderLayer.__init__V  sƒ   ø€ Ý‰Œ×ÒÑÔÐØ!Ô-ˆÔå)°À9ÐMÑMÔMˆŒå˜fÑ%Ô%ˆŒÝ-¨fÔ.@ÀfÔFYÐZÑZÔZˆÔÝ(6°vÔ7IÈvÔObÐ(cÑ(cÔ(cˆÔ%Ð%Ð%r-   NFrŠ   r§   ro   rÈ   Ú	use_cacherÇ   rª   rL   c           
      óÎ   — |}|                       |¦  «        } | j        d||||||dœ|¤Ž\  }}	||z   }|}|                      |¦  «        }|                      |¦  «        }||z   }|S )N)rŠ   r§   ro   rÈ   rí   rÇ   r,   )rë   rè   rì   ré   )
rG   rŠ   r§   ro   rÈ   rí   rÇ   rª   ÚresidualÚ_s
             r.   rt   zXcodec2DecoderLayer.forward`  s¡   € ð !ˆØ×,Ò,¨]Ñ;Ô;ˆà)˜4œ>ð 
Ø'Ø)Ø%Ø+ØØ 3ð
ð 
ð ð
ð 
Ñˆ�qð ! =Ñ0ˆð !ˆØ×5Ò5°mÑDÔDˆØŸš Ñ/Ô/ˆØ  =Ñ0ˆØÐr-   )NNNFN)r#   r$   r%   r   rx   r>   r'   r+   r*   r
   Úboolry   r   r   rt   r{   r|   s   @r.   rå   rå   U  sÿ   ø€ € € € € ðd˜}ð d¸ð dð dð dð dð dð dð /3Ø04Ø(,Ø!&ØHLðð à”|ðð œ tÑ+ðð Ô&¨Ñ-ð	ð
  ™ðð ˜$‘;ðð # 5¤<°´Ð#=Ô>ÀÑEðð Ð+Ô,ðð 
Œðð ð ð ð ð ð ð r-   rå   c                   ó*   ‡ — e Zd ZdZdˆ fd„	Zd„ Zˆ xZS )ÚXcodec2SnakeBetaa  
    A modified Snake function which uses separate parameters for the magnitude of the periodic components
    Shape:
        - Input: (B, C, T)
        - Output: (B, C, T), same shape as the input
    Parameters:
        - alpha - trainable parameter that controls frequency
        - beta - trainable parameter that controls magnitude
    References:
        - This activation function is a modified version based on this paper by Liu Ziyin, Tilman Hartwig, Masahito Ueda:
        https://huggingface.co/papers/2006.08195
    rP   c                 ó  •— t          ¦   «                              ¦   «          || _        t          t	          j        |¦  «        |z  ¦  «        | _        t          t	          j        |¦  «        |z  ¦  «        | _        d| _        d S )Ng•Ö&è.>)	r=   r>   Úin_featuresr   r'   ÚzerosÚalphaÚbetaÚno_div_by_zero)rG   rõ   r÷   rJ   s      €r.   r>   zXcodec2SnakeBeta.__init__Ž  sm   ø€ Ý‰Œ×ÒÑÔÐØ&ˆÔõ �uœ{¨;Ñ7Ô7¸%Ñ?Ñ@Ô@ˆŒ
Ý�eœk¨+Ñ6Ô6¸Ñ>Ñ?Ô?ˆŒ	à)ˆÔÐÐr-   c                 ó€  — | j                              d¦  «                             d¦  «        }| j                             d¦  «                             d¦  «        }t          j        |¦  «        }t          j        |¦  «        }|d|| j        z   z  t          j        t          j        ||z  ¦  «        d¦  «        z  z   }|S )u’   
        Forward pass of the function.
        Applies the function to the input elementwise.
        SnakeBeta âˆ¶= x + 1/b * sin^2 (xa)
        r   r^   rP   r   )r÷   r”   rø   r'   Úexprù   rÝ   rm   )rG   rŠ   r÷   rø   s       r.   rt   zXcodec2SnakeBeta.forward˜  s°   € ð ”
×$Ò$ QÑ'Ô'×1Ò1°"Ñ5Ô5ˆØŒy×"Ò" 1Ñ%Ô%×/Ò/°Ñ3Ô3ˆÝ”	˜%Ñ Ô ˆÝŒy˜‰ŒˆØ%¨°°tÔ7JÑ0JÑ)KÍuÌyÝŒI�m eÑ+Ñ,Ô,¨añP
ô P
ñ )
ñ 
ˆð Ðr-   )rP   ©r#   r$   r%   r&   r>   rt   r{   r|   s   @r.   ró   ró   €  sV   ø€ € € € € ðð ð*ð *ð *ð *ð *ð *ðð ð ð ð ð ð r-   ró   c                 óX  — |dz  dk    }|dz  }d|z  }d|dz
  z  t           j        z  |z  dz   }|dk    r	d|d	z
  z  }n|d
k    rd|dz
  dz  z  d|d
z
  z  z   }nd}t          j        ||dt          j        ¬¦  «        }|rt          j        | |¦  «        dz   }	nt          j        |¦  «        |z
  }	| dk    r#t          j        dd|ft          j        ¬¦  «        S t          j        d| z  |	z  ¦  «        }
d| z  |z  |
z  }||                     ¦   «         z  }| 	                    dd|¦  «        S )aB  Generates a 1D Kaiser-windowed sinc filter.

    Args:
        cutoff (float): Normalized cutoff frequency (0 to 0.5).
        half_width (float): Transition bandwidth.
        kernel_size (int): Number of filter taps.

    Returns:
        torch.Tensor: A tensor of shape (1, 1, kernel_size) representing the filter.
    r   r   é   gHáz®G@r   gÍÌÌÌÌÌ@g      I@gKê46¼?gffffff!@g      5@g¨WÊ2Ä±â?é   gš™™™™™Ù?gUjö@+0´?r¢   F)rø   ÚperiodicrR   ç      à?rQ   )
ÚmathÚpir'   Úkaiser_windowr´   rV   rö   ÚsincÚsumrÊ   )ÚcutoffÚ
half_widthÚkernel_sizeÚis_evenÚ	half_sizeÚdelta_fÚattenuationrø   r  Útime_indicesÚsinc_filterÚnormalized_filters               r.   Úkaiser_sinc_filter1dr  ©  ss  € ð ˜A‰o Ò"€GØ˜qÑ €Ið �*‰n€GØ˜9 q™=Ñ)­D¬GÑ3°gÑ=ÀÑD€Kà�TÒÐØ˜ sÑ*Ñ+ˆˆØ	˜Ò	Ð	Ø˜ rÑ)¨cÑ1Ñ1°G¸{ÈTÑ?QÑ4RÑRˆˆàˆåÔ'¨¸$ÈÕV[ÔVcÐdÑdÔd€Mð ð =Ý”| Y J°	Ñ:Ô:¸SÑ@ˆˆå”| KÑ0Ô0°9Ñ<ˆð �‚{€{ÝŒ{˜A˜q +Ð.µe´mÐDÑDÔDÐDå”*˜Q ™Z¨,Ñ6Ñ7Ô7€KØ˜F™
 ]Ñ2°[Ñ@Ðð Ð*×.Ò.Ñ0Ô0Ñ0Ðà×!Ò! ! Q¨Ñ4Ô4Ð4r-   c                   ó&   ‡ — e Zd Zdˆ fd„	Zd„ Zˆ xZS )ÚXcodec2DownSample1dr   Nc                 ó¨  •— t          ¦   «                              ¦   «          d|z  }d|z  }|| _        || _        || _        |dk     rt          d¦  «        ‚|dk    rt          d¦  «        ‚|dz  dk    | _        |dz  t          | j        ¦  «        z
  | _        |dz  | _	        || _
        t          |||¦  «        }|                      d|d	¬
¦  «         d S )Nr  ç333333ã?r¢   z(Minimum cutoff must be larger than zero.z'A cutoff above 0.5 does not make sense.r   r   ÚfilterFr:   )r=   r>   r  r  r	  Ú
ValueErrorÚevenrx   Úpad_leftÚ	pad_rightÚstrider  rE   )rG   Úratior	  r  r  r  rJ   s         €r.   r>   zXcodec2DownSample1d.__init__Ø  sß   ø€ Ý‰Œ×ÒÑÔÐØ�u‘ˆØ˜5‘[ˆ
ØˆŒØ$ˆŒØ&ˆÔà�CŠ<ˆ<ÝÐGÑHÔHÐHØ�CŠ<ˆ<ÝÐFÑGÔGÐGà !‘O qÒ(ˆŒ	Ø# qÑ(­3¨t¬y©>¬>Ñ9ˆŒØ$¨Ñ)ˆŒØˆŒÝ% f¨j¸+ÑFÔFˆØ×Ò˜X v¸%ÐÑ@Ô@Ð@Ð@Ð@r-   c                 ó  — |j         d         }t          j        || j        | j        fd¬¦  «        }t          j        || j                             |j        ¦  «         	                    |dd¦  «        | j
        |¬¦  «        }|S )Nr   Ú	replicate©Úmoder^   ©r  Úgroups)rf   ÚFÚpadr  r  Úconv1dr  rX   rR   re   r  )rG   rŠ   ÚchannelsÚouts       r.   rt   zXcodec2DownSample1d.forwardì  s}   € Ø Ô& qÔ)ˆÝœ˜m¨d¬m¸T¼^Ð-LÐS^Ð_Ñ_Ô_ˆÝŒhØàŒK�NŠN˜=Ô.Ñ/Ô/×6Ò6°xÀÀRÑHÔHØ”;Øð
ñ 
ô 
ˆð ˆ
r-   ©r   N©r#   r$   r%   r>   rt   r{   r|   s   @r.   r  r  ×  sR   ø€ € € € € ðAð Að Að Að Að Að(
ð 
ð 
ð 
ð 
ð 
ð 
r-   r  c                   ó&   ‡ — e Zd Zdˆ fd„	Zd„ Zˆ xZS )ÚXcodec2UpSample1dr   Nc                 óÖ  •— t          ¦   «                              ¦   «          || _        |€t          d|z  dz  ¦  «        dz  n|| _        || _        | j        |z  dz
  | _        | j        | j        z  | j        | j        z
  dz  z   | _        | j        | j        z  | j        | j        z
  dz   dz  z   | _        t          d|z  d|z  | j        ¬¦  «        }|  
                    d|d¬	¦  «         d S )
Né   r   r   r  r  )r  r  r	  r  Fr:   )r=   r>   r  rx   r	  r  r$  r  r  r  rE   )rG   r  r	  r  rJ   s       €r.   r>   zXcodec2UpSample1d.__init__ú  só   ø€ Ý‰Œ×ÒÑÔÐØˆŒ
Ø6AÐ6I�3˜q 5™y¨A™~Ñ.Ô.°Ñ2Ð2È{ˆÔØˆŒØÔ# uÑ,¨qÑ0ˆŒØœ 4¤;Ñ.°$Ô2BÀTÄ[Ñ2PÐUVÑ1VÑVˆŒØœ D¤KÑ/°4Ô3CÀdÄkÑ3QÐTUÑ3UÐZ[Ñ2[Ñ[ˆŒå%¨S°5©[ÀSÈ5Á[Ð^bÔ^nÐoÑoÔoˆØ×Ò˜X v¸%ÐÑ@Ô@Ð@Ð@Ð@r-   c           	      óB  — |j         d         }t          j        || j        | j        fd¬¦  «        }| j        t          j        || j                             |j        ¦  «                             |dd¦  «        | j	        |¬¦  «        z  }|d| j
        | j         …f         }|S )Nr   r  r  r^   r!  .)rf   r#  r$  r  Úconv_transpose1dr  rX   rR   re   r  r  r  )rG   rŠ   r&  s      r.   rt   zXcodec2UpSample1d.forward  s¡   € Ø Ô& qÔ)ˆÝœ˜m¨d¬h¸¼Ð-AÈÐTÑTÔTˆØœ
¥QÔ%7ØàŒK�NŠN˜=Ô.Ñ/Ô/×6Ò6°xÀÀRÑHÔHØ”;Øð&
ñ &
ô &
ñ 
ˆð & c¨4¬=¸D¼N¸?Ð+JÐ&JÔKˆØÐr-   r(  r)  r|   s   @r.   r+  r+  ù  sR   ø€ € € € € ð
Að 
Að 
Að 
Að 
Að 
Aðð ð ð ð ð ð r-   r+  c            	       ó@   ‡ — e Zd Z	 	 	 	 d	dedededefˆ fd„Zd„ Zˆ xZS )
ÚXcodec2AntiAliasedActivation1dr   é   Úup_ratioÚ
down_ratioÚup_kernel_sizeÚdown_kernel_sizec                 óæ   •— t          ¦   «                              ¦   «          t          |¦  «        st          d¦  «        ‚|| _        t          ||¦  «        | _        t          ||¦  «        | _        d S )Nz$Activation function must be callable)	r=   r>   ÚcallableÚ	TypeErrorÚactr+  Úupsampler  Ú
downsample)rG   Ú
activationr3  r4  r5  r6  rJ   s         €r.   r>   z'Xcodec2AntiAliasedActivation1d.__init__  si   ø€ õ 	‰Œ×ÒÑÔÐÝ˜
Ñ#Ô#ð 	DÝÐBÑCÔCÐCØˆŒÝ)¨(°NÑCÔCˆŒÝ-¨jÐ:JÑKÔKˆŒˆˆr-   c                 ó„   — |                       |¦  «        }|                      |¦  «        }|                      |¦  «        }|S ru   )r;  r:  r<  rŒ   s     r.   rt   z&Xcodec2AntiAliasedActivation1d.forward$  s;   € ØŸš mÑ4Ô4ˆØŸš Ñ/Ô/ˆØŸš¨Ñ6Ô6ˆàÐr-   )r   r   r2  r2  )r#   r$   r%   rx   r>   rt   r{   r|   s   @r.   r1  r1    s’   ø€ € € € € ð ØØ Ø "ðLð Lð ðLð ð	Lð
 ðLð ðLð Lð Lð Lð Lð Lðð ð ð ð ð ð r-   r1  c                   ó(   ‡ — e Zd ZdZˆ fd„Zd„ Zˆ xZS )ÚXcodec2ResidualUnitza
    A residual unit composed of Snake1d and weight-normalized Conv1d layers with dilations.
    c                 óT  •— t          ¦   «                              ¦   «          d|z  dz  }t          t          |¦  «        ¬¦  «        | _        t          j        ||d||¬¦  «        | _        t          t          |¦  «        ¬¦  «        | _        t          j        ||d¬¦  «        | _	        d S )Nr-  r   ©r=  é   )r	  ÚdilationÚpaddingr   )r	  )
r=   r>   r1  ró   Úsnake1r„   ÚConv1dÚconv1Úsnake2Úconv2)rG   Ú	dimensionrD  r$  rJ   s       €r.   r>   zXcodec2ResidualUnit.__init__1  s™   ø€ Ý‰Œ×ÒÑÔÐØ˜Ñ! aÑ'ˆÝ4Õ@PÐQZÑ@[Ô@[Ð\Ñ\Ô\ˆŒÝ”Y˜y¨)ÀÈXÐ_bÐcÑcÔcˆŒ
Ý4Õ@PÐQZÑ@[Ô@[Ð\Ñ\Ô\ˆŒÝ”Y˜y¨)ÀÐCÑCÔCˆŒ
ˆ
ˆ
r-   c                 ó  — |}|                       |                      |¦  «        ¦  «        }|                      |                      |¦  «        ¦  «        }|j        d         |j        d         z
  dz  }|dk    r|d|| …f         }||z   }|S )ar  
        Forward pass through the residual unit.

        Args:
            hidden_state (`torch.Tensor` of shape `(batch_size, channels, time_steps)`):
                Input tensor .

        Returns:
            output_tensor (`torch.Tensor` of shape `(batch_size, channels, time_steps)`):
                Input tensor after passing through the residual unit.
        r^   r   r   .)rH  rF  rJ  rI  rf   )rG   Úhidden_stateÚoutput_tensorrE  s       r.   rt   zXcodec2ResidualUnit.forward9  s‘   € ð %ˆØŸ
š
 4§;¢;¨}Ñ#=Ô#=Ñ>Ô>ˆØŸ
š
 4§;¢;¨}Ñ#=Ô#=Ñ>Ô>ˆàÔ% bÔ)¨MÔ,?ÀÔ,CÑCÈÑIˆØ�QŠ;ˆ;Ø'¨¨W°g°XÐ-=Ð(=Ô>ˆLØ$ }Ñ4ˆØÐr-   rü   r|   s   @r.   r@  r@  ,  sV   ø€ € € € € ðð ðDð Dð Dð Dð Dðð ð ð ð ð ð r-   r@  c                   ó8   ‡ — e Zd ZdZddededefˆ fd„Zd„ Zˆ xZS )	ÚXcodec2EncoderBlockz&Encoder block used in XCODEC2 encoder.r   r6   r  Ústride_indexc           
      ó´  •— t          ¦   «                              ¦   «          |j        d|z  z  }t          |dz  d¬¦  «        | _        t          |dz  d¬¦  «        | _        t          |dz  d¬¦  «        | _        t          t          |dz  ¦  «        ¬¦  «        | _	        t          j        |dz  |d|z  |t          j        |dz  ¦  «        ¬¦  «        | _        d S )Nr   r   )rD  r   é	   rB  ©r	  r  rE  )r=   r>   Úencoder_hidden_sizer@  Ú	res_unit1Ú	res_unit2Ú	res_unit3r1  ró   rF  r„   rG  r  ÚceilrH  )rG   r6   r  rQ  rK  rJ   s        €r.   r>   zXcodec2EncoderBlock.__init__S  sÖ   ø€ Ý‰Œ×ÒÑÔÐØÔ.°°L±Ñ@ˆ	Ý,¨Y¸!©^ÀaÐHÑHÔHˆŒÝ,¨Y¸!©^ÀaÐHÑHÔHˆŒÝ,¨Y¸!©^ÀaÐHÑHÔHˆŒÝ4Õ@PÐQZÐ^_ÑQ_Ñ@`Ô@`ÐaÑaÔaˆŒÝ”YØ˜‰N˜I°1°v±:ÀfÕVZÔV_Ð`fÐijÑ`jÑVkÔVkð
ñ 
ô 
ˆŒ
ˆ
ˆ
r-   c                 óÔ   — |                       |¦  «        }|                      |¦  «        }|                      |                      |¦  «        ¦  «        }|                      |¦  «        }|S ru   )rV  rW  rF  rX  rH  )rG   rM  s     r.   rt   zXcodec2EncoderBlock.forward^  sX   € Ø—~’~ lÑ3Ô3ˆØ—~’~ lÑ3Ô3ˆØ—{’{ 4§>¢>°,Ñ#?Ô#?Ñ@Ô@ˆØ—z’z ,Ñ/Ô/ˆàÐr-   )r   r   )	r#   r$   r%   r&   r   rx   r>   rt   r{   r|   s   @r.   rP  rP  P  sl   ø€ € € € € Ø0Ð0ð	
ð 	
˜}ð 	
°cð 	
ÈSð 	
ð 	
ð 	
ð 	
ð 	
ð 	
ðð ð ð ð ð ð r-   rP  c                   ó.   ‡ — e Zd ZdZdefˆ fd„Zd„ Zˆ xZS )ÚXcodec2EncoderzXCODEC2 Encoderr6   c                 ó  •— t          ¦   «                              ¦   «          t          j        d|j        dd¬¦  «        | _        g | _        t          |j        ¦  «        D ]+\  }}|dz   }| xj        t          |||¬¦  «        gz  c_        Œ,t          j
        | j        ¦  «        | _        |j        dt          |j        ¦  «        z  z  }t          t          |¦  «        ¬¦  «        | _        t          j        ||j        dd¬¦  «        | _        d S )Nr   rC  r   ©r	  rE  )r  rQ  r   rB  )r=   r>   r„   rG  rU  rH  ÚblockÚ	enumerateÚdownsampling_ratiosrP  Ú
ModuleListÚlenr1  ró   rF  rT   rJ  )rG   r6   rQ  r  Úd_modelrJ   s        €r.   r>   zXcodec2Encoder.__init__j  sÿ   ø€ Ý‰Œ×ÒÑÔÐõ ”Y˜q &Ô"<È!ÐUVÐWÑWÔWˆŒ
àˆŒ
å$-¨fÔ.HÑ$IÔ$Ið 	bð 	bÑ ˆL˜&Ø'¨!Ñ+ˆLØˆJŒJÕ.¨v¸fÐS_Ð`Ñ`Ô`ÐaÑaˆJŒJˆJå”] 4¤:Ñ.Ô.ˆŒ
ØÔ,¨qµC¸Ô8RÑ4SÔ4SÑ/SÑSˆÝ4Õ@PÐQXÑ@YÔ@YÐZÑZÔZˆŒÝ”Y˜w¨Ô(:ÈÐSTÐUÑUÔUˆŒ
ˆ
ˆ
r-   c                 ó®   — |                       |¦  «        }| j        D ]} ||¦  «        }Œ|                      |¦  «        }|                      |¦  «        }|S ru   )rH  r_  rF  rJ  )rG   rM  r£   s      r.   rt   zXcodec2Encoder.forward{  s]   € Ø—z’z ,Ñ/Ô/ˆà”jð 	0ð 	0ˆFØ!˜6 ,Ñ/Ô/ˆLˆLà—{’{ <Ñ0Ô0ˆØ—z’z ,Ñ/Ô/ˆàÐr-   )r#   r$   r%   r&   r   r>   rt   r{   r|   s   @r.   r\  r\  g  s`   ø€ € € € € ØÐðV˜}ð Vð Vð Vð Vð Vð Vð"	ð 	ð 	ð 	ð 	ð 	ð 	r-   r\  c                   óH   ‡ — e Zd Zdefˆ fd„Zdej        dej        fd„Zˆ xZS )ÚXcodec2ResNetBlockr6   c                 óè  •— t          ¦   «                              ¦   «          t          j        d|j        dd¬¦  «        | _        t          j        ¦   «         | _        t          j        |j        |j        ddd¬¦  «        | _	        t          j        d|j        dd¬¦  «        | _
        t          j        ¦   «         | _        |j        | _        t          j        |j        |j        ddd¬¦  «        | _        d S )Né    rÕ   T)Ú
num_groupsÚnum_channelsrÖ   Úaffiner   r   rT  )r=   r>   r„   Ú	GroupNormrT   Únorm1ÚSiLUÚactivation1rG  rH  Únorm2Úactivation2Úactivation_dropoutrJ  r‰   s     €r.   r>   zXcodec2ResNetBlock.__init__ˆ  sÌ   ø€ Ý‰Œ×ÒÑÔÐÝ”\¨R¸fÔ>PÐVZÐcgÐhÑhÔhˆŒ
Ýœ7™9œ9ˆÔÝ”Y˜vÔ1°6Ô3EÐSTÐ]^ÐhiÐjÑjÔjˆŒ
Ý”\¨R¸fÔ>PÐVZÐcgÐhÑhÔhˆŒ
Ýœ7™9œ9ˆÔØ"(Ô";ˆÔÝ”Y˜vÔ1°6Ô3EÐSTÐ]^ÐhiÐjÑjÔjˆŒ
ˆ
ˆ
r-   rŠ   rL   c                 ó¸  — |                      dd¦  «        }|}|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }t          j                             || j	        | j
        ¬¦  «        }|                      |¦  «        }||z                         dd¦  «        S )Nr   r   r­   )rj   rn  rp  rH  rq  rr  r„   r²   r©   rs  r¯   rJ  ©rG   rŠ   rï   s      r.   rt   zXcodec2ResNetBlock.forward’  sÄ   € Ø%×/Ò/°°1Ñ5Ô5ˆØ ˆØŸ
š
 =Ñ1Ô1ˆØ×(Ò(¨Ñ7Ô7ˆØŸ
š
 =Ñ1Ô1ˆØŸ
š
 =Ñ1Ô1ˆØ×(Ò(¨Ñ7Ô7ˆÝœ×-Ò-¨m¸tÔ?VÐaeÔanÐ-ÑoÔoˆØŸ
š
 =Ñ1Ô1ˆØ Ñ(×3Ò3°A°qÑ9Ô9Ð9r-   r�   r|   s   @r.   rg  rg  ‡  sq   ø€ € € € € ðk˜}ð kð kð kð kð kð kð
: U¤\ð 
:°e´lð 
:ð 
:ð 
:ð 
:ð 
:ð 
:ð 
:ð 
:r-   rg  c                   ó¼   ‡ — e Zd ZdZdefˆ fd„Zdd„Zdej        dej        fd„Z	dd
ej        de
dej        fd„Zd
ej        deej        ej        f         fd„Zˆ xZS )ÚXcodec2FiniteScalarQuantizationa!  
    Finite Scalar Quantization (FSQ) module that quantizes continuous latent representations into discrete codes.
    Original code: https://github.com/lucidrains/vector-quantize-pytorch/blob/353d46027888dfb140c3c65a67a7356f1492d71d/vector_quantize_pytorch/finite_scalar_quantization.py#L64

    Original modeling uses `ResidualFSQ` with a single quantizer: https://huggingface.co/HKUSTAudio/xcodec2/blob/main/vq/codec_decoder_vocos.py#L389
    But we can directly use FSQ since a main feature of Xcodec2 is that it uses a single codebook.
    r6   c                 ó:  •— t          ¦   «                              ¦   «          t          |j        ¦  «        | _        |                      ¦   «         \  }}}|                      d|d¬¦  «         |                      d|d¬¦  «         |                      d|d¬¦  «         d S )NÚlevelsFr:   ÚbasisÚcodebook)r=   r>   ÚlistÚquantization_levelsÚ_compute_buffersrE   )rG   r6   ry  rz  r{  rJ   s        €r.   r>   z(Xcodec2FiniteScalarQuantization.__init__¨  s›   ø€ Ý‰Œ×ÒÑÔÐÝ#'¨Ô(BÑ#CÔ#CˆÔ Ø"&×"7Ò"7Ñ"9Ô"9Ñˆ��xØ×Ò˜X v¸%ÐÑ@Ô@Ð@Ø×Ò˜W e¸ÐÑ>Ô>Ð>Ø×Ò˜Z¨¸eÐÑDÔDÐDÐDÐDr-   Nc                 ó¨  — t          j        | j        t           j        |¬¦  «        }t          j        t          j        dg| j        dd…         z   |¬¦  «        dt           j        ¬¦  «        }t          j        t          t          j        | j        ¦  «        ¦  «        |¬¦  «         	                    d¦  «        }||z  |z  }|dz  }||z
  |z  }|||fS )	zFCompute the levels, basis, and codebook buffers for the FSQ quantizer.)rR   rH   r   Nr^   ©rH   r   r¬   r   )
r'   Útensorr}  Úint32ÚcumprodrV   rx   ÚnpÚprodr”   )rG   rH   ry  rz  ÚindicesÚlevel_indicesr  r{  s           r.   r~  z0Xcodec2FiniteScalarQuantization._compute_buffers°  sÐ   € å”˜dÔ6½e¼kÐRXÐYÑYÔYˆÝ”ÝŒL˜!˜˜tÔ7¸¸¸Ô<Ñ<ÀVÐLÑLÔLÐRSÕ[`Ô[fð
ñ 
ô 
ˆõ ”,�s¥2¤7¨4Ô+CÑ#DÔ#DÑEÔEÈfÐUÑUÔU×_Ò_Ð`bÑcÔcˆØ  EÑ)¨VÑ3ˆØ˜q‘[ˆ
Ø! JÑ.°*Ñ<ˆØ�u˜hÐ&Ð&r-   r†  rL   c                 óx   — |                      d¦  «        }|| j        z  | j        z  }| j        dz  }||z
  |z  }|S )z`
        Convert integer codebook indices to normalized per-dimension codes in [-1, 1].
        r^   r   )r”   rz  ry  )rG   r†  r‡  r  Úcodess        r.   Ú_indices_to_codesz1Xcodec2FiniteScalarQuantization._indices_to_codes¼  sJ   € ð ×#Ò# BÑ'Ô'ˆØ  D¤JÑ.°$´+Ñ=ˆØ”[ AÑ%ˆ
Ø Ñ+¨zÑ9ˆØˆr-   çü©ñÒMbP?rŠ   rÖ   c                 óÔ   — | j         dz
  d|z   z  dz  }t          j        | j         dz  dk    dd¦  «        }||z                       ¦   «         }||z                        ¦   «         |z  |z
  S )að  
        Constrain `hidden_states` to the valid quantization range for each dimension.

        Uses a scaled tanh to soft-clip values into the interval
        $[-(L-1)/2, (L-1)/2]$ (offset by 0.5 for even-level dimensions), where $L$ is
        the number of quantization levels. The small `eps` margin prevents values from
        saturating exactly at the boundary, which would zero out gradients.

        Args:
            hidden_states (`torch.Tensor`): Continuous input to be bounded.
            eps (`float`, *optional*, defaults to `1e-3`):
                Small margin added to the level range to avoid gradient saturation at boundaries.

        Returns:
            `torch.Tensor`: Bounded values in the valid quantization range.
        r   r   r   r  r¢   )ry  r'   ÚwhereÚatanhÚtanh)rG   rŠ   rÖ   Ú
half_rangeÚoffsetÚshifts         r.   Úboundz%Xcodec2FiniteScalarQuantization.boundÆ  sr   € ð" ”k A‘o¨!¨c©'Ñ2°QÑ6ˆ
Ý”˜Tœ[¨1™_°Ò1°3¸Ñ<Ô<ˆØ˜*Ñ$×+Ò+Ñ-Ô-ˆØ Ñ%×+Ò+Ñ-Ô-°
Ñ:¸VÑCÐCr-   c                 ó\  — |j         }t          |j        j        t          ¦  «        r|j        j        dk    r|j        j        nd}t          |d¬¦  «        5  |                     ¦   «         }| j        dz  }|                      |¦  «        }| 	                    ¦   «         }|||z
   
                    ¦   «         z   }||z  }||z  |z   }|| j        z                       d¬¦  «                             t          j        ¦  «        }d d d ¦  «         n# 1 swxY w Y   |                     |¦  «        |fS )Nr_   r`   Fra   r   r^   rd   )rR   rg   rH   rh   ri   r   rY   ry  r“  ÚroundÚdetachrz  r  rX   r'   r‚  )	rG   rŠ   Úoriginal_dtyperb   r  Úroundedr‰  Úcode_scaledr†  s	            r.   rt   z'Xcodec2FiniteScalarQuantization.forwardÜ  s€  € à&Ô,ˆõ ˜-Ô.Ô3µSÑ9Ô9ðØ>KÔ>RÔ>WÐ[`Ò>`Ð>`ð Ô Ô%Ð%àð 	õ
 ¨¸UÐCÑCÔCð 
	Mð 
	MØ)×/Ò/Ñ1Ô1ˆMØœ¨Ñ)ˆJà ŸJšJ }Ñ5Ô5ˆMØ#×)Ò)Ñ+Ô+ˆGØ! W¨}Ñ%<×$DÒ$DÑ$FÔ$FÑFˆEØ˜JÑ&ˆEà  :Ñ-°Ñ;ˆKØ" T¤ZÑ/×4Ò4¸Ð4Ñ<Ô<×?Ò?ÅÄÑLÔLˆGð
	Mð 
	Mð 
	Mñ 
	Mô 
	Mð 
	Mð 
	Mð 
	Mð 
	Mð 
	Mð 
	Møøøð 
	Mð 
	Mð 
	Mð 
	Mð �xŠx˜Ñ'Ô'¨Ð0Ð0s   ÁB*DÄDÄDru   )r‹  )r#   r$   r%   r&   r   r>   r~  r'   r+   rŠ  rY   r“  ry   rt   r{   r|   s   @r.   rw  rw  Ÿ  sû   ø€ € € € € ðð ðE˜}ð Eð Eð Eð Eð Eð Eð
'ð 
'ð 
'ð 
'ð¨¬ð ¸%¼,ð ð ð ð ðDð D 5¤<ð D°eð DÀuÄ|ð Dð Dð Dð Dð,1 U¤\ð 1°e¸E¼LÈ%Ì,Ð<VÔ6Wð 1ð 1ð 1ð 1ð 1ð 1ð 1ð 1r-   rw  c                   óL   ‡ — e Zd ZdZdefˆ fd„Zdej        dej        fd„Zˆ xZ	S )ÚXcodec2ISTFTHeadzù
    Head for converting decoder outputs to waveform via STFT projection and ISTFT.

    Uses custom "same" padding ISTFT from Vocos:
    https://github.com/gemelo-ai/vocos/blob/c859e3b7b534f3776a357983029d34170ddd6fc3/vocos/spectral_ops.py#L47
    r6   c                 óV  •— t          ¦   «                              ¦   «          t          j        |j        |j        dz   ¦  «        | _        |j        | _        |j        | _        | j        | j        z
  dz  | _        t          j
        |j        ¦  «        }|                      d|d¬¦  «         d S )Nr   ÚwindowFr:   )r=   r>   r„   r…   rT   Ún_fftÚlinearÚ
hop_lengthrE  r'   Úhann_windowrE   )rG   r6   r�  rJ   s      €r.   r>   zXcodec2ISTFTHead.__init__ú  s‘   ø€ Ý‰Œ×ÒÑÔÐÝ”i Ô 2°F´LÀ1Ñ4DÑEÔEˆŒØ”\ˆŒ
Ø Ô+ˆŒØœ
 T¤_Ñ4¸Ñ:ˆŒÝÔ" 6¤<Ñ0Ô0ˆØ×Ò˜X v¸%ÐÑ@Ô@Ð@Ð@Ð@r-   rŠ   rL   c                 ó  — |                       |¦  «                             dd¦  «        }|                     dd¬¦  «        \  }}|                     ¦   «         }|                     ¦   «         }t	          j        |¦  «                             d¬¦  «        }|t	          j        d|z  ¦  «        z  }t          j                             || j	        dd¬¦  «        }|| j
        d d d …d f         z  }|j        d	         }|dz
  | j        z  | j	        z   }t          j        |d|fd| j	        fd| j        f¬
¦  «        d d …dd| j        | j         …f         }	t          j        | j
                             ¦   «                              d|d	¦  «                             dd¦  «        d|fd| j	        fd| j        f¬
¦  «                             ¦   «         | j        | j         …         }
|
                     d¬¦  «        }
|	|
z  }	|	                     d¦  «        S )Nr   r   rd   g      Y@)Úmaxy              ð?Úbackward)r[   Únormr^   )Úoutput_sizer	  r  r   g•dyáý¥=)Úmin)rŸ  rj   ÚchunkrY   r'   rû   ÚclampÚfftÚirfftrž  r�  rf   r   r#  ÚfoldrE  Úsquarere   Úsqueezer”   )rG   rŠ   Ú	stft_predÚ	magnitudeÚphaseÚspectrogram_complexÚtime_framesÚ
num_framesr¦  ÚaudioÚwindow_envelopes              r.   rt   zXcodec2ISTFTHead.forward  sÿ  € Ø—K’K Ñ.Ô.×8Ò8¸¸AÑ>Ô>ˆ	Ø$Ÿ?š?¨1°!˜?Ñ4Ô4Ñˆ	�5à—O’OÑ%Ô%ˆ	Ø—’‘”ˆå”I˜iÑ(Ô(×.Ò.°3Ð.Ñ7Ô7ˆ	Ø'­%¬)°B¸±JÑ*?Ô*?Ñ?Ðõ ”i—o’oÐ&9¸4¼:È1ÐS]�oÑ^Ô^ˆØ! D¤K°°a°a°a¸°Ô$>Ñ>ˆØ(Ô.¨rÔ2ˆ
Ø! A‘~¨¬Ñ8¸4¼:ÑEˆÝ”ØØ˜KÐ(Ø˜DœJ˜Ø�t”Ð'ð	
ñ 
ô 
ð
 ˆ!ˆ!ˆQ��4”< 4¤< -Ð/Ð
/ô1ˆõ œ&ØŒK×ÒÑ Ô ×'Ò'¨¨:°rÑ:Ô:×DÒDÀQÈÑJÔJØ˜KÐ(Ø˜DœJ˜Ø�t”Ð'ð	
ñ 
ô 
÷
 Š'‰)Œ)�D”L D¤L =Ð0ô2ˆð *×/Ò/°EÐ/Ñ:Ô:ˆØ˜Ñ'ˆØ�Š˜qÑ!Ô!Ð!r-   ©
r#   r$   r%   r&   r   r>   r'   r+   rt   r{   r|   s   @r.   r›  r›  ò  s{   ø€ € € € € ðð ðA˜}ð Að Að Að Að Að Að!" U¤\ð !"°e´lð !"ð !"ð !"ð !"ð !"ð !"ð !"ð !"r-   r›  c                   ó†   ‡ — e Zd Zdefˆ fd„Zdej        dej        fd„Zdej        deej        ej        f         fd„Z	ˆ xZ
S )ÚXcodec2Quantizerr6   c                 ó4  •— t          ¦   «                              ¦   «          t          |¦  «        | _        t	          j        |j        t          |j        ¦  «        ¦  «        | _	        t	          j        t          |j        ¦  «        |j        ¦  «        | _
        d S ru   )r=   r>   rw  Ú	quantizerr„   r…   Úquantization_dimrc  r}  Ú
project_inÚproject_outr‰   s     €r.   r>   zXcodec2Quantizer.__init__(  sq   ø€ Ý‰Œ×ÒÑÔÐÝ8¸Ñ@Ô@ˆŒÝœ) FÔ$;½SÀÔA[Ñ=\Ô=\Ñ]Ô]ˆŒÝœ9¥S¨Ô)CÑ%DÔ%DÀfÔF]Ñ^Ô^ˆÔÐÐr-   r†  rL   c                 óz   — |                      d¦  «        }| j        j        |         }|                      |¦  «        S ©Nr^   )r®  r»  r{  r¾  )rG   r†  r‰  s      r.   Ú
from_codeszXcodec2Quantizer.from_codes.  s6   € Ø—/’/ "Ñ%Ô%ˆØ”Ô'¨Ô0ˆØ×Ò Ñ&Ô&Ð&r-   rŠ   c                 ó   — |                       |¦  «        }|j        }| j                             |¦  «        }|                      |¦  «        \  }}|                      |                     |¦  «        ¦  «        }|                     d¦  «        }||fS rÀ  )r½  rR   r»  r“  r¾  rX   r”   )rG   rŠ   r—  Úquantized_outr†  s        r.   rt   zXcodec2Quantizer.forward3  s…   € ØŸš¨Ñ6Ô6ˆØ&Ô,ˆØœ×,Ò,¨]Ñ;Ô;ˆØ!%§¢°Ñ!>Ô!>Ñˆ�wØ×(Ò(¨×)9Ò)9¸.Ñ)IÔ)IÑJÔJˆØ×#Ò# BÑ'Ô'ˆØ˜gÐ%Ð%r-   )r#   r$   r%   r   r>   r'   r+   rÁ  ry   rt   r{   r|   s   @r.   r¹  r¹  '  s£   ø€ € € € € ð_˜}ð _ð _ð _ð _ð _ð _ð' %¤,ð '°5´<ð 'ð 'ð 'ð 'ð
& U¤\ð &°e¸E¼LÈ%Ì,Ð<VÔ6Wð &ð &ð &ð &ð &ð &ð &ð &r-   r¹  c                   óL   ‡ — e Zd ZdZdefˆ fd„Zdej        dej        fd„Zˆ xZ	S )ÚXcodec2DecoderzVVocos-based decoder with ResNet, Transformer, and ISTFT head for audio reconstruction.r6   c                 óæ  •‡— t          ¦   «                              ¦   «          t          j        ‰j        ‰j        j        z   ‰j        ¦  «        | _        t          j        ‰j        ‰j        dd¬¦  «        | _        t          j	        t          ‰¦  «        t          ‰¦  «        g¦  «        | _        ‰j        | _        t          ‰¬¦  «        | _        t          j	        ˆfd„t          ‰j        ¦  «        D ¦   «         ¦  «        | _        t          j	        t          ‰¦  «        t          ‰¦  «        g¦  «        | _        t          j        ‰j        d¬¦  «        | _        t+          ‰¦  «        | _        d S )NrC  r   r^  )r6   c                 ó0   •— g | ]}t          ‰|¦  «        ‘ŒS r,   )rå   )Ú.0r½   r6   s     €r.   ú
<listcomp>z+Xcodec2Decoder.__init__.<locals>.<listcomp>H  s$   ø€ ÐeÐeÐe¸	Õ  ¨Ñ3Ô3ÐeÐeÐer-   rÕ   rç   )r=   r>   r„   r…   rT   Úsemantic_model_configÚfcrG  Úembedrb  rg  Ú	prior_netrU   r4   Ú
rotary_embÚrangeÚnum_hidden_layersÚlayersÚpost_netÚ	LayerNormr¥  r›  Úheadr‰   s    `€r.   r>   zXcodec2Decoder.__init__@  s3  øø€ Ý‰Œ×ÒÑÔÐÝ”)˜FÔ.°Ô1MÔ1YÑYÐ[aÔ[mÑnÔnˆŒÝ”Y˜vÔ1°6Ô3EÐSTÐ^_Ð`Ñ`Ô`ˆŒ
ÝœÕ(:¸6Ñ(BÔ(BÕDVÐW]ÑD^ÔD^Ð'_Ñ`Ô`ˆŒØ#)Ô#=ˆÔ Ý0¸Ð?Ñ?Ô?ˆŒÝ”mØeÐeÐeÐeÅUÈ6ÔKcÑEdÔEdÐeÑeÔeñ
ô 
ˆŒõ œÕ'9¸&Ñ'AÔ'AÕCUÐV\ÑC]ÔC]Ð&^Ñ_Ô_ˆŒÝ”L Ô!3¸Ð>Ñ>Ô>ˆŒ	Ý$ VÑ,Ô,ˆŒ	ˆ	ˆ	r-   rŠ   rL   c                 ó  — |                       |¦  «        }|                     dd¦  «        }|                      |¦  «        }|                     dd¦  «        }| j        D ]} ||¦  «        }Œt	          j        | j        |j        ¬¦  «                             d¦  «        }|  	                    ||¦  «        }| j
        D ]} ||fd|i|¤Ž}Œ| j        D ]} ||¦  «        }Œ|                      |                      |¦  «        ¦  «        S )Nr   r   r€  r   rÇ   )rË  rj   rÌ  rÍ  r'   rV   rU   rH   r”   rÎ  rÑ  rÒ  rÔ  r¥  )rG   rŠ   rª   Úlayerro   rÇ   s         r.   rt   zXcodec2Decoder.forwardN  s&  € ØŸš Ñ.Ô.ˆØ%×/Ò/°°1Ñ5Ô5ˆØŸ
š
 =Ñ1Ô1ˆØ%×/Ò/°°1Ñ5Ô5ˆð ”^ð 	1ð 	1ˆEØ!˜E -Ñ0Ô0ˆMˆMõ ”| DÔ$<À]ÔEYÐZÑZÔZ×dÒdÐefÑgÔgˆØ"Ÿošo¨m¸\ÑJÔJÐØ”[ð 	dð 	dˆEØ!˜E -ÐcÐcÐEXÐcÐ\bÐcÐcˆMˆMð ”]ð 	1ð 	1ˆEØ!˜E -Ñ0Ô0ˆMˆMà�yŠy˜Ÿš =Ñ1Ô1Ñ2Ô2Ð2r-   r·  r|   s   @r.   rÅ  rÅ  =  sp   ø€ € € € € Ø`Ð`ð-˜}ð -ð -ð -ð -ð -ð -ð3 U¤\ð 3ÀÄð 3ð 3ð 3ð 3ð 3ð 3ð 3ð 3r-   rÅ  c                   óH   ‡ — e Zd Zdefˆ fd„Zdej        dej        fd„Zˆ xZS )ÚXcodec2SemanticAdapterr6   c                 ó8  •— t          ¦   «                              ¦   «          t          j        |j        j        |j        j        ddd¬¦  «        | _        t          j        ¦   «         | _        t          j        |j        j        |j        j        ddd¬¦  «        | _	        t          j        ¦   «         | _
        t          j        |j        j        |j        j        ddd¬¦  «        | _        t          j        |j        j        |j        j        ddd¬¦  «        | _        d S )Nr   r   F)Úin_channelsÚout_channelsr	  rE  r�   T)r	  rE  r�   )r=   r>   r„   rG  rÊ  rT   rH  ÚReLUÚact1rJ  Úact2Úconv3Úconv4r‰   s     €r.   r>   zXcodec2SemanticAdapter.__init__j  s	  ø€ Ý‰Œ×ÒÑÔÐÝ”YØÔ4Ô@ØÔ5ÔAØØØð
ñ 
ô 
ˆŒ
õ ”G‘I”IˆŒ	Ý”YØÔ(Ô4ØÔ(Ô4ØØØð
ñ 
ô 
ˆŒ
õ ”G‘I”IˆŒ	Ý”YØÔ(Ô4ØÔ(Ô4ØØØð
ñ 
ô 
ˆŒ
õ ”YØÔ4Ô@ØÔ5ÔAØØØð
ñ 
ô 
ˆŒ
ˆ
ˆ
r-   rŠ   rL   c                 ó  — |                       |¦  «        }|                      |¦  «        }|}|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }||z   }|                      |¦  «        }|S ru   )rH  rÝ  rJ  rÞ  rß  rà  ru  s      r.   rt   zXcodec2SemanticAdapter.forward‹  s}   € ØŸ
š
 =Ñ1Ô1ˆØŸ	š	 -Ñ0Ô0ˆØ ˆØŸ
š
 =Ñ1Ô1ˆØŸ	š	 -Ñ0Ô0ˆØŸ
š
 =Ñ1Ô1ˆØ%¨Ñ0ˆØŸ
š
 =Ñ1Ô1ˆØÐr-   r�   r|   s   @r.   rØ  rØ  i  sk   ø€ € € € € ð
˜}ð 
ð 
ð 
ð 
ð 
ð 
ðB	 U¤\ð 	°e´lð 	ð 	ð 	ð 	ð 	ð 	ð 	ð 	r-   rØ  c                   óf   ‡ — e Zd ZU eed<   dZdZdZdZdgZ	dZ
dZdZdZdZdZdZeedœZˆ fd	„Zˆ xZS )
ÚXcodec2PreTrainedModelr6   Úxcodec2)rµ  TNrÈ   Úinput_values)rŠ   Ú
attentionsc                 óÒ  •— t          ¦   «                              |¦  «         t          |t          ¦  «        r4t	          j        |j        ¦  «         t	          j        |j        ¦  «         d S t          |t          ¦  «        r5t          j
        |j        ¦  «        }t	          j        |j        |¦  «         d S t          |t          ¦  «        rt|                     |j        j        ¬¦  «        \  }}}t	          j        |j        |¦  «         t	          j        |j        |¦  «         t	          j        |j        |¦  «         d S t          |t(          ¦  «        rBt+          d|j        z  d|j        z  |j        ¦  «        }t	          j        |j        |¦  «         d S t          |t2          ¦  «        r<t+          |j        |j        |j        ¦  «        }t	          j        |j        |¦  «         d S d S )Nr€  r  r  )r=   Ú_init_weightsrg   ró   ÚinitÚzeros_r÷   rø   r›  r'   r¡  rž  Úcopy_r�  rw  r~  ry  rH   rz  r{  r+  r  r  r	  r  r  r  r  )rG   r£   r�  ry  rz  r{  Úfilter_tensorrJ   s          €r.   rè  z$Xcodec2PreTrainedModel._init_weights«  s¯  ø€ Ý‰Œ×Ò˜fÑ%Ô%Ð%Ý�fÕ.Ñ/Ô/ð 	5ÝŒK˜œÑ%Ô%Ð%ÝŒK˜œÑ$Ô$Ð$Ð$Ð$Ý˜Õ 0Ñ1Ô1ð 	5ÝÔ& v¤|Ñ4Ô4ˆFÝŒJ�v”} fÑ-Ô-Ð-Ð-Ð-Ý˜Õ ?Ñ@Ô@ð 
	5Ø&,×&=Ò&=ÀVÄ]ÔEYÐ&=Ñ&ZÔ&ZÑ#ˆF�E˜8ÝŒJ�v”} fÑ-Ô-Ð-ÝŒJ�v”| UÑ+Ô+Ð+ÝŒJ�v”¨Ñ1Ô1Ð1Ð1Ð1Ý˜Õ 1Ñ2Ô2ð 	5Ý0°°v´|Ñ1CÀSÈ6Ì<ÑEWÐY_ÔYkÑlÔlˆMÝŒJ�v”} mÑ4Ô4Ð4Ð4Ð4Ý˜Õ 3Ñ4Ô4ð 	5Ý0°´ÀÔ@QÐSYÔSeÑfÔfˆMÝŒJ�v”} mÑ4Ô4Ð4Ð4Ð4ð	5ð 	5r-   )r#   r$   r%   r   r)   Úbase_model_prefixÚinput_modalitiesÚsupports_gradient_checkpointingÚ_no_split_modulesÚ_skip_keys_device_placementÚ_supports_flash_attnÚ_supports_sdpaÚ_supports_flex_attnÚ_supports_cache_classÚ_supports_attention_backendÚ_can_compile_fullgraphÚmain_input_namerå   Ú_can_record_outputsrè  r{   r|   s   @r.   rã  rã  —  s¢   ø€ € € € € € àÐÐÑØ!ÐØ!ÐØ&*Ð#ØÐØ#4Ð"5ÐØÐØ€NØÐØ ÐØ"&ÐØ!ÐØ$€Oà,Ø)ðð Ðð
5ð 5ð 5ð 5ð 5ð 5ð 5ð 5ð 5r-   rã  z!Xcodec2 neural audio codec model.)Úcustom_introc                   ó¼  ‡ — e Zd ZeZdefˆ fd„Zee	 	 	 ddej	        dej	        dej	        dz  dej	        dz  d	e
d
ee         deez  fd„¦   «         ¦   «         Zee	 	 ddej	        dz  dej	        dz  d
ee         deez  fd„¦   «         ¦   «         Zee	 	 	 ddej	        dej	        dej	        dz  dej	        dz  d	e
d
ee         deez  fd„¦   «         ¦   «         Zˆ xZS )ÚXcodec2Modelr6   c                 óâ  •— t          ¦   «                              |¦  «         |j        | _        t          j        |j        ¦  «        | _        t          |¦  «        | _        t          |¦  «        | _
        t          j        |j        |j        j        z   |j        |j        j        z   ¦  «        | _        t          |¦  «        | _        t#          |¦  «        | _        |                      ¦   «          d S ru   )r=   r>   r   r   Úfrom_configrÊ  Úsemantic_encoderrØ  Úsemantic_adapterr\  Úacoustic_encoderr„   r…   rT   Ú
fc_encoderr¹  r»  rÅ  Úacoustic_decoderÚ	post_initr‰   s     €r.   r>   zXcodec2Model.__init__Ä  sÇ   ø€ Ý‰Œ×Ò˜Ñ Ô Ð à Ô+ˆŒÝ )Ô 5°fÔ6RÑ SÔ SˆÔÝ 6°vÑ >Ô >ˆÔÝ .¨vÑ 6Ô 6ˆÔÝœ)ØÔ Ô!=Ô!IÑIØÔ Ô!=Ô!IÑIñ
ô 
ˆŒõ *¨&Ñ1Ô1ˆŒÝ .¨vÑ 6Ô 6ˆÔà�ŠÑÔÐÐÐr-   NFrå  Úinput_featuresÚpadding_maskÚinput_features_maskÚoutput_latentsrª   rL   c                 ó@  — t          j        ¦   «         5  |                      ||¬¦  «        }ddd¦  «         n# 1 swxY w Y   |j                             dd¦  «        }|                      |¦  «        }|                      |¦  «        }	t          j        ||	gd¬¦  «        }
|                      |
                     dd¦  «        ¦  «        }
|  	                    |
¦  «        \  }}|                     dd¦  «        }|                     dd¦  «        }d}|�y| 
                    dd¬¦  «        }|| j        z  }t          j        |j        d         |j        ¬	¦  «                             dd¦  «        }||k                          |j        ¦  «        }t%          ||r|nd|¬
¦  «        S )a  
        input_values (`torch.Tensor` of shape `(batch_size, 1, sequence_length)`):
            Input audio waveform.
        input_features (`torch.Tensor` of shape `(batch_size, mel_bins, time_steps)`):
            Input audio mel spectrogram for semantic encoding.
        padding_mask (`torch.Tensor` of shape `(batch_size, 1, sequence_length)`):
            Padding mask used to pad `input_values`.
        input_features_mask (`torch.Tensor` of shape `(batch_size, time_steps)`, *optional*):
            Attention mask for the spectrogram input to the semantic encoder. `1` for valid frames, `0` for padding.
        output_latents (`bool`, *optional*, defaults to `False`):
            Whether to return the continuous latent representation from the quantizer.
        )r§   Nr   r   rd   r^   T)r[   rÜ   r€  )r    r!   r"   )r'   rz   rÿ  Úlast_hidden_staterj   r   r  rk   r  r»  r  r   rV   rf   rH   rÊ   rX   rR   r0   )rG   rå  r  r  r  r  rª   Úsemantic_outputÚsemantic_hidden_statesÚacoustic_hidden_statesrŠ   r!   r    r"   Úaudio_lengthÚtoken_lengthÚidxs                    r.   ÚencodezXcodec2Model.encodeÔ  sû  € õ2 Œ]‰_Œ_ð 	hð 	hØ"×3Ò3°NÐSfÐ3ÑgÔgˆOð	hð 	hð 	hñ 	hô 	hð 	hð 	hð 	hð 	hð 	hð 	høøøð 	hð 	hð 	hð 	hà!0Ô!B×!LÒ!LÈQÐPQÑ!RÔ!RÐØ!%×!6Ò!6Ð7MÑ!NÔ!NÐð "&×!6Ò!6°|Ñ!DÔ!DÐÝœ	Ð#9Ð;QÐ"RÐXYÐZÑZÔZˆØŸš¨×(?Ò(?ÀÀ1Ñ(EÔ(EÑFÔFˆð  $Ÿ~š~¨mÑ<Ô<Ñˆ�Ø×#Ò# A qÑ)Ô)ˆØ!×+Ò+¨A¨qÑ1Ô1ˆð  ÐØÐ#Ø'×+Ò+°¸DÐ+ÑAÔAˆLØ'¨4¬?Ñ:ˆLÝ”,˜{Ô0°Ô4¸\Ô=PÐQÑQÔQ×VÒVÐWXÐZ\Ñ]Ô]ˆCØ # lÒ 2×6Ò6°|Ô7IÑJÔJÐå#Ø#Ø-Ð7�G�G°4Ø-ð
ñ 
ô 
ð 	
s   ”8¸<¿<r    r!   c                 óò   — |€|€t          d¦  «        ‚|�/| j                             |                     dd¦  «        ¦  «        }n|                     dd¦  «        } | j        |fi |¤Ž}t          |¬¦  «        S )a3  
        audio_codes (`torch.LongTensor`  of shape `(batch_size, 1, codes_length)`):
            Discrete code indices computed using `model.encode`.
        latents (torch.Tensor of shape `(batch_size, dimension, time_steps)`, *optional*):
            Quantized continuous representation of input.
        Nz3Either `latents` or `audio_codes` must be provided.r   r   )r   )r  r»  rÁ  rj   r  r2   )rG   r    r!   rª   Úrecon_audios        r.   ÚdecodezXcodec2Model.decode
  sŠ   € ð ˆ?˜{Ð2ÝÐRÑSÔSÐSàÐ"Ø”n×/Ò/°×0EÒ0EÀaÈÑ0KÔ0KÑLÔLˆGˆGà×'Ò'¨¨1Ñ-Ô-ˆGà+�dÔ+¨GÐ>Ð>°vÐ>Ð>ˆÝ#°Ð=Ñ=Ô=Ð=r-   c                 óè   — |j         d         }|                      ||||dd¬¦  «        } | j        d	|j        ddœ|¤Žd         dd|…f         }	t	          |	|j        |r|j        nd|j        ¬¦  «        S )
a  
        input_values (`torch.Tensor` of shape `(batch_size, 1, sequence_length)`):
            Input audio waveform.
        input_features (`torch.Tensor` of shape `(batch_size, mel_bins, time_steps)`):
            Input audio mel spectrogram for semantic encoding.
        padding_mask (`torch.Tensor` of shape `(batch_size, 1, sequence_length)`):
            Padding mask used to pad `input_values`.
        input_features_mask (`torch.Tensor` of shape `(batch_size, time_steps)`, *optional*):
            Attention mask for the spectrogram input to the semantic encoder. `1` for valid frames, `0` for padding.
        output_latents (`bool`, *optional*, defaults to `False`):
            Whether to return the continuous latent representation from the quantizer.

        Examples:

        ```python
        >>> from datasets import load_dataset
        >>> from transformers import AutoFeatureExtractor, Xcodec2Model

        >>> dataset = load_dataset("hf-internal-testing/librispeech_asr_dummy", "clean", split="validation")
        >>> audio = dataset["train"]["audio"][0]["array"]

        >>> model_id = "HKUSTAudio/xcodec2-hf"
        >>> model = Xcodec2Model.from_pretrained(model_id)
        >>> feature_extractor = AutoFeatureExtractor.from_pretrained(model_id)

        >>> inputs = feature_extractor(audio=audio, sampling_rate=feature_extractor.sampling_rate, return_tensors="pt")

        >>> outputs = model(**inputs)
        >>> audio_codes = outputs.audio_codes
        >>> audio_values = outputs.audio_values
        ```r^   T)r  r  r  r  Úreturn_dict)r!   r  r   .N)r   r    r!   r"   r,   )rf   r  r  r!   r   r    r"   )
rG   rå  r  r  r  r  rª   ÚlengthÚencoder_outputsr   s
             r.   rt   zXcodec2Model.forward#  s¯   € ðV Ô# BÔ'ˆàŸ+š+ØØ)Ø%Ø 3ØØð &ñ 
ô 
ˆð #�t”{Ð_¨?Ô+BÐPTÐ_Ð_ÐX^Ð_Ð_Ð`aÔbÐcfÐhoÐioÐhoÐcoÔpˆåØ%Ø'Ô3Ø/=ÐG�OÔ+Ð+À4Ø,Ô=ð	
ñ 
ô 
ð 	
r-   )NNF)NN)r#   r$   r%   r   Úconfig_classr>   r   r   r'   r+   rñ   r   r   ry   r0   r  r2   r  r   rt   r{   r|   s   @r.   rü  rü  À  sû  ø€ € € € € à €Lð˜}ð ð ð ð ð ð ð  Øð
 -1Ø37Ø$ð2
ð 2
à”lð2
ð œð2
ð ”l TÑ)ð	2
ð
 #œ\¨DÑ0ð2
ð ð2
ð Ð+Ô,ð2
ð 
Ð%Ñ	%ð2
ð 2
ð 2
ñ Ôñ „^ð2
ðh Øð ,0Ø'+ð>ð >à”\ DÑ(ð>ð ” Ñ$ð>ð Ð+Ô,ð	>ð
 
Ð%Ñ	%ð>ð >ð >ñ Ôñ „^ð>ð. Øð
 -1Ø37Ø$ð:
ð :
à”lð:
ð œð:
ð ”l TÑ)ð	:
ð
 #œ\¨DÑ0ð:
ð ð:
ð Ð+Ô,ð:
ð 
�Ñ	ð:
ð :
ð :
ñ Ôñ „^ð:
ð :
ð :
ð :
ð :
r-   rü  )r   )r¢   )Qr  Úcollections.abcr   Údataclassesr   Útypingr   Únumpyr„  r'   Útorch.nnr„   Útorch.nn.functionalr²   r#  r   Ú r   ré  Úactivationsr	   Úcache_utilsr
   Úintegrationsr   r   r   Úmodeling_layersr   Úmodeling_rope_utilsr   r   Úmodeling_utilsr   r   Úprocessing_utilsr   Úutilsr   r   r   r   Úutils.genericr   Úautor   Úconfiguration_xcodec2r   r   r0   r2   ÚModuler4   r~   r‘   rš   r+   rx   r¡   rY   rº   r¼   rÔ   rå   ró   r  r  r+  r1  r@  rP  r\  rg  rw  r›  r¹  rÅ  rØ  rã  rü  Ú__all__r,   r-   r.   ú<module>r.     s^  ðð( €€€Ø $Ð $Ð $Ð $Ð $Ð $Ø !Ð !Ð !Ð !Ð !Ð !Ø Ð Ð Ð Ð Ð à Ð Ð Ð Ø €€€Ø Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð à &Ð &Ð &Ð &Ð &Ð &Ø !Ð !Ð !Ð !Ð !Ð !Ø  Ð  Ð  Ð  Ð  Ð  Ø fÐ fÐ fÐ fÐ fÐ fÐ fÐ fÐ fÐ fØ 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø KÐ KÐ KÐ KÐ KÐ KÐ KÐ KØ FÐ FÐ FÐ FÐ FÐ FÐ FÐ FØ &Ð &Ð &Ð &Ð &Ð &Ø VÐ VÐ VÐ VÐ VÐ VÐ VÐ VÐ VÐ VÐ VÐ VØ +Ð +Ð +Ð +Ð +Ð +Ø Ð Ð Ð Ð Ð Ø 0Ð 0Ð 0Ð 0Ð 0Ð 0ð Ø
ð1ð 1ð 1ð 1ð 1�Kñ 1ô 1ñ „ñ „ð1ð( Ø
ð1ð 1ð 1ð 1ð 1˜;ñ 1ô 1ñ „ñ „ð1ð" Ø
ð2ð 2ð 2ð 2ð 2˜;ñ 2ô 2ñ „ñ „ð2ð><ð ><ð ><ð ><ð ><˜RœYñ ><ô ><ð ><ðBð ð ð ð �”ñ ô ð ð(ð (ð (ð ÐÐ*Ñ+Ô+ðð ð ñ ,Ô+ðð2	U˜Uœ\ð 	U°#ð 	U¸%¼,ð 	Uð 	Uð 	Uð 	Uð& ð%ð %ØŒIð%àŒ<ð%ð 
Œð%ð Œ<ð	%ð
 ”L 4Ñ'ð%ð ð%ð ð%ð Ð'Ô(ð%ð %ð %ð %ð2 ÐÐ)Ñ*Ô*ðC)ð C)ð C)ð C)ð C)�r”yñ C)ô C)ñ +Ô*ðC)ðL Ð˜YÑ'Ô'ðJð Jð Jð Jð J�R”Yñ Jô Jñ (Ô'ðJð((ð (ð (ð (ð (Ð4ñ (ô (ð (ðV&ð &ð &ð &ð &�r”yñ &ô &ð &ðR+5ð +5ð +5ð\ð ð ð ð ˜"œ)ñ ô ð ðDð ð ð ð ˜œ	ñ ô ð ð6ð ð ð ð  R¤Yñ ô ð ð0!ð !ð !ð !ð !˜"œ)ñ !ô !ð !ðHð ð ð ð ˜"œ)ñ ô ð ð.ð ð ð ð �R”Yñ ô ð ð@:ð :ð :ð :ð :˜œñ :ô :ð :ð0P1ð P1ð P1ð P1ð P1 b¤iñ P1ô P1ð P1ðf2"ð 2"ð 2"ð 2"ð 2"�r”yñ 2"ô 2"ð 2"ðj&ð &ð &ð &ð &�r”yñ &ô &ð &ð,)3ð )3ð )3ð )3ð )3�R”Yñ )3ô )3ð )3ðX+ð +ð +ð +ð +˜RœYñ +ô +ð +ð\ ð%5ð %5ð %5ð %5ð %5˜_ñ %5ô %5ñ „ð%5ðP €Ð@ÐAÑAÔAð^
ð ^
ð ^
ð ^
ð ^
Ð)ñ ^
ô ^
ñ BÔAð^
ðB Ð3Ð
4€€€r-   