Ë
    S^(hëv  ã                   óŽ  — d Z ddl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mZmZmZ dd	lmZ d
Ze G d„ de«      «       Ze G d„ de«      «       Ze G d„ de«      «       Z G d„ de
j6                  «      Z G d„ de
j6                  «      Z G d„ de
j6                  «      Z G d„ de
j6                  «      Z G d„ de
j6                  «      Z  G d„ de
j6                  «      Z! G d„ de
j6                  «      Z" G d„ d e
j6                  «      Z# G d!„ d"e«      Z$d#Z%d$Z& ed%e%«       G d&„ d'e$«      «       Z'd'd"gZ(y)(zTransformers DAC model.é    N)Ú	dataclass)ÚOptionalé   )ÚPreTrainedModel)ÚModelOutputÚadd_start_docstringsÚ%add_start_docstrings_to_model_forwardÚreplace_return_docstringsé   )Ú	DacConfigr   c                   óÚ   — e Zd ZU dZdZeej                     ed<   dZ	eej                     ed<   dZ
eej                     ed<   dZeej                     ed<   dZeej                     ed<   y)Ú	DacOutputa.  
    Args:
        loss (`torch.Tensor`):
            Loss from the encoder model, comprising the weighted combination of the commitment and codebook losses.
        audio_values (`torch.Tensor` of shape `(batch_size, input_length)`):
            Reconstructed audio data.
        quantized_representation (`torch.Tensor` of shape `(batch_size, dimension, time_steps)`):
            Quantized continuous representation of input.
        audio_codes (`torch.LongTensor` of shape `(batch_size, num_codebooks, time_steps)`):
            Codebook indices for each codebook (quantized discrete representation of input).
        projected_latents (`torch.Tensor` of shape `(batch_size, num_codebooks * dimension, time_steps)`):
            Projected latents (continuous representation of input before quantization).
    NÚlossÚaudio_valuesÚquantized_representationÚaudio_codesÚprojected_latents)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   ÚtorchÚFloatTensorÚ__annotations__r   r   r   Ú
LongTensorr   © ó    úb/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/dac/modeling_dac.pyr   r   (   st   … ñð )-€Dˆ(�5×$Ñ$Ñ
%Ó,Ø04€L�(˜5×,Ñ,Ñ-Ó4Ø<@Ð˜h u×'8Ñ'8Ñ9Ó@Ø.2€K�˜%×*Ñ*Ñ+Ó2Ø59Ð�x × 1Ñ 1Ñ2Ô9r   r   c                   ó²   — e Zd ZU dZdZeej                     ed<   dZ	eej                     ed<   dZ
eej                     ed<   dZeej                     ed<   y)ÚDacEncoderOutputaÛ  
    Args:
        loss (`torch.Tensor`):
            Loss from the encoder model, comprising the weighted combination of the commitment and codebook losses.
        quantized_representation (`torch.Tensor` of shape `(batch_size, dimension, time_steps)`, *optional*):
            Quantized continuous representation of input.
        audio_codes (`torch.Tensor` of shape `(batch_size, num_codebooks, time_steps)`, *optional*):
            Codebook indices for each codebook (quantized discrete representation of input).
        projected_latents (`torch.Tensor` of shape `(batch_size, num_codebooks * dimension, time_steps)`, *optional*):
            Projected latents (continuous representation of input before quantization).
    Nr   r   r   r   )r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r    r    ?   s_   … ñ
ð )-€Dˆ(�5×$Ñ$Ñ
%Ó,Ø<@Ð˜h u×'8Ñ'8Ñ9Ó@Ø/3€K�˜%×+Ñ+Ñ,Ó3Ø59Ð�x × 1Ñ 1Ñ2Ô9r   r    c                   ó:   — e Zd ZU dZdZeej                     ed<   y)ÚDacDecoderOutputz¸
    Args:
        audio_values (`torch.FloatTensor`  of shape `(batch_size, input_length)`, *optional*):
            Decoded audio values, obtained using the decoder part of Dac.
    Nr   )	r   r   r   r   r   r   r   r   r   r   r   r   r"   r"   S   s   … ñð 15€L�(˜5×,Ñ,Ñ-Ô4r   r"   c                   ó(   ‡ — e Zd ZdZˆ fd„Zd„ Zˆ xZS )ÚSnake1dz;
    A 1-dimensional Snake activation function module.
    c                 ó€   •— t         ‰| �  «        t        j                  t	        j
                  d|d«      «      | _        y )Nr   )ÚsuperÚ__init__ÚnnÚ	Parameterr   ÚonesÚalpha)ÚselfÚ
hidden_dimÚ	__class__s     €r   r'   zSnake1d.__init__d   s+   ø€ Ü‰ÑÔÜ—\‘\¤%§*¡*¨Q°
¸AÓ">Ó?ˆ�
r   c                 ó  — |j                   }|j                  |d   |d   d«      }|| j                  dz   j                  «       t	        j
                  | j                  |z  «      j                  d«      z  z   }|j                  |«      }|S )Nr   r   éÿÿÿÿg•Ö&è.>é   )ÚshapeÚreshaper+   Ú
reciprocalr   ÚsinÚpow)r,   Úhidden_statesr2   s      r   ÚforwardzSnake1d.forwardh   s‚   € Ø×#Ñ#ˆØ%×-Ñ-¨e°A©h¸¸a¹À"ÓEˆØ%¨¯©°dÑ):×(FÑ(FÓ(HÌ5Ï9É9ÐUY×U_ÑU_ÐboÑUoÓKp×KtÑKtÐuvÓKwÑ(wÑwˆØ%×-Ñ-¨eÓ4ˆØÐr   )r   r   r   r   r'   r8   Ú__classcell__©r.   s   @r   r$   r$   _   s   ø„ ñô@ör   r$   c                   ó4   ‡ — e Zd ZdZdefˆ fd„Zd„ Zd„ Zˆ xZS )ÚDacVectorQuantizeaÕ  
    Implementation of VQ similar to Karpathy's repo (https://github.com/karpathy/deep-vector-quantization)

    Additionally uses following tricks from improved VQGAN
    (https://arxiv.org/pdf/2110.04627.pdf):
        1. Factorized codes: Perform nearest neighbor lookup in low-dimensional space
            for improved codebook usage
        2. l2-normalized codes: Converts euclidean distance to cosine similarity which
            improves training stability
    Úconfigc                 óD  •— t         ‰| �  «        t        j                  |j                  |j
                  d¬«      | _        t        j                  |j
                  |j                  d¬«      | _        t        j                  |j                  |j
                  «      | _
        y )Nr   ©Úkernel_size)r&   r'   r(   ÚConv1dÚhidden_sizeÚcodebook_dimÚin_projÚout_projÚ	EmbeddingÚcodebook_sizeÚcodebook©r,   r=   r.   s     €r   r'   zDacVectorQuantize.__init__|   sn   ø€ Ü‰ÑÔä—y‘y ×!3Ñ!3°V×5HÑ5HÐVWÔXˆŒÜŸ	™	 &×"5Ñ"5°v×7IÑ7IÐWXÔYˆŒÜŸ™ V×%9Ñ%9¸6×;NÑ;NÓOˆ�r   c                 ó@  — | j                  |«      }| j                  |«      \  }}t        j                  ||j	                  «       d¬«      }t        j                  ||j	                  «       d¬«      }|||z
  j	                  «       z   }| j                  |«      }|||||fS )aJ  
        Quantizes the input tensor using a fixed codebook and returns the corresponding codebook vectors.

        Args:
            hidden_state (`torch.FloatTensor` of shape `(batch_size, dimension, time_steps)`):
                Input tensor.

        Returns:
            quantized_representation (`torch.Tensor`of shape `(batch_size, dimension, time_steps)`):
                Quantized continuous representation of input.
            commitment_loss (`torch.FloatTensor`of shape `(1)`):
                Commitment loss to train encoder to predict vectors closer to codebook entries.
            codebook_loss (`torch.FloatTensor`of shape `(1)`):
                Codebook loss to update the codebook.
            audio_codes (`torch.LongTensor` of shape `(batch_size, time_steps)`):
                Codebook indices for each codebook, quantized discrete representation of input.
            projected_latents (torch.FloatTensor of shape `(batch_size, num_codebooks * dimension, time_steps)`):
                Projected latents (continuous representation of input before quantization).
        Úmean)Ú	reduction)rD   Údecode_latentsÚFÚmse_lossÚdetachrE   )r,   Úhidden_stater   r   r   Úcommitment_lossÚcodebook_losss          r   r8   zDacVectorQuantize.forwardƒ   s£   € ð* !ŸL™L¨Ó6ÐØ04×0CÑ0CÐDUÓ0VÑ-Ð  +äŸ*™*Ð%6Ð8P×8WÑ8WÓ8YÐekÔlˆÜŸ
™
Ð#;Ð=N×=UÑ=UÓ=WÐciÔjˆà#4Ð8PÐSdÑ8d×7lÑ7lÓ7nÑ#nÐ Ø#'§=¡=Ð1IÓ#JÐ à'¨¸-ÈÐVgÐgÐgr   c                 ó|  — |j                   \  }}}|j                  ddd«      j                  ||z  |«      }| j                  j                  }t        j                  |«      }t        j                  |«      }|j                  d«      j                  dd¬«      }|d|z  |j                  «       z  z
   |j                  d«      j                  dd¬«      j                  «       z   }|j                  d«      d   }	|	j                  |j                  d«      d«      }	| j                  |	«      j                  dd«      }
|
|	fS )Nr   r1   r   T)Úkeepdimr0   )r2   Úpermuter3   rH   ÚweightrN   Ú	normalizer6   ÚsumÚtÚmaxÚsizeÚ	transpose)r,   r7   Ú
batch_sizer-   Úsequence_lengthÚ	encodingsrH   Úl2_normÚdistÚindicesr   s              r   rM   z DacVectorQuantize.decode_latents£   s  € Ø2?×2EÑ2EÑ/ˆ
�J Ø!×)Ñ)¨!¨Q°Ó2×:Ñ:¸:ÈÑ;WÐYcÓdˆ	Ø—=‘=×'Ñ'ˆô —K‘K 	Ó*ˆ	Ü—;‘;˜xÓ(ˆð —-‘- Ó"×&Ñ& q°$Ð&Ó7ˆØ˜1˜y™=¨8¯:©:«<Ñ7Ñ7Ð8¸8¿<¹<È»?×;NÑ;NÈqÐZ^Ð;NÓ;_×;aÑ;aÓ;cÑcˆà—(‘(˜1“+˜a‘.ˆØ—/‘/ -×"4Ñ"4°QÓ"7¸Ó<ˆØ#'§=¡=°Ó#9×#CÑ#CÀAÀqÓ#IÐ Ø'¨Ð0Ð0r   )	r   r   r   r   r   r'   r8   rM   r9   r:   s   @r   r<   r<   p   s"   ø„ ñ	ðP˜yõ Pòhö@1r   r<   c                   ó4   ‡ — e Zd ZdZddedefˆ fd„Zd„ Zˆ xZS )ÚDacResidualUnitza
    A residual unit composed of Snake1d and weight-normalized Conv1d layers with dilations.
    Ú	dimensionÚdilationc                 óê   •— t         ‰| �  «        d|z  dz  }t        |«      | _        t	        j
                  ||d||¬«      | _        t        |«      | _        t	        j
                  ||d¬«      | _        y )Né   r1   é   )r@   rg   Úpaddingr   r?   )	r&   r'   r$   Úsnake1r(   rA   Úconv1Úsnake2Úconv2)r,   rf   rg   Úpadr.   s       €r   r'   zDacResidualUnit.__init__»   sb   ø€ Ü‰ÑÔØ˜Ñ! aÑ'ˆä˜iÓ(ˆŒÜ—Y‘Y˜y¨)ÀÈXÐ_bÔcˆŒ
Ü˜iÓ(ˆŒÜ—Y‘Y˜y¨)ÀÔCˆ�
r   c                 óö   — |}| j                  | j                  |«      «      }| j                  | j                  |«      «      }|j                  d   |j                  d   z
  dz  }|dkD  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.
        r0   r1   r   .)rm   rl   ro   rn   r2   )r,   rQ   Úoutput_tensorrk   s       r   r8   zDacResidualUnit.forwardÄ   s‰   € ð %ˆØŸ
™
 4§;¡;¨}Ó#=Ó>ˆØŸ
™
 4§;¡;¨}Ó#=Ó>ˆà×%Ñ% bÑ)¨M×,?Ñ,?ÀÑ,CÑCÈÑIˆØ�QŠ;Ø'¨¨W°g°XÐ-=Ð(=Ñ>ˆLØ$ }Ñ4ˆØÐr   )é   r   )r   r   r   r   Úintr'   r8   r9   r:   s   @r   re   re   ¶   s#   ø„ ññD #ð D°cõ Dör   re   c                   ó8   ‡ — e Zd ZdZddededefˆ fd„Zd„ Zˆ xZS )ÚDacEncoderBlockz"Encoder block used in DAC encoder.r=   ÚstrideÚstride_indexc           
      ó`  •— t         ‰| �  «        |j                  d|z  z  }t        |dz  d¬«      | _        t        |dz  d¬«      | _        t        |dz  d¬«      | _        t        |dz  «      | _        t        j                  |dz  |d|z  |t        j                  |dz  «      ¬«      | _        y )Nr1   r   ©rg   r   é	   ©r@   rw   rk   )r&   r'   Úencoder_hidden_sizere   Ú	res_unit1Ú	res_unit2Ú	res_unit3r$   rl   r(   rA   ÚmathÚceilrm   )r,   r=   rw   rx   rf   r.   s        €r   r'   zDacEncoderBlock.__init__Þ   sž   ø€ Ü‰ÑÔà×.Ñ.°°L±Ñ@ˆ	Ü(¨°a©À!ÔDˆŒÜ(¨°a©À!ÔDˆŒÜ(¨°a©À!ÔDˆŒÜ˜i¨1™nÓ-ˆŒÜ—Y‘YØ˜‰N˜I°1°v±:ÀfÔVZ×V_ÑV_Ð`fÐijÑ`jÓVkô
ˆ�
r   c                 ó¬   — | j                  |«      }| j                  |«      }| j                  | j                  |«      «      }| j	                  |«      }|S ©N)r~   r   rl   r€   rm   ©r,   rQ   s     r   r8   zDacEncoderBlock.forwardê   sI   € Ø—~‘~ lÓ3ˆØ—~‘~ lÓ3ˆØ—{‘{ 4§>¡>°,Ó#?Ó@ˆØ—z‘z ,Ó/ˆàÐr   ©r   r   ©	r   r   r   r   r   rt   r'   r8   r9   r:   s   @r   rv   rv   Û   s%   ø„ Ù,ñ

˜yð 

°#ð 

Èõ 

ör   rv   c                   ó8   ‡ — e Zd ZdZddededefˆ fd„Zd„ Zˆ xZS )ÚDacDecoderBlockz"Decoder block used in DAC decoder.r=   rw   rx   c           
      ól  •— t         ‰| �  «        |j                  d|z  z  }|j                  d|dz   z  z  }t        |«      | _        t        j                  ||d|z  |t        j                  |dz  «      ¬«      | _	        t        |d¬«      | _        t        |d¬«      | _        t        |d¬«      | _        y )Nr1   r   r|   rz   r   r{   )r&   r'   Údecoder_hidden_sizer$   rl   r(   ÚConvTranspose1dr�   r‚   Úconv_t1re   r~   r   r€   )r,   r=   rw   rx   Ú	input_dimÚ
output_dimr.   s         €r   r'   zDacDecoderBlock.__init__ö   s¦   ø€ Ü‰ÑÔà×.Ñ.°!°\±/ÑAˆ	Ø×/Ñ/°1¸ÈÑ9IÑ3JÑJˆ
Ü˜iÓ(ˆŒÜ×)Ñ)ØØØ˜F™
ØÜ—I‘I˜f q™jÓ)ô
ˆŒô )¨¸aÔ@ˆŒÜ(¨¸aÔ@ˆŒÜ(¨¸aÔ@ˆ�r   c                 ó°   — | j                  |«      }| j                  |«      }| j                  |«      }| j                  |«      }| j	                  |«      }|S r„   )rl   r�   r~   r   r€   r…   s     r   r8   zDacDecoderBlock.forward  sN   € Ø—{‘{ <Ó0ˆØ—|‘| LÓ1ˆØ—~‘~ lÓ3ˆØ—~‘~ lÓ3ˆØ—~‘~ lÓ3ˆàÐr   r†   r‡   r:   s   @r   r‰   r‰   ó   s)   ø„ Ù,ñA˜yð A°#ð AÈõ Aö$r   r‰   c                   ó|   ‡ — e Zd ZdZdefˆ fd„Zd
dee   fd„Zde	j                  fd„Zde	j                  fd	„Zˆ xZS )ÚDacResidualVectorQuantizez„
    ResidualVectorQuantize block - Introduced in SoundStream: An end2end neural audio codec (https://arxiv.org/abs/2107.03312)
    r=   c                 ó   •— t         ‰| �  «        |j                  }|j                  }|| _        t	        j
                  t        |j                  «      D �cg c]  }t        |«      ‘Œ c}«      | _        || _        y c c}w r„   )	r&   r'   Ún_codebooksÚquantizer_dropoutr(   Ú
ModuleListÚranger<   Ú
quantizers)r,   r=   r”   r•   Úir.   s        €r   r'   z"DacResidualVectorQuantize.__init__  sh   ø€ Ü‰ÑÔà×(Ñ(ˆØ"×4Ñ4Ðà&ˆÔäŸ-™-ÌEÐRX×RdÑRdÓLeÖ(fÀqÔ):¸6Õ)BÒ(fÓgˆŒØ!2ˆÕùò )gs   ÁA;Ún_quantizersc                 óŠ  — d}|}d}d}g }g }|�|n| j                   }| j                  r­t        j                  |j                  d   f«      | j                   z  dz   }t        j
                  d| j                   dz   |j                  d   f«      }	t        |j                  d   | j                  z  «      }
|	d|
 |d|
 |j                  |j                  «      }t        | j                  «      D ]¢  \  }}| j                  du r||k\  r nŠ ||«      \  }}}}}t        j                  |j                  d   f||j                  ¬«      |k  }|||dd…ddf   z  z   }||z
  }|||z  z  }|||z  z  }|j                  |«       |j                  |«       Œ¤ t        j                  |d¬«      }t        j                  |d¬«      }|||||fS )aQ  
        Quantizes the input tensor using a fixed set of codebooks and returns corresponding codebook vectors.
        Args:
            hidden_state (`torch.Tensor` of shape `(batch_size, dimension, time_steps)`):
                Input tensor to be quantized.
            n_quantizers (`int`, *optional*):
                Number of quantizers to use. If specified and `self.quantizer_dropout` is True,
                this argument is ignored during training, and a random number of quantizers is used.

        Returns:
            quantized_representation (`torch.Tensor` of shape `(batch_size, dimension, time_steps)`):
                Quantized continuous representation of input.
            audio_codes (`torch.Tensor` of shape `(batch_size, num_codebooks, time_steps)`):
                Codebook indices for each codebook (quantized discrete representation of input).
            projected_latents (`torch.Tensor` of shape `(batch_size, num_codebooks * dimension, time_steps)`):
                Projected latents (continuous representation of input before quantization).
            commitment_loss (`torch.Tensor` of shape `(1)`):
                Commitment loss to train the encoder to predict vectors closer to codebook entries.
            codebook_loss (`torch.Tensor` of shape `(1)`):
                Codebook loss to update the codebook.
        r   Nr   F)Ú
fill_valueÚdevice©Údim)r”   Útrainingr   r*   r2   Úrandintrt   r•   Útor�   Ú	enumerater˜   ÚfullÚappendÚstackÚcat)r,   rQ   rš   r   ÚresidualrR   rS   r   r   ÚdropoutÚ	n_dropoutr™   Ú	quantizerÚquantized_representation_iÚcommitment_loss_iÚcodebook_loss_iÚ	indices_iÚprojected_latents_iÚmasks                      r   r8   z!DacResidualVectorQuantize.forward"  sú  € ð. $%Ð ØˆØˆØˆàˆØÐà'3Ð'?‘|ÀT×EUÑEUˆØ�=Š=Ü Ÿ:™: |×'9Ñ'9¸!Ñ'<Ð&>Ó?À$×BRÑBRÑRÐUVÑVˆLÜ—m‘m A t×'7Ñ'7¸!Ñ';¸l×>PÑ>PÐQRÑ>SÐ=UÓVˆGÜ˜L×.Ñ.¨qÑ1°D×4JÑ4JÑJÓKˆIØ'.¨z°	Ð':ˆL˜˜)Ð$Ø'Ÿ?™?¨<×+>Ñ+>Ó?ˆLä% d§o¡oÓ6ò 	:‰LˆAˆyØ�}‰} Ñ%¨!¨|Ò*;ÙámvØónÑjÐ&Ð(9¸?ÈIÐWjô
 —:‘:˜|×1Ñ1°!Ñ4Ð6À1È\×M`ÑM`ÔaÐdpÑpˆDØ'?ÐB\Ð_cÒdeÐgkÐmqÐdqÑ_rÑBrÑ'rÐ$ØÐ"<Ñ<ˆHð Ð0°4Ñ7Ñ7ˆOØ˜_¨tÑ3Ñ3ˆMà×Ñ˜yÔ)Ø×$Ñ$Ð%8Õ9ð%	:ô( —k‘k +°1Ô5ˆÜ!ŸI™IÐ&7¸QÔ?Ðà'¨Ð6GÈÐZgÐgÐgr   r   c                 óP  — d}g }|j                   d   }t        |«      D ]l  }| j                  |   j                  |dd…|dd…f   «      j	                  dd«      }|j                  |«       || j                  |   j                  |«      z  }Œn |t        j                  |d¬«      |fS )a–  
        Reconstructs the continuous representation from quantized codes.

        Args:
            audio_codes (`torch.Tensor` of shape `(batch_size, num_codebooks, time_steps)`):
                Quantized discrete representation of input.

        Returns:
            quantized_representation (`torch.Tensor`):
                Quantized continuous representation of input.
            projected_latents (`torch.Tensor`):
                List of projected latents (continuous representations of input before quantization)
                for each codebook.
            audio_codes (`torch.Tensor`):
                Codebook indices for each codebook.
        g        r   Nr1   rž   )	r2   r—   r˜   rH   r]   r¥   rE   r   r§   )r,   r   r   r   r”   r™   r°   s          r   Ú
from_codesz$DacResidualVectorQuantize.from_codesb  sµ   € ð" $'Ð ØÐØ!×'Ñ'¨Ñ*ˆÜ�{Ó#ò 	YˆAØ"&§/¡/°!Ñ"4×"=Ñ"=¸kÊ!ÈQÒPQÈ'Ñ>RÓ"S×"]Ñ"]Ð^_ÐabÓ"cÐØ×$Ñ$Ð%8Ô9Ø$¨¯©¸Ñ(:×(CÑ(CÐDWÓ(XÑXÑ$ð	Yð (¬¯©Ð3DÈ!Ô)LÈkÐYÐYr   Úlatentsc                 ó„  — d}g }g }t        j                  dg| j                  D �cg c]  }|j                  ‘Œ c}z   «      }t        j                  |d¬«      }t        j                  ||j                  d   k  «      d   j                  dd¬«      d   }t        |«      D ]�  }	||	   ||	dz      }}
| j                  |	   j                  |dd…|
|…dd…f   «      \  }}|j                  |«       |j                  |«       | j                  |	   j                  |«      }||z   }Œƒ |t        j                  |d¬«      fS c c}w )a�  Reconstructs the quantized representation from unquantized latents.

        Args:
            latents (`torch.Tensor` of shape `(batch_size, total_latent_dimension, time_steps)`):
                Continuous representation of input after projection.

        Returns:
            quantized_representation (`torch.Tensor` of shape `(batch_size, dimension, time_steps)`):
                Quantized representation of the full-projected space.
            quantized_latents (`torch.Tensor` of shape `(batch_size, dimension, time_steps)`):
                Quantized representation of the latent space (continuous representation before quantization).
        r   rž   r   T)ÚaxisÚkeepdimsN)r   Útensorr˜   rC   ÚcumsumÚnpÚwherer2   r[   r—   rM   r¥   rE   r§   )r,   r´   r   Úquantized_latentsÚcodesÚqÚcodebook_dims_tensorÚdimsr”   r™   Úhidden_dim_jÚhidden_dim_kÚquantized_latents_iÚcodes_ir¬   s                  r   Úfrom_latentsz&DacResidualVectorQuantize.from_latents|  sG  € ð $%Ð ØÐØˆÜ$Ÿ|™|¨Q¨CÈ4Ï?É?Ö2[Àa°1·>³>Ò2[Ñ,[Ó\ÐÜ�|‰|Ð0°aÔ8ˆä—h‘h˜t w§}¡}°QÑ'7Ñ7Ó8¸Ñ;×?Ñ?ÀQÐQUÐ?ÓVÐWXÑYˆÜ�{Ó#ò 	]ˆAØ)-¨a©°$°q¸1±u±+˜,ˆLØ+/¯?©?¸1Ñ+=×+LÑ+LÈWÒUVÐXdÐeqÐXqÒstÐUtÑMuÓ+vÑ(Ð Ø×$Ñ$Ð%8Ô9Ø�L‰L˜Ô!à)-¯©¸Ñ);×)DÑ)DÐEXÓ)YÐ&Ø'?ÐB\Ñ'\Ñ$ð	]ð (¬¯©Ð3DÈ!Ô)LÐLÐLùò 3\s   ¦D=
r„   )r   r   r   r   r   r'   r   rt   r8   r   ÚTensorr³   rÅ   r9   r:   s   @r   r’   r’     sK   ø„ ñð	3˜yõ 	3ñ>h°(¸3±-ó >hð@Z e§l¡ló Zð4M E§L¡L÷ Mr   r’   c                   ó.   ‡ — e Zd ZdZdefˆ fd„Zd„ Zˆ xZS )Ú
DacDecoderzDAC Decoderr=   c                 óÞ  •— t         ‰	| �  «        |j                  }|j                  }|j                  }t        j                  ||dd¬«      | _        g }t        |«      D ]  \  }}|t        |||«      gz  }Œ t        j                  |«      | _        |j                  ddz   z  z  }t        |«      | _        t        j                  |ddd¬«      | _        t        j                  «       | _        y )Nrj   r   ©r@   rk   r1   r   )r&   r'   rB   r‹   Úupsampling_ratiosr(   rA   rm   r£   r‰   r–   Úblockr$   rl   ro   ÚTanhÚtanh)
r,   r=   Úinput_channelÚchannelsÚstridesrÌ   rx   rw   r�   r.   s
            €r   r'   zDacDecoder.__init__Ÿ  sÚ   ø€ Ü‰ÑÔà×*Ñ*ˆØ×-Ñ-ˆØ×*Ñ*ˆô —Y‘Y˜}¨hÀAÈqÔQˆŒ
ð ˆÜ$-¨gÓ$6ò 	EÑ ˆL˜&Ø”o f¨f°lÓCÐDÑD‰Eð	Eô —]‘] 5Ó)ˆŒ
Ø×/Ñ/°1¸ÈÑ9IÑ3JÑJˆ
Ü˜jÓ)ˆŒÜ—Y‘Y˜z¨1¸!ÀQÔGˆŒ
Ü—G‘G“Iˆ�	r   c                 óÀ   — | j                  |«      }| j                  D ]
  } ||«      }Œ | j                  |«      }| j                  |«      }| j	                  |«      }|S r„   )rm   rÌ   rl   ro   rÎ   )r,   rQ   Úlayers      r   r8   zDacDecoder.forward´  s_   € Ø—z‘z ,Ó/ˆà—Z‘Zò 	/ˆEÙ  Ó.‰Lð	/ð —{‘{ <Ó0ˆØ—z‘z ,Ó/ˆØ—y‘y Ó.ˆàÐr   ©r   r   r   r   r   r'   r8   r9   r:   s   @r   rÈ   rÈ   œ  s   ø„ Ùð˜yõ ö*
r   rÈ   c                   ó.   ‡ — e Zd ZdZdefˆ fd„Zd„ Zˆ xZS )Ú
DacEncoderzDAC Encoderr=   c                 óè  •— t         ‰| �  «        |j                  }t        j                  d|j
                  dd¬«      | _        g | _        t        |«      D ],  \  }}|dz   }| xj                  t        |||¬«      gz  c_        Œ. t        j                  | j                  «      | _        |j
                  dz  z  }t        |«      | _        t        j                  ||j                  dd¬«      | _        y )Nr   rj   r   rÊ   )rw   rx   r1   )r&   r'   Údownsampling_ratiosr(   rA   r}   rm   rÌ   r£   rv   r–   r$   rl   rB   ro   )r,   r=   rÑ   rx   rw   Úd_modelr.   s         €r   r'   zDacEncoder.__init__Ä  sÏ   ø€ Ü‰ÑÔà×,Ñ,ˆä—Y‘Y˜q &×"<Ñ"<È!ÐUVÔWˆŒ
àˆŒ
ä$-¨gÓ$6ò 	^Ñ ˆL˜&Ø'¨!Ñ+ˆLØ�JŠJœ?¨6¸&È|Ô\Ð]Ñ]ŽJð	^ô —]‘] 4§:¡:Ó.ˆŒ
Ø×,Ñ,¨q°,©Ñ>ˆÜ˜gÓ&ˆŒÜ—Y‘Y˜w¨×(:Ñ(:ÈÐSTÔUˆ�
r   c                 óž   — | j                  |«      }| j                  D ]
  } ||«      }Œ | j                  |«      }| j                  |«      }|S r„   )rm   rÌ   rl   ro   )r,   rQ   Úmodules      r   r8   zDacEncoder.forwardÖ  sQ   € Ø—z‘z ,Ó/ˆà—j‘jò 	0ˆFÙ! ,Ó/‰Lð	0ð —{‘{ <Ó0ˆØ—z‘z ,Ó/ˆàÐr   rÔ   r:   s   @r   rÖ   rÖ   Á  s   ø„ ÙðV˜yõ Vö$	r   rÖ   c                   ó.   — e Zd ZdZeZdZdZd„ Zd„ Z	d„ Z
y)ÚDacPreTrainedModelz‚
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained models.
    ÚdacÚinput_valuesc                 óä   — t        |t        j                  «      rVt        j                  j	                  |j
                  d¬«       t        j                  j                  |j                  d«       y y )Ng{®Gáz”?)Ústdr   )Ú
isinstancer(   rA   ÚinitÚtrunc_normal_rW   Ú	constant_Úbias)r,   rÛ   s     r   Ú_init_weightsz DacPreTrainedModel._init_weightsë  sH   € Ü�fœbŸi™iÔ(Ü�G‰G×!Ñ! &§-¡-°TÐ!Ô:Ü�G‰G×Ñ˜fŸk™k¨1Õ-ð )r   c                 óz  — t         j                  j                  }t        t         j                  j                  d«      r$t         j                  j                  j                  }| j
                  j                  D ]&  } ||j                  «        ||j                  «       Œ(  || j                  j                  «        || j                  j                  «       | j                  j                  D ]¼  } ||j                  «        ||j                  j                  «        ||j                  j                  «        ||j                  j                  «        ||j                  j                  «        ||j                  j                  «        ||j                  j                  «       Œ¾  || j                   j                  «        || j                   j                  «       | j                   j                  D ]¼  } ||j"                  «        ||j                  j                  «        ||j                  j                  «        ||j                  j                  «        ||j                  j                  «        ||j                  j                  «        ||j                  j                  «       Œ¾ y )NÚweight_norm)r(   Úutilsré   ÚhasattrÚparametrizationsr«   r˜   rD   rE   Úencoderrm   ro   rÌ   r~   r   r€   Údecoderr�   )r,   ré   rÓ   s      r   Úapply_weight_normz$DacPreTrainedModel.apply_weight_normð  sÙ  € Ü—h‘h×*Ñ*ˆÜ”2—8‘8×,Ñ,¨mÔ<ÜŸ(™(×3Ñ3×?Ñ?ˆKà—^‘^×.Ñ.ò 	(ˆEÙ˜Ÿ™Ô&Ù˜Ÿ™Õ'ð	(ñ 	�D—L‘L×&Ñ&Ô'Ù�D—L‘L×&Ñ&Ô'à—\‘\×'Ñ'ò 	/ˆEÙ˜Ÿ™Ô$Ù˜Ÿ™×-Ñ-Ô.Ù˜Ÿ™×-Ñ-Ô.Ù˜Ÿ™×-Ñ-Ô.Ù˜Ÿ™×-Ñ-Ô.Ù˜Ÿ™×-Ñ-Ô.Ù˜Ÿ™×-Ñ-Õ.ð	/ñ 	�D—L‘L×&Ñ&Ô'Ù�D—L‘L×&Ñ&Ô'à—\‘\×'Ñ'ò 	/ˆEÙ˜Ÿ™Ô&Ù˜Ÿ™×-Ñ-Ô.Ù˜Ÿ™×-Ñ-Ô.Ù˜Ÿ™×-Ñ-Ô.Ù˜Ÿ™×-Ñ-Ô.Ù˜Ÿ™×-Ñ-Ô.Ù˜Ÿ™×-Ñ-Õ.ñ	/r   c                 óV  — | j                   j                  D ]T  }t        j                  j	                  |j
                  «       t        j                  j	                  |j                  «       ŒV t        j                  j	                  | j                  j                  «       t        j                  j	                  | j                  j                  «       | j                  j                  D �]^  }t        j                  j	                  |j                  «       t        j                  j	                  |j                  j                  «       t        j                  j	                  |j                  j                  «       t        j                  j	                  |j                  j                  «       t        j                  j	                  |j                  j                  «       t        j                  j	                  |j                  j                  «       t        j                  j	                  |j                  j                  «       �Œa t        j                  j	                  | j                  j                  «       t        j                  j	                  | j                  j                  «       | j                  j                  D �]^  }t        j                  j	                  |j                  «       t        j                  j	                  |j                  j                  «       t        j                  j	                  |j                  j                  «       t        j                  j	                  |j                  j                  «       t        j                  j	                  |j                  j                  «       t        j                  j	                  |j                  j                  «       t        j                  j	                  |j                  j                  «       �Œa y r„   )r«   r˜   r(   rê   Úremove_weight_normrD   rE   rí   rm   ro   rÌ   r~   r   r€   rî   r�   )r,   rÓ   s     r   rñ   z%DacPreTrainedModel.remove_weight_norm  si  € Ø—^‘^×.Ñ.ò 	8ˆEÜ�H‰H×'Ñ'¨¯©Ô6Ü�H‰H×'Ñ'¨¯©Õ7ð	8ô 	�‰×#Ñ# D§L¡L×$6Ñ$6Ô7Ü
�‰×#Ñ# D§L¡L×$6Ñ$6Ô7à—\‘\×'Ñ'ó 	?ˆEÜ�H‰H×'Ñ'¨¯©Ô4Ü�H‰H×'Ñ'¨¯©×(=Ñ(=Ô>Ü�H‰H×'Ñ'¨¯©×(=Ñ(=Ô>Ü�H‰H×'Ñ'¨¯©×(=Ñ(=Ô>Ü�H‰H×'Ñ'¨¯©×(=Ñ(=Ô>Ü�H‰H×'Ñ'¨¯©×(=Ñ(=Ô>Ü�H‰H×'Ñ'¨¯©×(=Ñ(=Ö>ð	?ô 	�‰×#Ñ# D§L¡L×$6Ñ$6Ô7Ü
�‰×#Ñ# D§L¡L×$6Ñ$6Ô7à—\‘\×'Ñ'ó 	?ˆEÜ�H‰H×'Ñ'¨¯©Ô6Ü�H‰H×'Ñ'¨¯©×(=Ñ(=Ô>Ü�H‰H×'Ñ'¨¯©×(=Ñ(=Ô>Ü�H‰H×'Ñ'¨¯©×(=Ñ(=Ô>Ü�H‰H×'Ñ'¨¯©×(=Ñ(=Ô>Ü�H‰H×'Ñ'¨¯©×(=Ñ(=Ô>Ü�H‰H×'Ñ'¨¯©×(=Ñ(=Ö>ñ	?r   N)r   r   r   r   r   Úconfig_classÚbase_model_prefixÚmain_input_namerç   rï   rñ   r   r   r   rÝ   rÝ   â  s)   „ ñð €LØÐØ$€Oò.ò
/óB?r   rÝ   aH  
    This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the
    library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads
    etc.)

    This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.
    Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage
    and behavior.

    Parameters:
        config ([`DacConfig`]):
            Model configuration class with all the parameters of the model. Initializing with a config file does not
            load the weights associated with the model, only the configuration. Check out the
            [`~PreTrainedModel.from_pretrained`] method to load the model weights.
a‡  
    Args:
        input_values (`torch.Tensor` of shape `(batch_size, 1, time_steps)`).
            Audio data to encode,
        n_quantizers (`int`, *optional*):
            Number of quantizers to use. If `None`, all quantizers are used. Default is `None`.
        return_dict (`bool`, *optional*):
            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
z%The DAC (Descript Audio Codec) model.c            
       óR  ‡ — e Zd Zdefˆ fd„Z eee¬«      	 	 ddej                  de
e   de
e   fd„«       Z eee¬«      	 	 	 dde
ej                     d	e
ej                     de
e   fd
„«       Z ee«       eee¬«      	 	 ddej                  de
e   de
e   fd„«       «       Zˆ xZS )ÚDacModelr=   c                 ó‚  •— t         ‰| �  |«       || _        t        |«      | _        t        |«      | _        t        |«      | _        t        t        j                  | j                  j                  «      «      | _        d| j                  z  | j                  j                  k7  rt        d«      ‚| j                  «        y )Nr1   z'The codebook_size must be a power of 2.)r&   r'   r=   rÖ   rí   rÈ   rî   r’   r«   rt   r�   Úlog2rG   Úbits_per_codebookÚ
ValueErrorÚ	post_initrI   s     €r   r'   zDacModel.__init__O  s�   ø€ Ü‰Ñ˜Ô ØˆŒä! &Ó)ˆŒÜ! &Ó)ˆŒä2°6Ó:ˆŒä!$¤T§Y¡Y¨t¯{©{×/HÑ/HÓ%IÓ!JˆÔØˆd×$Ñ$Ñ$¨¯©×(AÑ(AÒAÜÐFÓGÐGð 	�‰Õr   )Úoutput_typerò   rß   rš   Úreturn_dictc                 ó  — |�|n| j                   j                  }| j                  |«      }| j                  ||«      \  }}}}}| j                   j                  |z  | j                   j
                  |z  z   }	|s|	|||fS t        |	|||«      S )aÿ  
        Encode given audio data and return quantized latent codes

        Args:
            input_values (`torch.Tensor of shape `(batch_size, 1, time_steps)`):
                Input audio data to encode,
            n_quantizers (int, *optional*):
                Number of quantizers to use. If None, all quantizers are used. Default is None.
            return_dict (`bool`, *optional*):
                Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
        Returns:

        )r=   rý   rí   r«   Úcommitment_loss_weightÚcodebook_loss_weightr    )
r,   rß   rš   rý   r   r   r   rR   rS   r   s
             r   ÚencodezDacModel.encode_  sŸ   € ð( &1Ð%<‘kÀ$Ç+Á+×BYÑBYˆà#'§<¡<°Ó#=Ð Øcg×cqÑcqØ$ lód
Ñ`Ð  +Ð/@À/ÐS`ð �{‰{×1Ñ1°OÑCÀdÇkÁk×FfÑFfÐivÑFvÑvˆáØÐ2°KÐARÐSÐSä Ð&>ÀÐM^Ó_Ð_r   r   r   c                 óô   — |€|€t        d«      ‚|�|n| j                  j                  }|�| j                  j	                  |«      d   }| j                  |«      j                  d«      }|s|fS t        |«      S )a  Decode given latent codes and return audio data

        Args:
            quantized_representation (torch.Tensor of shape `(batch_size, dimension, time_steps)`, *optional*):
                Quantized continuous representation of input.
            audio_codes (`torch.Tensor` of shape `(batch_size, num_codebooks, time_steps)`, *optional*):
                The codebook indices for each codebook, representing the quantized discrete
                representation of the input. This parameter should be provided if you want
                to decode directly from the audio codes (it will overwrite quantized_representation).
            return_dict (`bool`, *optional*):
                Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.

        Returns:

        zDEither `quantized_representation` or `audio_codes` must be provided.r   r   )rú   r=   rý   r«   r³   rî   Úsqueezer"   )r,   r   r   rý   r   s        r   ÚdecodezDacModel.decode�  s�   € ð. $Ð+°Ð0CÜÐcÓdÐdà%0Ð%<‘kÀ$Ç+Á+×BYÑBYˆàÐ"Ø'+§~¡~×'@Ñ'@ÀÓ'MÈaÑ'PÐ$à—|‘|Ð$<Ó=×EÑEÀaÓHˆáØ �?Ð"ä Ó-Ð-r   c                 óð   — |�|n| j                   j                  }|j                  d   }| j                  ||d¬«      \  }}}}| j	                  |d¬«      d   dd|…f   }	|s||	|||fS t        ||	|||«      S )a¡  
        Returns:
        Examples:

        ```python
        >>> from datasets import load_dataset, Audio
        >>> from transformers import DacModel, AutoProcessor
        >>> librispeech_dummy = load_dataset("hf-internal-testing/librispeech_asr_dummy", "clean", split="validation")

        >>> model = DacModel.from_pretrained("descript/dac_16khz")
        >>> processor = AutoProcessor.from_pretrained("descript/dac_16khz")
        >>> librispeech_dummy = librispeech_dummy.cast_column("audio", Audio(sampling_rate=processor.sampling_rate))
        >>> audio_sample = librispeech_dummy[-1]["audio"]["array"]
        >>> inputs = processor(raw_audio=audio_sample, sampling_rate=processor.sampling_rate, return_tensors="pt")

        >>> encoder_outputs = model.encode(inputs["input_values"])
        >>> # Get the intermediate audio codes
        >>> audio_codes = encoder_outputs.audio_codes
        >>> # Reconstruct the audio from its quantized representation
        >>> audio_values = model.decode(encoder_outputs.quantized_representation)
        >>> # or the equivalent with a forward pass
        >>> audio_values = model(inputs["input_values"]).audio_values
        ```Nr0   F)rý   r   .)r=   rý   r2   r  r  r   )
r,   rß   rš   rý   Úlengthr   r   r   r   r   s
             r   r8   zDacModel.forward§  s«   € ð@ &1Ð%<‘kÀ$Ç+Á+×BYÑBYˆØ×#Ñ# BÑ'ˆØIMÏÉØ˜,°Eð JUó J
ÑFˆÐ&¨Ð5Fð —{‘{Ð#;È�{ÓOÐPQÑRÐSVÐX_ÐY_ÐX_ÐS_Ñ`ˆáØ˜,Ð(@À+ÐO`ÐaÐaä˜˜|Ð-EÀ{ÐTeÓfÐfr   )NN)NNN)r   r   r   r   r'   r
   r    Ú_CONFIG_FOR_DOCr   rÆ   r   rt   Úboolr  r"   r  r	   ÚDAC_INPUTS_DOCSTRINGr   r8   r9   r:   s   @r   rö   rö   J  s  ø„ ð
˜yõ ñ  Ð+;È/ÔZð '+Ø&*ñ	`à—l‘lð`ð ˜s‘mð`ð ˜d‘^ò	`ó [ð`ñB Ð+;È/ÔZð <@Ø.2Ø&*ñ	#.à"*¨5¯<©<Ñ"8ð#.ð ˜eŸl™lÑ+ð#.ð ˜d‘^ò	#.ó [ð#.ñJ +Ð+?Ó@Ù¨9À?ÔSð '+Ø&*ñ	(gà—l‘lð(gð ˜s‘mð(gð ˜d‘^ò	(gó Tó Aô(gr   rö   ))r   r�   Údataclassesr   Útypingr   Únumpyrº   r   Útorch.nnr(   Útorch.nn.functionalÚ
functionalrN   Úmodeling_utilsr   rê   r   r   r	   r
   Úconfiguration_dacr   r  r   r    r"   ÚModuler$   r<   re   rv   r‰   r’   rÈ   rÖ   rÝ   ÚDAC_START_DOCSTRINGr	  rö   Ú__all__r   r   r   ú<module>r     sq  ðñ ã Ý !Ý ã Û Ý ß Ð å -÷ó õ )ð €ð ô:�ó :ó ð:ð, ô:�{ó :ó ð:ð& ô5�{ó 5ó ð5ôˆb�i‰iô ô"C1˜Ÿ	™	ô C1ôL"�b—i‘iô "ôJ�b—i‘iô ô0�b—i‘iô ô>GM §	¡	ô GMôT"�—‘ô "ôJ�—‘ô ôBJ?˜ô J?ðZÐ ð Ð ñ Ø+ØóôCgÐ!ó Cgó	ðCgðL Ð+Ð
,�r   