§
    ‚ŠtjE
 ã                   ó*  — d Z 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m	Z	 ddl
mZ dd	lmZ dd
lmZmZ ddlmZ ddl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(  e$j)        e*¦  «        Z+e#e G d„ de!¦  «        ¦   «         ¦   «         Z, G d„ d¦  «        Z-e#e G d„ de!¦  «        ¦   «         ¦   «         Z.e#e G d„ de!¦  «        ¦   «         ¦   «         Z/ G d„ de	j0        ¦  «        Z1 G d„ d e	j0        ¦  «        Z2 G d!„ d"e	j0        ¦  «        Z3 G d#„ d$e	j0        ¦  «        Z4 G d%„ d&e	j0        ¦  «        Z5 G d'„ d(e	j0        ¦  «        Z6d)„ Z7dQd*„Z8 G d+„ d,e	j0        ¦  «        Z9d-ej:        d.e;d/ej:        fd0„Z<	 dRd2e	j0        d3ej:        d4ej:        d5ej:        d6ej:        dz  d7e=d8e=d9ee"         fd:„Z> G d;„ d<e	j0        ¦  «        Z? G d=„ d>e¦  «        Z@ G d?„ d@e	j0        ¦  «        ZA G dA„ dBe	j0        ¦  «        ZB G dC„ dDe	j0        ¦  «        ZC G dE„ dFe	j0        ¦  «        ZD G dG„ dHe	j0        ¦  «        ZE G dI„ dJe	j0        ¦  «        ZFe# G dK„ dLe¦  «        ¦   «         ZG e#dM¬N¦  «         G dO„ dPeG¦  «        ¦   «         ZHdPdLgZIdS )SzPyTorch Mimi model.é    N)ÚCallable)Ú	dataclass)ÚOptional)Únné   )Úinitialization)ÚACT2FN)ÚCacheÚDynamicCache)Ú!create_sliding_window_causal_mask)ÚGradientCheckpointingLayer)ÚBaseModelOutputWithPast)ÚROPE_INIT_FUNCTIONSÚdynamic_rope_update)ÚALL_ATTENTION_FUNCTIONSÚPreTrainedModel)ÚUnpack)ÚModelOutputÚTransformersKwargsÚauto_docstringÚlogging)Úmaybe_autocasté   )Ú
MimiConfigc                   óx   — 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dz  ed<   dZedz  ed<   dS )Ú
MimiOutputaV  
    audio_codes (`torch.LongTensor`  of shape `(batch_size, num_quantizers, codes_length)`, *optional*):
        Discrete code embeddings computed using `model.encode`.
    audio_values (`torch.FloatTensor` of shape `(batch_size, sequence_length)`, *optional*):
        Decoded audio values, obtained using the decoder part of Mimi.
    encoder_past_key_values (`Cache`, *optional*):
        Pre-computed hidden-states (key and values in the self-attention blocks) that can be used to speed up sequential decoding of the encoder transformer.
        This typically consists in the `past_key_values` returned by the model at a previous stage of decoding, when `use_cache=True` or `config.use_cache=True`.

        The model will output the same cache format that is fed as input.

        If `past_key_values` are used, the user can optionally input only the last `audio_values` or `audio_codes (those that don't
        have their past key value states given to this model).
    decoder_past_key_values (`Cache`, *optional*):
        Pre-computed hidden-states (key and values in the self-attention blocks) that can be used to speed up sequential decoding of the decoder transformer.
        This typically consists in the `past_key_values` returned by the model at a previous stage of decoding, when `use_cache=True` or `config.use_cache=True`.

        The model will output the same cache format that is fed as input.

        If `past_key_values` are used, the user can optionally input only the last `audio_values` or `audio_codes (those that don't
        have their past key value states given to this model).
    NÚaudio_codesÚaudio_valuesÚencoder_past_key_valuesÚdecoder_past_key_values)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚtorchÚ
LongTensorÚ__annotations__r   ÚFloatTensorr   r
   r    © ó    úd/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/mimi/modeling_mimi.pyr   r   )   sx   € € € € € € ðð ð. ,0€K�Ô! DÑ(Ð/Ð/Ñ/Ø-1€L�%Ô# dÑ*Ð1Ð1Ñ1Ø,0Ð˜U T™\Ð0Ð0Ñ0Ø,0Ð˜U T™\Ð0Ð0Ñ0Ð0Ð0r*   r   c            	       ó‚   — e Zd ZdZdedee         dee         dee         fd„Zdej	        defd	„Z
dej	        defd
„ZdS )ÚMimiConv1dPaddingCachea�  
    Padding cache for MimiConv1d causal convolutions in order to support streaming via cache padding.
    See: https://huggingface.co/papers/2005.06720 & https://huggingface.co/papers/2204.07064

    A padding cache is a list of cached partial hidden states for each convolution layer.
    Hidden states are cached from the previous call to the MimiConv1d forward pass, given the padding size.
    Ú
num_layersÚper_layer_paddingÚper_layer_padding_modeÚper_layer_in_channelsc                 ó  — t          |¦  «        t          |¦  «        t          |¦  «        h}t          |¦  «        dk    s|                     ¦   «         |k    rt          d|› d�¦  «        ‚|| _        || _        || _        d g|z  | _        d S )Nr   zExpected `num_layers` (zU) values in `per_layer_padding`, `per_layer_padding_mode` and `per_layer_in_channels`)ÚlenÚpopÚ
ValueErrorr/   r0   r1   Úpadding_cache)Úselfr.   r/   r0   r1   Úfrom_args_num_layerss         r+   Ú__init__zMimiConv1dPaddingCache.__init__R   s³   € õ !$Ð$5Ñ 6Ô 6½Ð<RÑ8SÔ8SÕUXÐYnÑUoÔUoÐpÐåÐ#Ñ$Ô$¨Ò)Ð)Ð-A×-EÒ-EÑ-GÔ-GÈ:Ò-UÐ-UÝð L¨*ð  Lð  Lð  Lñô ð ð "3ˆÔØ&<ˆÔ#Ø%:ˆÔ"à"˜V jÑ0ˆÔÐÐr*   Úhidden_statesÚ	layer_idxc                 óJ  — |j         d         |j        |j        }}}| j        |         | j        |         | j        |         }}}|dk    rt          j        |||||¬¦  «        }	n@|dk    r't          j        |||||¬¦  «        |ddd…f         z  }	nt          d|› d	�¦  «        ‚|	S )
ad  
        Initialize the cache for a specific layer.

        Parameters:
            hidden_states (`torch.Tensor`):
                The hidden states to initialize the cache with.
            layer_idx (`int`):
                The index of the layer to initialize the cache for.
        Returns:
            `torch.Tensor`, the initialized cache.
        r   Úconstant©ÚdeviceÚdtypeÚ	replicate.Nr   zPadding mode z not supported)
Úshaper@   r?   r/   r0   r1   r%   ÚzerosÚonesÚNotImplementedError)
r7   r:   r;   Ú
batch_sizer@   r?   ÚpaddingÚpadding_modeÚin_channelsÚcurrent_caches
             r+   Ú_cache_initz"MimiConv1dPaddingCache._cache_initg   sÚ   € ð %2Ô$7¸Ô$:¸MÔ<OÐQ^ÔQe˜6�Eˆ
àÔ" 9Ô-ØÔ'¨	Ô2ØÔ& yÔ1ð  +�ˆð ˜:Ò%Ð%Ý!œK¨
°KÀÐQWÐ_dÐeÑeÔeˆMˆMØ˜[Ò(Ð(å”
˜: {°GÀFÐRWÐXÑXÔXÐ[hÐilÐnpÐopÐnpÐipÔ[qÑqð ˆMõ &Ð&R°lÐ&RÐ&RÐ&RÑSÔSÐSàÐr*   c                 óä  — |j         d         |j        |j        }}}| j        |         | j        |         }}| j        |         €|                      ||¦  «        }n| j        |         }|dk    r`t          d||j         d         z
  ¦  «        }	|	dk    r)t          j	        |dd…dd…|	 d…f         |gd¬¦  «        }
n,|dd…dd…| d…f         }
nt          j
        ||d||¬¦  «        }
|
| j        |<   |S )a¬  
        Updates the padding cache with the new padding states for the layer `layer_idx` and returns the current cache.

        Parameters:
            hidden_states (`torch.Tensor`):
                The hidden states to be partially cached.
            layer_idx (`int`):
                The index of the layer to cache the states for.
        Returns:
            `torch.Tensor` or `None`, the current padding cache.
        r   Néÿÿÿÿ©Údim)r@   r?   )rB   r@   r?   r/   r1   r6   rK   Úmaxr%   ÚcatÚempty)r7   r:   r;   rF   r@   r?   rG   rI   rJ   Ú	shortfallÚpadding_statess              r+   ÚupdatezMimiConv1dPaddingCache.update…   s   € ð %2Ô$7¸Ô$:¸MÔ<OÐQ^ÔQe˜6�Eˆ
Ø#Ô5°iÔ@À$ÔB\Ð]fÔBg�ˆàÔ˜iÔ(Ð0Ø ×,Ò,¨]¸IÑFÔFˆMˆMà Ô.¨yÔ9ˆMð �QŠ;ˆ;Ý˜A˜w¨Ô)<¸RÔ)@Ñ@ÑAÔAˆIØ˜1Š}ˆ}Ý!&¤¨M¸!¸!¸!¸Q¸Q¸QÀÀ
ÀÀÐ:KÔ,LÈmÐ+\ÐbdÐ!eÑ!eÔ!e��à!.¨q¨q¨q°!°!°!°g°X°Y°Y¨Ô!?��å"œ[¨°[À!È5ÐY_Ð`Ñ`Ô`ˆNà(6ˆÔ˜9Ñ%ØÐr*   N)r!   r"   r#   r$   ÚintÚlistÚstrr9   r%   ÚTensorrK   rU   r)   r*   r+   r-   r-   I   s¬   € € € € € ðð ð1àð1ð   œ9ð1ð !% S¤	ð	1ð
  $ Cœyð1ð 1ð 1ð 1ð*¨¬ð À#ð ð ð ð ð< E¤Lð ¸Sð ð ð ð ð ð r*   r-   c                   óZ   — e Zd ZU dZdZej        dz  ed<   dZe	dz  ed<   dZ
edz  ed<   dS )ÚMimiEncoderOutputaÔ  
    audio_codes (`torch.LongTensor`  of shape `(batch_size, num_quantizers, codes_length)`, *optional*):
        Discrete code embeddings computed using `model.encode`.
    encoder_past_key_values (`Cache`, *optional*):
        Pre-computed hidden-states (key and values in the self-attention blocks) that can be used to speed up sequential decoding of the encoder transformer.
        This typically consists in the `past_key_values` returned by the model at a previous stage of decoding, when `use_cache=True` or `config.use_cache=True`.

        The model will output the same cache format that is fed as input.

        If `past_key_values` are used, the user can optionally input only the last `audio_values` or `audio_codes (those that don't
        have their past key value states given to this model).
    padding_cache (`MimiConv1dPaddingCache`, *optional*):
        Padding cache for MimiConv1d causal convolutions in order to support streaming via cache padding.
    Nr   r   r6   )r!   r"   r#   r$   r   r%   r&   r'   r   r
   r6   r-   r)   r*   r+   r[   r[   §   sa   € € € € € € ðð ð ,0€K�Ô! DÑ(Ð/Ð/Ñ/Ø,0Ð˜U T™\Ð0Ð0Ñ0Ø37€MÐ)¨DÑ0Ð7Ð7Ñ7Ð7Ð7r*   r[   c                   óF   — e Zd ZU dZdZej        dz  ed<   dZe	dz  ed<   dS )ÚMimiDecoderOutputa+  
    audio_values (`torch.FloatTensor`  of shape `(batch_size, segment_length)`, *optional*):
        Decoded audio values, obtained using the decoder part of Mimi.
    decoder_past_key_values (`Cache`, *optional*):
        Pre-computed hidden-states (key and values in the self-attention blocks) that can be used to speed up sequential decoding of the decoder transformer.
        This typically consists in the `past_key_values` returned by the model at a previous stage of decoding, when `use_cache=True` or `config.use_cache=True`.

        The model will output the same cache format that is fed as input.

        If `past_key_values` are used, the user can optionally input only the last `audio_values` or `audio_codes (those that don't
        have their past key value states given to this model).
    Nr   r    )
r!   r"   r#   r$   r   r%   r(   r'   r    r
   r)   r*   r+   r]   r]   ¾   sK   € € € € € € ðð ð .2€L�%Ô# dÑ*Ð1Ð1Ñ1Ø,0Ð˜U T™\Ð0Ð0Ñ0Ð0Ð0r*   r]   c                   ó  ‡ — e Zd ZdZ	 	 	 	 	 	 ddedededed	ed
ededz  dededz  fˆ fd„Zd„ Zd„ Z	de
j        de
j        fd„Zedde
j        deeef         dedefd„¦   «         Zde
j        de
j        fd„Zdd„Zˆ xZS ) Ú
MimiConv1dz;Conv1d with asymmetric or causal padding and normalization.r   NTrI   Úout_channelsÚkernel_sizeÚstrideÚdilationÚgroupsÚpad_modeÚbiasr;   c           	      ó  •— t          ¦   «                              ¦   «          |j        | _        |€|j        n|| _        |
| _        || _        |dk    r*|dk    r$t                               d|› d|› d|› d�¦  «         t          j
        |||||||	¬¦  «        | _        | j        j        d         }t          j        | j        j        d         t          j        ¬¦  «        }| j        j        d         }t          j        |dz
  |z  dz   t          j        ¬¦  «        }|                      d	|d
¬¦  «         |                      d|d
¬¦  «         |                      d||z
  d
¬¦  «         | j        dz  | _        | j        | j        z
  | _        d S )Nr   zNMimiConv1d has been initialized with stride > 1 and dilation > 1 (kernel_size=z stride=z, dilation=ú).)rc   rd   rf   r   ©r@   rb   F©Ú
persistentra   Úpadding_totalé   )Úsuperr9   Úuse_causal_convÚcausalre   r;   rI   ÚloggerÚwarningr   ÚConv1dÚconvra   r%   Útensorrb   Úint64rc   Úregister_bufferrl   Úpadding_rightÚpadding_left)r7   ÚconfigrI   r`   ra   rb   rc   rd   re   rf   r;   Ú	__class__s              €r+   r9   zMimiConv1d.__init__Õ   sº  ø€ õ 	‰Œ×ÒÑÔÐØÔ,ˆŒØ+3Ð+;˜œ˜ÀˆŒØ"ˆŒØ&ˆÔð �AŠ:ˆ:˜( Qš,˜,Ý�NŠNðVØ!,ðVð VØ6<ðVð VØIQðVð Vð Vñô ð õ
 ”IØ˜ {°FÀXÐV\Ðcgð
ñ 
ô 
ˆŒ	ð ”iÔ+¨AÔ.ˆÝ”˜dœiÔ.¨qÔ1½¼ÐEÑEÔEˆØ”9Ô% aÔ(ˆõ ”l K°!¡O°xÑ#?À!Ñ#CÍ5Ì;ÐWÑWÔWˆà×Ò˜X v¸%ÐÑ@Ô@Ð@Ø×Ò˜]¨KÀEÐÑJÔJÐJØ×Ò˜_¨k¸FÑ.BÈuÐÑUÔUÐUð "Ô/°1Ñ4ˆÔØ Ô.°Ô1CÑCˆÔÐÐr*   c                 ó²   — t           j        j        }t          t           j        j        d¦  «        rt           j        j        j        } || j        ¦  «         d S ©NÚweight_norm©r   Úutilsr~   ÚhasattrÚparametrizationsrt   ©r7   r~   s     r+   Úapply_weight_normzMimiConv1d.apply_weight_norm  óI   € Ý”hÔ*ˆÝ•2”8Ô,¨mÑ<Ô<ð 	@Ýœ(Ô3Ô?ˆKàˆ�D”IÑÔÐÐÐr*   c                 óN   — t           j                             | j        ¦  «         d S ©N©r   r€   Úremove_weight_normrt   ©r7   s    r+   r‰   zMimiConv1d.remove_weight_norm	  ó    € Ý
Œ×#Ò# D¤IÑ.Ô.Ð.Ð.Ð.r*   r:   Úreturnc                 óü   — |j         d         }|| j        z
  | j        z   | j        z  dz   }t	          j        |¦  «                             t          j        ¦  «        dz
  }|| j        z  | j        z   | j        z
  }||z
  S )zSee `pad_for_conv1d`.rM   r   )rB   ra   rl   rb   r%   ÚceilÚtorv   )r7   r:   ÚlengthÚn_framesÚideal_lengths        r+   Ú_get_extra_padding_for_conv1dz(MimiConv1d._get_extra_padding_for_conv1d  s~   € ð
 Ô$ RÔ(ˆØ˜TÔ-Ñ-°Ô0BÑBÀdÄkÑQÐTUÑUˆÝ”:˜hÑ'Ô'×*Ò*­5¬;Ñ7Ô7¸!Ñ;ˆØ $¤+Ñ-°Ô0@Ñ@À4ÔCUÑUˆà˜fÑ$Ð$r*   Úzeroç        ÚpaddingsÚmodeÚvaluec                 óv  — | j         d         }|\  }}|dk    r"t          j                             | |||¦  «        S t	          ||¦  «        }d}||k    r*||z
  dz   }t          j                             | d|f¦  «        } t          j                             | |||¦  «        }	|	j         d         |z
  }
|	dd|
…f         S )zÊTiny wrapper around torch.nn.functional.pad, just to allow for reflect padding on small input.
        If this is the case, we insert extra 0 padding to the right before the reflection happens.
        rM   Úreflectr   r   .N)rB   r   Ú
functionalÚpadrP   )r:   r–   r—   r˜   r�   ry   rx   Úmax_padÚ	extra_padÚpaddedÚends              r+   Ú_pad1dzMimiConv1d._pad1d  sÊ   € ð Ô$ RÔ(ˆØ&.Ñ#ˆ�mØ�9ÒÐÝ”=×$Ò$ ]°H¸dÀEÑJÔJÐJå�l MÑ2Ô2ˆØˆ	Ø�WÒÐØ &Ñ(¨1Ñ,ˆIÝœM×-Ò-¨m¸aÀ¸^ÑLÔLˆMÝ”×"Ò" =°(¸DÀ%ÑHÔHˆØŒl˜2Ô Ñ*ˆØ�c˜4˜C˜4�iÔ Ð r*   Úinput_lengthc                 óî  — || j         z
  | j        z   | j        z  dz   }t          j        |¦  «                             t          j        ¦  «        dz
  }|| j        z  | j         z   | j        z
  }||z
  }| j        r
| j        }|}n| j        }| j	        |z   }||z   |z   }|d| j
        j        d         z  z   | j
        j        d         | j
        j         d         dz
  z  z
  dz
  | j
        j        d         z  dz   }|S )zD
        Return the length of the output of the MimiConv1d.
        r   rm   r   )ra   rl   rb   r%   rŽ   r�   rv   rp   ry   rx   rt   rG   rc   )r7   r¢   r‘   r’   Úextra_paddingry   rx   Úoutput_lengths           r+   Ú_get_output_lengthzMimiConv1d._get_output_length-  s  € ð
 ! 4Ô#3Ñ3°dÔ6HÑHÈDÌKÑWÐZ[Ñ[ˆÝ”:˜hÑ'Ô'×*Ò*­5¬;Ñ7Ô7¸!Ñ;ˆØ $¤+Ñ-°Ô0@Ñ@À4ÔCUÑUˆØ$ |Ñ3ˆàŒ;ð 	?ØÔ-ˆLØ)ˆMˆMàÔ,ˆLØ Ô.°Ñ>ˆMð $ lÑ2°]ÑBˆð ˜1˜tœyÔ0°Ô3Ñ3Ñ3°d´iÔ6HÈÔ6KÈtÌyÔOdÐefÔOgÐjkÑOkÑ6lÑlÐopÑpØŒYÔ˜aÔ ñ!à#$ñ%ˆð Ðr*   c                 ó´  — |                       |¦  «        }| j        s|�t          d¦  «        ‚| j        r6|�4|                     || j        ¦  «        }t          j        ||gd¬¦  «        }nX| j        r%|                      || j        |f| j	        ¬¦  «        }n,|                      || j
        | j        |z   f| j	        ¬¦  «        }|                      |¦  «        }|S )Nz=`padding_cache` is not supported for non-causal convolutions.rm   rN   )r—   )r“   rp   r5   rU   r;   r%   rQ   r¡   rl   re   ry   rx   rt   )r7   r:   r6   r¤   Úlayer_padding_caches        r+   ÚforwardzMimiConv1d.forwardG  sõ   € Ø×:Ò:¸=ÑIÔIˆàŒ{ð 	^˜}Ð8ÝÐ\Ñ]Ô]Ð]àŒ;ð 	˜=Ð4Ø"/×"6Ò"6°}ÀdÄnÑ"UÔ"UÐÝ!œIÐ':¸MÐ&JÐPQÐRÑRÔRˆMˆMàŒ[ð 	à ŸKšK¨¸Ô8JÈMÐ7ZÐaeÔan˜KÑoÔoˆMˆMð !ŸKšKØ Ô 1°4Ô3EÈÑ3UÐVÐ]aÔ]jð (ñ ô ˆMð Ÿ	š	 -Ñ0Ô0ˆØÐr*   )r   r   r   NTN)r”   r•   r‡   )r!   r"   r#   r$   rV   rX   Úboolr9   r„   r‰   r%   rY   r“   ÚstaticmethodÚtupleÚfloatr¡   r&   r¦   r©   Ú__classcell__©r{   s   @r+   r_   r_   Ò   sŸ  ø€ € € € € ØEÐEð ØØØ#ØØ $ð+Dð +Dð ð+Dð ð	+Dð
 ð+Dð ð+Dð ð+Dð ð+Dð ˜‘*ð+Dð ð+Dð ˜‘:ð+Dð +Dð +Dð +Dð +Dð +DðZð ð ð/ð /ð /ð
%à”|ð
%ð 
Œð
%ð 
%ð 
%ð 
%ð ð!ð !˜eœlð !°e¸CÀ¸H´oð !ÈSð !Ðbgð !ð !ð !ñ „\ð!ð$¨uÔ/?ð ÀEÔDTð ð ð ð ð4ð ð ð ð ð ð ð r*   r_   c                   óR   ‡ — e Zd ZdZ	 	 	 ddededededef
ˆ fd	„Zd
„ Zd„ Zd„ Zˆ xZ	S )ÚMimiConvTranspose1dzDConvTranspose1d with asymmetric or causal padding and normalization.r   TrI   r`   ra   rb   rd   c                 óÎ  •— t          ¦   «                              ¦   «          |j        | _        |j        | _        t          j        ||||||¬¦  «        | _        | j        s| j        dk    st          d¦  «        ‚| j        j	        d         }| j        j
        d         }||z
  }| j        r"t          j        || j        z  ¦  «        | _        n
|dz  | _        || j        z
  | _        d S )N)rd   rf   ç      ð?zB`trim_right_ratio` != 1.0 only makes sense for causal convolutionsr   rm   )rn   r9   ro   rp   Útrim_right_ratior   ÚConvTranspose1drt   r5   ra   rb   ÚmathrŽ   rx   ry   )
r7   rz   rI   r`   ra   rb   rd   rf   rl   r{   s
            €r+   r9   zMimiConvTranspose1d.__init__a  sð   ø€ õ 	‰Œ×ÒÑÔÐØÔ,ˆŒØ &Ô 7ˆÔÝÔ& {°LÀ+ÈvÐ^dÐkoÐpÑpÔpˆŒ	à”ð 	c˜tÔ4¸Ò;Ð;ÝÐaÑbÔbÐbà”iÔ+¨AÔ.ˆØ”Ô! !Ô$ˆØ# fÑ,ˆð Œ;ð 	4õ "&¤¨=¸4Ô;PÑ+PÑ!QÔ!QˆDÔÐð "/°!Ñ!3ˆDÔà)¨DÔ,>Ñ>ˆÔÐÐr*   c                 ó²   — t           j        j        }t          t           j        j        d¦  «        rt           j        j        j        } || j        ¦  «         d S r}   r   rƒ   s     r+   r„   z%MimiConvTranspose1d.apply_weight_norm…  r…   r*   c                 óN   — t           j                             | j        ¦  «         d S r‡   rˆ   rŠ   s    r+   r‰   z&MimiConvTranspose1d.remove_weight_normŒ  r‹   r*   c                 ó|   — |                       |¦  «        }|j        d         | j        z
  }|d| j        |…f         }|S )NrM   .)rt   rB   rx   ry   )r7   r:   r    s      r+   r©   zMimiConvTranspose1d.forward�  sG   € ØŸ	š	 -Ñ0Ô0ˆð Ô! "Ô%¨Ô(:Ñ:ˆØ% c¨4Ô+<¸sÐ+BÐ&BÔCˆØÐr*   )r   r   T)
r!   r"   r#   r$   rV   r9   r„   r‰   r©   r®   r¯   s   @r+   r±   r±   ^  s³   ø€ € € € € ØNÐNð ØØð"?ð "?ð ð"?ð ð	"?ð
 ð"?ð ð"?ð ð"?ð "?ð "?ð "?ð "?ð "?ðHð ð ð/ð /ð /ðð ð ð ð ð ð r*   r±   c                   óD   ‡ — e Zd ZdZdededee         fˆ fd„Zdd„Zˆ xZ	S )	ÚMimiResnetBlockz;
    Residual block from SEANet model as used by Mimi.
    rz   rO   Ú	dilationsc           	      óf  •— t          ¦   «                              ¦   «          |j        df}t          |¦  «        t          |¦  «        k    rt	          d¦  «        ‚||j        z  }g }t          t          ||¦  «        ¦  «        D ][\  }\  }}	|dk    r|n|}
|t          |¦  «        dz
  k    r|n|}|t          j	        ¦   «         gz  }|t          ||
|||	¬¦  «        gz  }Œ\t          j        |¦  «        | _        |j        rt          |||d¬¦  «        | _        d S t          j        ¦   «         | _        d S )Nr   z7Number of kernel sizes should match number of dilationsr   )rc   )ra   )rn   r9   Úresidual_kernel_sizer3   r5   ÚcompressÚ	enumerateÚzipr   ÚELUr_   Ú
ModuleListÚblockÚuse_conv_shortcutÚshortcutÚIdentity)r7   rz   rO   r¼   Úkernel_sizesÚhiddenrÄ   Úira   rc   Úin_chsÚout_chsr{   s               €r+   r9   zMimiResnetBlock.__init__�  s;  ø€ Ý‰Œ×ÒÑÔÐØÔ3°QÐ7ˆÝˆ|ÑÔ¥ I¡¤Ò.Ð.ÝÐVÑWÔWÐWà˜œÑ'ˆØˆÝ*3µC¸ÀiÑ4PÔ4PÑ*QÔ*Qð 	[ð 	[Ñ&ˆAÑ&�˜XØ šF˜F�S�S¨ˆFØ¥# lÑ"3Ô"3°aÑ"7Ò7Ð7�c�c¸VˆGØ•b”f‘h”h�ZÑˆEØ•j ¨°¸+ÐPXÐYÑYÔYÐZÑZˆEˆEÝ”] 5Ñ)Ô)ˆŒ
àÔ#ð 	*Ý& v¨s°CÀQÐGÑGÔGˆDŒMˆMˆMåœK™MœMˆDŒMˆMˆMr*   Nc                 ó  — |}| j         D ]0}t          |t          ¦  «        r |||¬¦  «        }Œ% ||¦  «        }Œ1t          | j        t          ¦  «        r|                      ||¬¦  «        }n|                      |¦  «        }||z   S ©N©r6   )rÄ   Ú
isinstancer_   rÆ   )r7   r:   r6   ÚresidualÚlayers        r+   r©   zMimiResnetBlock.forward±  s˜   € Ø ˆà”Zð 	5ð 	5ˆEÝ˜%¥Ñ,Ô,ð 5Ø %  mÀ=Ð QÑ QÔ Q��à %  mÑ 4Ô 4��å�d”m¥ZÑ0Ô0ð 	/Ø—}’} X¸]�}ÑKÔKˆHˆHà—}’} XÑ.Ô.ˆHà˜-Ñ'Ð'r*   r‡   )
r!   r"   r#   r$   r   rV   rW   r9   r©   r®   r¯   s   @r+   r»   r»   ˜  st   ø€ € € € € ðð ð*˜zð *°ð *ÀÀSÄ	ð *ð *ð *ð *ð *ð *ð((ð (ð (ð (ð (ð (ð (ð (r*   r»   c                   ó0   ‡ — e Zd ZdZdefˆ fd„Zdd„Zˆ xZS )ÚMimiEncoderzSEANet encoder as used by Mimi.rz   c           	      óü  •— t          ¦   «                              ¦   «          t          ||j        |j        |j        ¦  «        g}d}dg}t          |j        ¦  «        D ]Ú}||j        z  }t          |j	        ¦  «        D ]Z}| 
                    dt          |¦  «        › d�dt          |¦  «        › d�g¦  «         |t          |||j        |z  dg¦  «        gz  }Œ[|t          j        ¦   «         gz  }|                     dt          |¦  «        › �¦  «         |t          |||dz  |dz  |¬¦  «        gz  }|dz  }ŒÛ|t          j        ¦   «         gz  }|                     dt          |¦  «        › �¦  «         |t          |||j        z  |j        |j        ¦  «        gz  }t          j        |¦  «        | _        || _        t-          | j        ¦  «        D ]+\  }}	|                      |	¦  «        }
t1          |
d|¦  «         Œ,d S )	Nr   zlayers.0zlayers.z.block.1z.block.3rm   ©ra   rb   r;   )rn   r9   r_   Úaudio_channelsÚnum_filtersra   ÚreversedÚupsampling_ratiosÚrangeÚnum_residual_layersÚextendr3   r»   Údilation_growth_rater   rÂ   ÚappendÚhidden_sizeÚlast_kernel_sizerÃ   ÚlayersÚ_mimiconv1d_layer_namesrÀ   Úget_submoduleÚsetattr)r7   rz   ÚmodelÚscalingÚmimiconv1d_layer_namesÚratioÚcurrent_scaleÚjr;   Ú	layernameÚ
conv_layerr{   s              €r+   r9   zMimiEncoder.__init__Å  s/  ø€ Ý‰Œ×ÒÑÔÐÝ˜F FÔ$9¸6Ô;MÈvÔOaÑbÔbÐcˆØˆð #- Ðõ ˜fÔ6Ñ7Ô7ð 
	ð 
	ˆEØ# fÔ&8Ñ8ˆMå˜6Ô5Ñ6Ô6ð gð g�Ø&×-Ò-Ð/M½¸U¹¼Ð/MÐ/MÐ/MÐOmÕY\Ð]bÑYcÔYcÐOmÐOmÐOmÐ.nÑoÔoÐoØ�/¨&°-À&ÔB]Ð_`ÑB`ÐbcÐAdÑeÔeÐfÑf��à•b”f‘h”h�ZÑˆEØ"×)Ò)Ð*@µC¸±J´JÐ*@Ð*@ÑAÔAÐAØ•j ¨¸ÈÑ8IÐW\Ð_`ÑW`ÐinÐoÑoÔoÐpÑpˆEØ�q‰LˆGˆGà•"”&‘(”(�ÑˆØ×%Ò%Ð&<µ°E±
´
Ð&<Ð&<Ñ=Ô=Ð=Ø•*˜V W¨vÔ/AÑ%AÀ6ÔCUÐW]ÔWnÑoÔoÐpÑpˆå”m EÑ*Ô*ˆŒØ'=ˆÔ$õ %.¨dÔ.JÑ$KÔ$Kð 	8ð 	8Ñ ˆI�yØ×+Ò+¨IÑ6Ô6ˆJÝ�J ¨YÑ7Ô7Ð7Ð7ð	8ð 	8r*   Nc                 ó„   — | j         D ]7}t          |t          t          f¦  «        r |||¬¦  «        }Œ, ||¦  «        }Œ8|S rÎ   )râ   rÐ   r_   r»   )r7   r:   r6   rÒ   s       r+   r©   zMimiEncoder.forwardæ  sW   € Ø”[ð 	5ð 	5ˆEÝ˜%¥*­oÐ!>Ñ?Ô?ð 5Ø %  mÀ=Ð QÑ QÔ Q��à %  mÑ 4Ô 4��ØÐr*   r‡   ©r!   r"   r#   r$   r   r9   r©   r®   r¯   s   @r+   rÔ   rÔ   Â  s_   ø€ € € € € Ø)Ð)ð8˜zð 8ð 8ð 8ð 8ð 8ð 8ðBð ð ð ð ð ð ð r*   rÔ   c                   ó8   ‡ — e Zd ZdZˆ fd„Zdej        fd„Zˆ xZS )ÚMimiLayerScalez©Layer scale from [Touvron et al 2021] (https://huggingface.co/papers/2103.17239).
    This rescales diagonally the residual outputs close to 0, with a learnt scale.
    c                 óÂ   •— t          ¦   «                              ¦   «          |j        }|j        }t	          j        t          j        |f|d¬¦  «        ¦  «        | _        d S )NT)Úrequires_grad)	rn   r9   rà   Úlayer_scale_initial_scaler   Ú	Parameterr%   ÚfullÚscale)r7   rz   ÚchannelsÚinitial_scaler{   s       €r+   r9   zMimiLayerScale.__init__ô  sR   ø€ Ý‰Œ×ÒÑÔÐØÔ%ˆØÔ8ˆÝ”\¥%¤*¨h¨[¸-ÐW[Ð"\Ñ"\Ô"\Ñ]Ô]ˆŒ
ˆ
ˆ
r*   Úxc                 ó   — | j         |z  S r‡   )r÷   )r7   rú   s     r+   r©   zMimiLayerScale.forwardú  s   € ØŒz˜A‰~Ðr*   )	r!   r"   r#   r$   r9   r%   rY   r©   r®   r¯   s   @r+   rñ   rñ   ï  sd   ø€ € € € € ðð ð^ð ^ð ^ð ^ð ^ð˜œð ð ð ð ð ð ð ð r*   rñ   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 )ÚMimiRotaryEmbeddingÚinv_freqNrz   c                 ó²  •— t          ¦   «                              ¦   «          |j        | _        |j        | _        || _        | j        j        d         | _        | j        }| j        dk    rt          | j                 } || j        |¦  «        \  }| _
        |                      d|d¬¦  «         |                      d|                     ¦   «         d¬¦  «         d S )NÚ	rope_typeÚdefaultrþ   Frj   Úoriginal_inv_freq)rn   r9   Úmax_position_embeddingsÚmax_seq_len_cachedÚoriginal_max_seq_lenrz   Úrope_parametersr   Úcompute_default_rope_parametersr   Úattention_scalingrw   Úclone)r7   rz   r?   Úrope_init_fnrþ   r{   s        €r+   r9   zMimiRotaryEmbedding.__init__  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*   r?   ztorch.deviceÚseq_lenrŒ   z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_dimNr³   r   rm   ri   r>   )	r  Úgetattrrà   Únum_attention_headsr%   Úarangerv   r�   r­   )rz   r?   r  ÚbaserO   Úattention_factorrþ   s          r+   r  z3MimiRotaryEmbedding.compute_default_rope_parameters  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   rM   r   ÚmpsÚcpuF)Údevice_typeÚenabledrm   rN   ri   )rþ   r­   ÚexpandrB   r�   r?   rÐ   ÚtyperX   r   Ú	transposer%   rQ   Úcosr  Úsinr@   )
r7   rú   Úposition_idsÚinv_freq_expandedÚposition_ids_expandedr  ÚfreqsÚembr  r  s
             r+   r©   zMimiRotaryEmbedding.forward0  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*r‡   ©NNN)r!   r"   r#   r%   rY   r'   r   r9   r«   r   rV   r¬   r­   r  Úno_gradr   r©   r®   r¯   s   @r+   rý   rý   ÿ  sù   ø€ € € € € € ØŒlÐÐÑðVð V˜zð Vð Vð Vð Vð Vð Vð  à$(Ø+/Ø"ð*ð *Ø˜TÑ!ð*à˜Ô(ð*ð �t‘ð*ð 
ˆ~˜uÐ$Ô	%ð	*ð *ð *ñ „\ð*ð: €U„]�_„_Øð<ð <ñ Ôñ „_ð<ð <ð <ð <ð <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..NrM   rm   rN   )rB   r%   rQ   )rú   Úx1Úx2s      r+   Úrotate_halfr(  A  s]   € à	
ˆ3Ð"�!”'˜"”+ Ñ"Ð"Ð"Ô	#€BØ	
ˆ3�”˜”˜qÑ Ð"Ð"Ð"Ô	#€BÝŒ9�r�c˜2�Y BÐ'Ñ'Ô'Ð'r*   c                 ó¾   — |                      |¦  «        }|                      |¦  «        }| |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Úkr  r  Úunsqueeze_dimÚq_embedÚk_embeds          r+   Úapply_rotary_pos_embr0  I  sc   € ð$ �-Š-˜Ñ
&Ô
&€CØ
�-Š-˜Ñ
&Ô
&€CØ�3‰w�; q™>œ>¨CÑ/Ñ0€GØ�3‰w�; q™>œ>¨CÑ/Ñ0€GØ�GÐÐr*   c                   óB   ‡ — e Zd Zˆ fd„Zdej        dej        fd„Zˆ xZS )ÚMimiMLPc                 ó  •— t          ¦   «                              ¦   «          || _        t          |j                 | _        t          j        |j        |j	        d¬¦  «        | _
        t          j        |j	        |j        d¬¦  «        | _        d S )NF©rf   )rn   r9   rz   r	   Ú
hidden_actÚactivation_fnr   ÚLinearrà   Úintermediate_sizeÚfc1Úfc2©r7   rz   r{   s     €r+   r9   zMimiMLP.__init__c  sr   ø€ Ý‰Œ×ÒÑÔÐØˆŒÝ# FÔ$5Ô6ˆÔÝ”9˜VÔ/°Ô1IÐPUÐVÑVÔVˆŒÝ”9˜VÔ5°vÔ7IÐPUÐVÑVÔVˆŒˆˆr*   r:   rŒ   c                 ó„   — |                       |¦  «        }|                      |¦  «        }|                      |¦  «        }|S r‡   )r9  r6  r:  )r7   r:   s     r+   r©   zMimiMLP.forwardk  s=   € ØŸš Ñ/Ô/ˆØ×*Ò*¨=Ñ9Ô9ˆØŸš Ñ/Ô/ˆØÐr*   )r!   r"   r#   r9   r%   rY   r©   r®   r¯   s   @r+   r2  r2  b  sc   ø€ € € € € ðWð Wð Wð Wð Wð U¤\ð °e´lð ð ð ð ð ð ð ð r*   r2  r:   Ún_reprŒ   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)rB   r  Úreshape)r:   r=  ÚbatchÚnum_key_value_headsÚslenr  s         r+   Ú	repeat_kvrC  s  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*   r•   ÚmoduleÚqueryÚkeyr˜   Úattention_maskrç   Ú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 )Nrm   r   rM   )rO   r@   )ÚpÚtrainingr   )rC  Únum_key_value_groupsr%   Úmatmulr  r   r›   ÚsoftmaxÚfloat32r�   r@   rH  rL  Ú
contiguous)rD  rE  rF  r˜   rG  rç   rH  rI  Ú
key_statesÚvalue_statesÚattn_weightsÚattn_outputs               r+   Úeager_attention_forwardrV  €  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dededz  fˆ fd„Z	 	 	 ddej        dej        dz  de	dz  d	e
ej        ej        f         dz  d
e
ej        ej        f         f
d„Zˆ xZS )ÚMimiAttentionz=Multi-headed attention from 'Attention Is All You Need' paperNrz   r;   c                 ó‚  •— t          ¦   «                              ¦   «          || _        || _        |j        | _        |j        | _        |j        | _        |j        | _        |j	        | _	        | j        | j	        z  | _
        |j        | _        d| _        dt          j        |j        ¦  «        z  | _        | j        | j        z  dk    r t!          d| j        › d| 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        ¬¦  «        | _        |j        | _        d S )NTr   r   z?hidden_size must be divisible by num_heads (got `hidden_size`: z and `num_heads`: rh   r4  )rn   r9   rz   r;   Úattention_dropoutrà   r  Ú	num_headsr  rA  rM  r  Ú	is_causalr¶   Úsqrtrç   r5   r   r7  Úattention_biasÚq_projÚk_projÚv_projÚo_projÚsliding_window©r7   rz   r;   r{   s      €r+   r9   zMimiAttention.__init__œ  sœ  ø€ Ý‰Œ×ÒÑÔÐØˆŒØ"ˆŒà!'Ô!9ˆÔØ!Ô-ˆÔØÔ3ˆŒØœˆŒØ#)Ô#=ˆÔ Ø$(¤N°dÔ6NÑ$NˆÔ!Ø'-Ô'EˆÔ$ØˆŒØ�4œ9 V¤_Ñ5Ô5Ñ5ˆŒàÔ˜dœnÑ,°Ò1Ð1Ýð8ÐRVÔRbð 8ð 8Ø%)¤^ð8ð 8ð 8ñô ð õ
 ”i Ô 0°$´.À4Ä=Ñ2PÐW]ÔWlÐmÑmÔmˆŒÝ”i Ô 0°$Ô2JÈTÌ]Ñ2ZÐagÔavÐwÑwÔwˆŒÝ”i Ô 0°$Ô2JÈTÌ]Ñ2ZÐagÔavÐwÑwÔwˆŒÝ”i ¤°´Ñ >ÀÔ@PÐW]ÔWlÐmÑmÔmˆŒØ$Ô3ˆÔÐÐr*   r:   rG  Úpast_key_valuesÚposition_embeddingsrŒ   c                 ó.  — |j         d d…         }g |¢d‘| j        ‘R }|                      |¦  «                             |¦  «                             dd¦  «        }|                      |¦  «                             |¦  «                             dd¦  «        }	|                      |¦  «                             |¦  «                             dd¦  «        }
|\  }}t          ||	||¦  «        \  }}	|�|                     |	|
| j	        ¦  «        \  }	}
t          j        | j        j        t          ¦  «        } || ||	|
|f| j        sdn| j        | j        | j        dœ|¤Ž\  }} |j        g |¢d‘R Ž                      ¦   «         }|                      |¦  «        }||fS )NrM   r   rm   r•   )rH  rç   rc  )rB   r  r_  Úviewr  r`  ra  r0  rU   r;   r   Úget_interfacerz   Ú_attn_implementationrV  rL  rZ  rç   rc  r?  rQ  rb  )r7   r:   rG  re  rf  rI  Úinput_shapeÚhidden_shapeÚquery_statesrR  rS  r  r  Úattention_interfacerU  rT  s                   r+   r©   zMimiAttention.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Ý#7¸ÀjÐRUÐWZÑ#[Ô#[Ñ ˆ�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*   r‡   r#  )r!   r"   r#   r$   r   rV   r9   r%   rY   r
   r¬   r©   r®   r¯   s   @r+   rX  rX  ™  sÜ   ø€ € € € € ØGÐGð4ð 4˜zð 4°c¸D±jð 4ð 4ð 4ð 4ð 4ð 4ð< /3Ø(,ØHLð')ð ')à”|ð')ð œ tÑ+ð')ð  ™ð	')ð
 # 5¤<°´Ð#=Ô>ÀÑEð')ð 
ˆuŒ|˜Uœ\Ð)Ô	*ð')ð ')ð ')ð ')ð ')ð ')ð ')ð ')r*   rX  c                   óò   ‡ — e Zd Zdedefˆ fd„Z	 	 	 	 	 ddej        dej        dz  dedz  d	e	dz  d
e	dz  de
ej        ej        f         dz  de
ej        e
ej        ej        f         dz  f         fd„Zˆ xZS )ÚMimiTransformerLayerrz   r;   c                 ó˜  •— t          ¦   «                              ¦   «          |j        | _        t          ||¬¦  «        | _        t          |¦  «        | _        t          j        |j        |j	        ¬¦  «        | _
        t          j        |j        |j	        ¬¦  «        | _        t          |¦  «        | _        t          |¦  «        | _        d S )N)rz   r;   )Úeps)rn   r9   rà   rX  Ú	self_attnr2  Úmlpr   Ú	LayerNormÚnorm_epsÚinput_layernormÚpost_attention_layernormrñ   Úself_attn_layer_scaleÚmlp_layer_scalerd  s      €r+   r9   zMimiTransformerLayer.__init__â  s£   ø€ Ý‰Œ×ÒÑÔÐØ!Ô-ˆÔå&¨fÀ	ÐJÑJÔJˆŒå˜6‘?”?ˆŒÝ!œ|¨FÔ,>ÀFÄOÐTÑTÔTˆÔÝ(*¬°VÔ5GÈVÌ_Ð(]Ñ(]Ô(]ˆÔ%Ý%3°FÑ%;Ô%;ˆÔ"Ý-¨fÑ5Ô5ˆÔÐÐr*   NFr:   rG  re  Úoutput_attentionsÚ	use_cacherf  rŒ   c           
      ó0  — |}|                       |¦  «        } | j        d||||||dœ|¤Ž\  }}	||                      |¦  «        z   }|}|                      |¦  «        }|                      |¦  «        }||                      |¦  «        z   }|f}
|r|
|	fz  }
|
S )N)r:   rG  re  r{  r|  rf  r)   )rw  rs  ry  rx  rt  rz  )r7   r:   rG  re  r{  r|  rf  rI  rÑ   Úself_attn_weightsÚoutputss              r+   r©   zMimiTransformerLayer.forwardî  sÝ   € ð !ˆà×,Ò,¨]Ñ;Ô;ˆð ,:¨4¬>ð ,
Ø'Ø)Ø+Ø/ØØ 3ð,
ð ,
ð ð,
ð ,
Ñ(ˆÐ(ð ! 4×#=Ò#=¸mÑ#LÔ#LÑLˆð !ˆØ×5Ò5°mÑDÔDˆØŸš Ñ/Ô/ˆØ  4×#7Ò#7¸Ñ#FÔ#FÑFˆà Ð"ˆàð 	,ØÐ)Ð+Ñ+ˆGàˆr*   )NNFFN)r!   r"   r#   r   rV   r9   r%   rY   r
   rª   r¬   r(   r©   r®   r¯   s   @r+   rp  rp  á  s  ø€ € € € € ð
6˜zð 
6°cð 
6ð 
6ð 
6ð 
6ð 
6ð 
6ð /3Ø(,Ø).Ø!&ØHLð%ð %à”|ð%ð œ tÑ+ð%ð  ™ð	%ð
   $™;ð%ð ˜$‘;ð%ð # 5¤<°´Ð#=Ô>ÀÑEð%ð 
ˆuÔ  %¨Ô(9¸5Ô;LÐ(LÔ"MÐPTÑ"TÐTÔ	Uð%ð %ð %ð %ð %ð %ð %ð %r*   rp  c                   óº   ‡ — e Zd ZdZdefˆ fd„Z	 	 	 	 	 	 	 	 ddej        dz  dej        dz  dej        dz  de	dz  d	e
dz  d
e
dz  de
dz  de
dz  deez  fd„Zˆ xZS )ÚMimiTransformerModelz�
    Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`MimiTransformerLayer`]

    Args:
        config: MimiConfig
    rz   c                 óü   •‡— t          ¦   «                              ¦   «          t          j        ˆfd„t	          ‰j        ¦  «        D ¦   «         ¦  «        | _        t          ‰¦  «        | _        d| _	        ‰| _
        d S )Nc                 ó0   •— g | ]}t          ‰|¦  «        ‘ŒS r)   )rp  )Ú.0r;   rz   s     €r+   ú
<listcomp>z1MimiTransformerModel.__init__.<locals>.<listcomp>"  s$   ø€ ÐfÐfÐf¸Õ! &¨)Ñ4Ô4ÐfÐfÐfr*   F)rn   r9   r   rÃ   rÛ   Únum_hidden_layersrâ   rý   Ú
rotary_embÚgradient_checkpointingrz   r;  s    `€r+   r9   zMimiTransformerModel.__init__  ss   øø€ Ý‰Œ×ÒÑÔÐå”mØfÐfÐfÐfÅeÈFÔLdÑFeÔFeÐfÑfÔfñ
ô 
ˆŒõ .¨fÑ5Ô5ˆŒà&+ˆÔ#ØˆŒˆˆr*   Nr:   rG  r  re  r|  r{  Úoutput_hidden_statesÚreturn_dictrŒ   c	           
      ó  — |�|n| j         j        }|�|n| j         j        }|�|n| j         j        }|�|n| j         j        }| j        r%| j        r|rt                               d¦  «         d}|r|€t          | j         ¬¦  «        }|€V|�| 
                    ¦   «         nd}
t          j        |j        d         |j        ¬¦  «        |
z   }|                     d¦  «        }t!          | j         ||||¬¦  «        }|                      ||¦  «        }|rd	nd}|rd	nd}| j        D ]2}|r||fz  } ||||||||¬
¦  «        }|d         }|r||d         fz  }Œ3|r||fz  }|st'          d„ ||||fD ¦   «         ¦  «        S t)          ||||¬¦  «        S )a|  
        Args:
            hidden_states (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`, *optional*):
                Embedded representation that will be contextualized by the model
            attention_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
                Mask to avoid performing attention on padding token indices. Mask values selected in `[0, 1]`:

                - 1 for tokens that are **not masked**,
                - 0 for tokens that are **masked**.

                [What are attention masks?](../glossary#attention-mask)

                Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
                [`PreTrainedTokenizer.__call__`] for details.

                If `past_key_values` is used, optionally only the last `decoder_input_ids` have to be input (see
                `past_key_values`).

                If you want to change padding behavior, you should read [`modeling_opt._prepare_decoder_attention_mask`]
                and modify to your needs. See diagram 1 in [the paper](https://huggingface.co/papers/1910.13461) for more
                information on the default strategy.

                - 1 indicates the head is **not masked**,
                - 0 indicates the head is **masked**.
            position_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
                Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
                config.n_positions - 1]`.

                [What are position IDs?](../glossary#position-ids)
            past_key_values (`Cache`, *optional*):
                It is a [`~cache_utils.Cache`] instance. For more details, see our [kv cache guide](https://huggingface.co/docs/transformers/en/kv_cache).

                If `past_key_values` are used, the user can optionally input only the last `input_ids` (those that don't
                have their past key value states given to this model) of shape `(batch_size, 1)` instead of all `input_ids`
                of shape `(batch_size, sequence_length)`.
            use_cache (`bool`, *optional*):
                If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see
                `past_key_values`).
            output_attentions (`bool`, *optional*):
                Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned
                tensors for more detail.
            output_hidden_states (`bool`, *optional*):
                Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
                more detail.
            return_dict (`bool`, *optional*):
                Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
        NzX`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`.F)rz   r   r   ©r?   )rz   Úinputs_embedsrG  re  r  r)   )rG  r  re  r{  r|  rf  c              3   ó   K  — | ]}|®|V — Œ	d S r‡   r)   )r„  Úvs     r+   ú	<genexpr>z/MimiTransformerModel.forward.<locals>.<genexpr>Ÿ  s1   è è € ð ð ØÐbcÐbo�ÐboÐboÐboÐboðð r*   )Úlast_hidden_statere  r:   Ú
attentions)rz   r{  r‰  r|  rŠ  rˆ  rL  rq   Úwarning_oncer   Úget_seq_lengthr%   r  rB   r?   r*  r   r‡  râ   r¬   r   )r7   r:   rG  r  re  r|  r{  r‰  rŠ  rI  Úpast_seen_tokensÚcausal_maskrf  Úall_hidden_statesÚall_self_attnsÚdecoder_layerÚlayer_outputss                    r+   r©   zMimiTransformerModel.forward)  se  € ðv 2CÐ1NÐ-Ð-ÐTXÔT_ÔTqÐà$8Ð$DÐ Ð È$Ì+ÔJjð 	ð "+Ð!6�I�I¸D¼KÔ<Qˆ	à%0Ð%<�k�kÀ$Ä+ÔBYˆàÔ&ð 	¨4¬=ð 	¸Yð 	Ý×ÒØjñô ð ð ˆIàð 	?˜Ð0Ý*°$´+Ð>Ñ>Ô>ˆOàÐØCRÐC^˜×=Ò=Ñ?Ô?Ð?ÐdeÐÝ œ<¨Ô(;¸AÔ(>À}ÔG[Ð\Ñ\Ô\Ð_oÑoˆLØ'×1Ò1°!Ñ4Ô4ˆLå7Ø”;Ø'Ø)Ø+Ø%ð
ñ 
ô 
ˆð #Ÿošo¨m¸\ÑJÔJÐð #7Ð@˜B˜B¸DÐØ0Ð:˜˜°dˆà!œ[ð 	6ð 	6ˆMØ#ð 6Ø! mÐ%5Ñ5Ð!à)˜MØØ*Ø)Ø /Ø"3Ø#Ø$7ðñ ô ˆMð *¨!Ô,ˆMà ð 6Ø =°Ô#3Ð"5Ñ5�øð  ð 	2Ø -Ð!1Ñ1Ðàð 	Ýð ð Ø)¨?Ð<MÈ~Ð^ðñ ô ñ ô ð õ 'Ø+Ø+Ø+Ø%ð	
ñ 
ô 
ð 	
r*   )NNNNNNNN)r!   r"   r#   r$   r   r9   r%   r&   rY   r
   rª   r¬   r   r©   r®   r¯   s   @r+   r�  r�    s  ø€ € € € € ðð ð	˜zð 	ð 	ð 	ð 	ð 	ð 	ð 26Ø.2Ø04Ø(,Ø!%Ø)-Ø,0Ø#'ð
ð 
àÔ'¨$Ñ.ð
ð œ tÑ+ð
ð Ô&¨Ñ-ð	
ð
  ™ð
ð ˜$‘;ð
ð   $™;ð
ð # T™kð
ð ˜D‘[ð
ð 
Ð(Ñ	(ð
ð 
ð 
ð 
ð 
ð 
ð 
ð 
r*   r�  c                   ó.   ‡ — e Zd ZdZdefˆ fd„Zd„ Zˆ xZS )ÚMimiDecoderzSEANet decoder as used by Mimi.rz   c           	      ó’  •— t          ¦   «                              ¦   «          t          dt          |j        ¦  «        z  ¦  «        }t          ||j        ||j        z  |j        ¦  «        g}|j        D ]}||j        z  }|t          j
        ¦   «         gz  }|t          |||dz  |dz  |¬¦  «        gz  }t          |j        ¦  «        D ]$}|t          ||dz  |j        |z  df¦  «        gz  }Œ%|dz  }Œ€|t          j
        ¦   «         gz  }|t          ||j        |j        |j        ¦  «        gz  }t          j        |¦  «        | _        d S )Nrm   rÖ   r   )rn   r9   rV   r3   rÚ   r_   rà   rØ   ra   r   rÂ   r±   rÛ   rÜ   r»   rÞ   r×   rá   rÃ   râ   )r7   rz   rç   ræ   ré   rê   rë   r{   s          €r+   r9   zMimiDecoder.__init__®  sf  ø€ Ý‰Œ×ÒÑÔÐÝ�a�3˜vÔ7Ñ8Ô8Ñ8Ñ9Ô9ˆÝ˜F FÔ$6¸À&ÔBTÑ8TÐV\ÔVhÑiÔiÐjˆð Ô-ð 
	ð 
	ˆEØ# fÔ&8Ñ8ˆMà•b”f‘h”h�ZÑˆEØÝ# F¨M¸=ÈAÑ;MÐ[`ÐcdÑ[dÐmrÐsÑsÔsðñ ˆEõ ˜6Ô5Ñ6Ô6ð lð l�Ø�/¨&°-À1Ñ2DÀvÔGbÐdeÑGeÐghÐFiÑjÔjÐkÑk��Ø˜‰MˆGˆGð 	•"”&‘(”(�ÑˆØ•*˜V VÔ%7¸Ô9NÐPVÔPgÑhÔhÐiÑiˆÝ”m EÑ*Ô*ˆŒˆˆr*   c                 ó0   — | j         D ]} ||¦  «        }Œ|S r‡   )râ   )r7   r:   rÒ   s      r+   r©   zMimiDecoder.forwardÆ  s*   € Ø”[ð 	1ð 	1ˆEØ!˜E -Ñ0Ô0ˆMˆMØÐr*   rï   r¯   s   @r+   rœ  rœ  «  sY   ø€ € € € € Ø)Ð)ð+˜zð +ð +ð +ð +ð +ð +ð0ð ð ð ð ð ð r*   rœ  c                   óf   ‡ — e Zd ZdZddedefˆ fd„Zedej	        fd„¦   «         Z
d„ Zd	„ Zd
„ Zˆ xZS )ÚMimiEuclideanCodebookz!Codebook with Euclidean distance.çñhãˆµøä>rz   Úepsilonc                 óª  •— t          ¦   «                              ¦   «          t          j        |j        |j        ¦  «        }|j        | _        |                      dt          j        dgt          j        ¬¦  «        ¦  «         |                      dt          j	        |j        ¦  «        ¦  «         |                      d|¦  «         d | _
        || _        d S )NÚinitializedTri   Úcluster_usageÚ	embed_sum)rn   r9   r%   rC   Úcodebook_sizeÚcodebook_dimrw   ru   rP  rD   Ú_embedr¢  )r7   rz   r¢  Úembedr{   s       €r+   r9   zMimiEuclideanCodebook.__init__Ï  s¯   ø€ Ý‰Œ×ÒÑÔÐÝ”˜FÔ0°&Ô2EÑFÔFˆà#Ô1ˆÔà×Ò˜]­E¬L¸$¸ÅuÄ}Ð,UÑ,UÔ,UÑVÔVÐVØ×Ò˜_­e¬j¸Ô9MÑ.NÔ.NÑOÔOÐOØ×Ò˜[¨%Ñ0Ô0Ð0ØˆŒØˆŒˆˆr*   rŒ   c                 óŒ   — | j         €7| j        | j                             | j        ¬¦  «        d d …d f         z  | _         | j         S )N)Úmin)r©  r¦  r¥  Úclampr¢  rŠ   s    r+   rª  zMimiEuclideanCodebook.embedÛ  sH   € àŒ;ÐØœ.¨4Ô+=×+CÒ+CÈÌÐ+CÑ+UÔ+UÐVWÐVWÐVWÐY]ÐV]Ô+^Ñ^ˆDŒKØŒ{Ðr*   c                 óÖ   — t          j        |d                               ¦   «         | j        d                               ¦   «         d¬¦  «        d         }|                     d¬¦  «        }|S )Nrm   )rK  r   rM   rN   )r%   Úcdistr­   rª  Úargmin)r7   r:   ÚdistsÚ	embed_inds       r+   ÚquantizezMimiEuclideanCodebook.quantizeá  s^   € õ ”˜M¨$Ô/×5Ò5Ñ7Ô7¸¼ÀDÔ9I×9OÒ9OÑ9QÔ9QÐUVÐWÑWÔWÐXYÔZˆØ—L’L R�LÑ(Ô(ˆ	ØÐr*   c                 óœ   — |j         }|                     d|d         f¦  «        }|                      |¦  «        } |j        |d d…         Ž }|S )NrM   )rB   r?  r³  rh  )r7   r:   rB   r²  s       r+   ÚencodezMimiEuclideanCodebook.encodeé  sR   € ØÔ#ˆà%×-Ò-¨r°5¸´9¨oÑ>Ô>ˆà—M’M -Ñ0Ô0ˆ	à"�I”N E¨#¨2¨#¤JÐ/ˆ	ØÐr*   c                 óP   — t           j                             || j        ¦  «        }|S r‡   )r   r›   Ú	embeddingrª  ©r7   r²  r³  s      r+   ÚdecodezMimiEuclideanCodebook.decodeô  s    € Ý”=×*Ò*¨9°d´jÑAÔAˆØˆr*   )r¡  )r!   r"   r#   r$   r   r­   r9   Úpropertyr%   rY   rª  r³  rµ  r¹  r®   r¯   s   @r+   r   r   Ì  s¬   ø€ € € € € Ø+Ð+ð
ð 
˜zð 
°Eð 
ð 
ð 
ð 
ð 
ð 
ð ð�u”|ð ð ð ñ „Xðð
ð ð ðð ð ðð ð ð ð ð ð r*   r   c                   ó4   ‡ — e Zd ZdZdefˆ fd„Zd„ Zd„ Zˆ xZS )ÚMimiVectorQuantizationzY
    Vector quantization implementation. Currently supports only euclidean distance.
    rz   c                 óp   •— t          ¦   «                              ¦   «          t          |¦  «        | _        d S r‡   )rn   r9   r   Úcodebookr;  s     €r+   r9   zMimiVectorQuantization.__init__ÿ  s,   ø€ Ý‰Œ×ÒÑÔÐÝ-¨fÑ5Ô5ˆŒˆˆr*   c                 óh   — |                      ddd¦  «        }| j                             |¦  «        }|S ©Nr   rm   r   )Úpermuter¾  rµ  )r7   r:   Úembed_ins      r+   rµ  zMimiVectorQuantization.encode  s3   € Ø%×-Ò-¨a°°AÑ6Ô6ˆØ”=×'Ò'¨Ñ6Ô6ˆØˆr*   c                 óh   — | j                              |¦  «        }|                     ddd¦  «        }|S rÀ  )r¾  r¹  rÁ  r¸  s      r+   r¹  zMimiVectorQuantization.decode  s3   € Ø”=×'Ò'¨	Ñ2Ô2ˆØ×#Ò# A q¨!Ñ,Ô,ˆØˆr*   )	r!   r"   r#   r$   r   r9   rµ  r¹  r®   r¯   s   @r+   r¼  r¼  ú  sl   ø€ € € € € ðð ð6˜zð 6ð 6ð 6ð 6ð 6ð 6ðð ð ð
ð ð ð ð ð ð r*   r¼  c                   óˆ   ‡ — e Zd ZdZddededz  fˆ fd„Zddej        dedz  dej        fd„Z	d	ej        dej        fd
„Z
ˆ xZS )ÚMimiResidualVectorQuantizerzResidual Vector Quantizer.Nrz   Únum_quantizersc                 ó  •‡— t          ¦   «                              ¦   «          ‰j        | _        ‰j        | _        |�|n‰j        | _        t          j        ˆfd„t          | j        ¦  «        D ¦   «         ¦  «        | _        d | _	        d | _
        ‰j        ‰j        k    rft          j                             ‰j        ‰j        dd¬¦  «        | _	        t          j                             ‰j        ‰j        dd¬¦  «        | _
        d S d S )Nc                 ó.   •— g | ]}t          ‰¦  «        ‘ŒS r)   )r¼  )r„  Ú_rz   s     €r+   r…  z8MimiResidualVectorQuantizer.__init__.<locals>.<listcomp>  s"   ø€ Ð$hÐ$hÐ$hÈÕ%;¸FÑ%CÔ%CÐ$hÐ$hÐ$hr*   r   Fr4  )rn   r9   r§  Ú
frame_raterÆ  r   rÃ   rÛ   râ   Ú
input_projÚoutput_projÚ$vector_quantization_hidden_dimensionrà   r%   rs   )r7   rz   rÆ  r{   s    ` €r+   r9   z$MimiResidualVectorQuantizer.__init__  s  øø€ Ý‰Œ×ÒÑÔÐØ#Ô1ˆÔØ Ô+ˆŒØ0>Ð0J˜n˜nÐPVÔPeˆÔÝ”mÐ$hÐ$hÐ$hÐ$hÍUÐSWÔSfÑMgÔMgÐ$hÑ$hÔ$hÑiÔiˆŒàˆŒØˆÔØÔ6¸&Ô:LÒLÐLÝ#œhŸošoØÔ" FÔ$OÐQRÐY^ð .ñ ô ˆDŒOõ  %œxŸšØÔ;¸VÔ=OÐQRÐY^ð  /ñ  ô  ˆDÔÐÐð	 MÐLr*   Ú
embeddingsrŒ   c                 ó0  — | j         �|                       |¦  «        }|�|n| j        }|}g }| j        d|…         D ]F}|                     |¦  «        }|                     |¦  «        }||z
  }|                     |¦  «         ŒGt          j        |¦  «        }|S )úñ
        Encode a given input tensor with the specified frame rate at the given number of quantizers / codebooks. The RVQ encode method sets
        the appropriate number of quantizers to use and returns indices for each quantizer.
        N)rË  rÆ  râ   rµ  r¹  rß   r%   Ústack)	r7   rÎ  rÆ  rÑ   Úall_indicesrÒ   ÚindicesÚ	quantizedÚout_indicess	            r+   rµ  z"MimiResidualVectorQuantizer.encode"  sª   € ð
 Œ?Ð&ØŸš¨Ñ4Ô4ˆJà+9Ð+E˜˜È4ÔK^ˆàˆØˆØ”[  . Ô1ð 	(ð 	(ˆEØ—l’l 8Ñ,Ô,ˆGØŸš WÑ-Ô-ˆIØ )Ñ+ˆHØ×Ò˜wÑ'Ô'Ð'Ð'Ý”k +Ñ.Ô.ˆØÐr*   Úcodesc                 ó  — t          j        d|j        ¬¦  «        }|                     dd¦  «        }t	          |¦  «        D ],\  }}| j        |         }|                     |¦  «        }||z   }Œ-| j        �|                      |¦  «        }|S )zJDecode the given codes of shape [B, K, T] to the quantized representation.r•   rŒ  r   r   )r%   ru   r?   r  rÀ   râ   r¹  rÌ  )r7   rÖ  Úquantized_outrÊ   rÓ  rÒ   rÔ  s          r+   r¹  z"MimiResidualVectorQuantizer.decode6  s�   € åœ S°´Ð>Ñ>Ô>ˆØ—’  1Ñ%Ô%ˆÝ# EÑ*Ô*ð 	6ð 	6‰JˆAˆwØ”K ”NˆEØŸš WÑ-Ô-ˆIØ)¨IÑ5ˆMˆMàÔÐ'Ø ×,Ò,¨]Ñ;Ô;ˆMØÐr*   r‡   )r!   r"   r#   r$   r   rV   r9   r%   rY   rµ  r¹  r®   r¯   s   @r+   rÅ  rÅ    s¸   ø€ € € € € Ø$Ð$ðð ˜zð ¸3À¹:ð ð ð ð ð ð ð"ð  ¤ð ¸sÀT¹zð ÐUZÔUað ð ð ð ð(˜EœLð ¨U¬\ð ð ð ð ð ð ð ð r*   rÅ  c                   ó|   ‡ — e Zd ZdZdefˆ fd„Zddej        dedz  dej        fd„Z	d	ej        dej        fd
„Z
ˆ xZS )Ú MimiSplitResidualVectorQuantizerz Split Residual Vector Quantizer.rz   c                 ó8  •— t          ¦   «                              ¦   «          |j        | _        |j        | _        |j        | _        |j        | _        |j        |j        z
  | _        t          || j        ¦  «        | _	        t          || j        ¦  «        | _
        d S r‡   )rn   r9   r§  rÊ  rÆ  Úmax_num_quantizersÚnum_semantic_quantizersÚnum_acoustic_quantizersrÅ  Ú"semantic_residual_vector_quantizerÚ"acoustic_residual_vector_quantizerr;  s     €r+   r9   z)MimiSplitResidualVectorQuantizer.__init__G  s‰   ø€ Ý‰Œ×ÒÑÔÐØ#Ô1ˆÔØ Ô+ˆŒØ"(Ô"7ˆÔà'-Ô'EˆÔ$Ø'-Ô'<¸vÔ?]Ñ']ˆÔ$å2MÈfÐVZÔVrÑ2sÔ2sˆÔ/Ý2MÈfÐVZÔVrÑ2sÔ2sˆÔ/Ð/Ð/r*   NrÎ  rÆ  rŒ   c                 óv  — |€| j         n|}|| j         k    rt          d| j         › d|› d�¦  «        ‚|| j        k     rt          d| j        › d|› d�¦  «        ‚| j                             |¦  «        }|| j        k    r<| j                             ||| j        z
  ¬¦  «        }t          j        ||gd¬¦  «        }|S )	rÐ  NúcThe number of quantizers (i.e codebooks) asked should be lower than the total number of quantizers ú, but is currently ú.zgThe number of quantizers (i.e codebooks) asked should be higher than the number of semantic quantizers )rÆ  r   rN   )rÜ  r5   rÝ  rß  rµ  rà  r%   rQ   )r7   rÎ  rÆ  rÖ  Úacoustic_codess        r+   rµ  z'MimiSplitResidualVectorQuantizer.encodeS  s5  € ð 5CÐ4J˜Ô0Ð0ÐP^ˆà˜DÔ3Ò3Ð3Ýð tÐvzô  wNð  tð  tð  cqð  tð  tð  tñô ð ð ˜DÔ8Ò8Ð8Ýð }Ðz~ô  {Wð  }ð  }ð  lzð  }ð  }ð  }ñô ð ð
 Ô7×>Ò>¸zÑJÔJˆà˜DÔ8Ò8Ð8Ø!ÔD×KÒKØ¨>¸DÔ<XÑ+Xð Lñ ô ˆNõ ”I˜u nÐ5¸1Ð=Ñ=Ô=ˆEàˆr*   rÖ  c                 óä   — | j                              |dd…d| j        …f         ¦  «        }|j        d         | j        k    r.|| j                             |dd…| j        d…f         ¦  «        z  }|S )z7Decode the given codes to the quantized representation.Nr   )rß  r¹  rÝ  rB   rà  )r7   rÖ  rØ  s      r+   r¹  z'MimiSplitResidualVectorQuantizer.decodep  s„   € ð Ô?×FÒFÀuÈQÈQÈQÐPnÐRVÔRnÐPnÐMnÔGoÑpÔpˆð Œ;�qŒ>˜DÔ8Ò8Ð8Ø˜TÔD×KÒKÈEÐRSÐRSÐRSÐUYÔUqÐUsÐUsÐRsÔLtÑuÔuÑuˆMØÐr*   r‡   )r!   r"   r#   r$   r   r9   r%   rY   r­   rµ  r¹  r®   r¯   s   @r+   rÚ  rÚ  D  s¯   ø€ € € € € Ø*Ð*ð
t˜zð 
tð 
tð 
tð 
tð 
tð 
tðð  ¤ð ¸uÀt¹|ð ÐW\ÔWcð ð ð ð ð:	˜EœLð 	¨U¬\ð 	ð 	ð 	ð 	ð 	ð 	ð 	ð 	r*   rÚ  c                   ó†   ‡ — e Zd ZU eed<   dZdZdZdZddgZ	dgZ
dZdZdZdZdZ ej        ¦   «         ˆ fd	„¦   «         Zˆ xZS )
ÚMimiPreTrainedModelrz   ÚmimiÚinput_valuesÚaudioTrÚ  rp  re  c                 ó  •— t          ¦   «                              |¦  «         t          |t          j        t          j        f¦  «        rpt          j        |j        ¦  «         |j	        �Nt          j        |j        |j        |j        d         z  z  ¦  «        }t          j        |j	        | |¬¦  «         dS dS t          |t           ¦  «        r&t          j        |j        | j        j        ¦  «         dS t          |t*          ¦  «        r”|j        j        d         }|j        j        d         }|j        j        d         }|dz
  |z  dz   }t          j        |j        |¦  «         t          j        |j        |¦  «         t          j        |j        ||z
  ¦  «         dS t          |t4          ¦  «        rMt          j        |j        ¦  «         t          j        |j        ¦  «         t          j        |j        ¦  «         dS dS )zInitialize the weightsNr   )ÚaÚbr   ) rn   Ú_init_weightsrÐ   r   rs   rµ   ÚinitÚkaiming_normal_Úweightrf   r¶   r]  rd   rI   ra   Úuniform_rñ   Ú	constant_r÷   rz   rô   r_   rt   rb   rc   rl   r   Úones_r¤  r¥  Úzeros_r¦  )r7   rD  r,  ra   rb   rc   r{   s         €r+   rï  z!MimiPreTrainedModel._init_weightsŒ  sÎ  ø€ õ 	‰Œ×Ò˜fÑ%Ô%Ð%Ý�f�rœy­"Ô*<Ð=Ñ>Ô>ð 	*ÝÔ  ¤Ñ/Ô/Ð/ØŒ{Ð&Ý”I˜fœm¨vÔ/AÀFÔDVÐWXÔDYÑ/YÑZÑ[Ô[�Ý”˜fœk¨a¨R°1Ð5Ñ5Ô5Ð5Ð5Ð5ð 'Ð&õ ˜¥Ñ/Ô/ð 	*ÝŒN˜6œ<¨¬Ô)NÑOÔOÐOÐOÐOÝ˜¥
Ñ+Ô+ð 	*Ø œ+Ô1°!Ô4ˆKØ”[Ô'¨Ô*ˆFØ”{Ô+¨AÔ.ˆHØ&¨™?¨hÑ6¸Ñ:ˆKÝŒN˜6œ=¨&Ñ1Ô1Ð1ÝŒN˜6Ô-¨{Ñ;Ô;Ð;ÝŒN˜6Ô/°¸vÑ1EÑFÔFÐFÐFÐFÝ˜Õ 5Ñ6Ô6ð 	*ÝŒJ�vÔ)Ñ*Ô*Ð*ÝŒJ�vÔ+Ñ,Ô,Ð,ÝŒK˜Ô(Ñ)Ô)Ð)Ð)Ð)ð	*ð 	*r*   )r!   r"   r#   r   r'   Úbase_model_prefixÚmain_input_nameÚinput_modalitiesÚsupports_gradient_checkpointingÚ_no_split_modulesÚ_skip_keys_device_placementÚ_supports_flash_attnÚ_supports_sdpaÚ_supports_flex_attnÚ_supports_attention_backendÚ_can_compile_fullgraphr%   r$  rï  r®   r¯   s   @r+   rè  rè  |  s¡   ø€ € € € € € àÐÐÑØÐØ$€OØÐØ&*Ð#Ø;Ð=SÐTÐØ#4Ð"5ÐØÐØ€NØÐØ"&Ðà!Ðà€U„]�_„_ð*ð *ð *ð *ñ „_ð*ð *ð *ð *ð *r*   rè  z,
    The Mimi neural audio codec model.
    )Úcustom_introc                   óä  ‡ — e Zd Zdefˆ fd„Z	 	 	 	 ddej        dedededz  de	dz  d	e
dz  d
e
dz  deej        ej        dz  f         fd„Zdej        dej        fd„Zddej        defd„Z	 	 	 	 	 	 ddej        dej        dz  dedz  dedz  de	dz  d	e
dz  d
e
dz  deej        ej        dz  f         ez  fd„Z	 	 ddej        dedz  d
e
dz  dej        fd„Z	 	 	 ddej        dej        dz  dedz  d
e
dz  deej        ej        f         ez  f
d„Ze	 	 	 	 	 	 ddej        dej        dz  dedz  dej        dz  dedz  dedz  d
e
dz  deej        ej        f         ez  fd„¦   «         Zˆ xZS )Ú	MimiModelrz   c                 ót  •— t          ¦   «                              |¦  «         || _        t          |¦  «        | _        t          |¦  «        | _        d | _        d | _        |j	        |j
        k    r¡t          ||j        |j        dt          |j
        |j	        z  ¦  «        z  dddt          | j        j        ¦  «        ¬¦  «        | _        t!          ||j        |j        dt          |j
        |j	        z  ¦  «        z  dd|j        ¬¦  «        | _        t          |¦  «        | _        t'          |¦  «        | _        t+          |¦  «        | _        t          t/          j        | j        j        ¦  «        ¦  «        | _        d| j        z  | j        j        k    rt7          d¦  «        ‚|                      ¦   «          d S )Nrm   FrA   )ra   rb   rf   re   r;   )ra   rb   rf   rd   z'The codebook_size must be a power of 2.)rn   r9   rz   rÔ   Úencoderr�  Úencoder_transformerÚ
downsampleÚupsamplerÊ  Úencodec_frame_rater_   rà   rV   r3   rã   r±   Úupsample_groupsÚdecoder_transformerrœ  ÚdecoderrÚ  Ú	quantizerr¶   Úlog2r§  Úbits_per_codebookr5   Ú	post_initr;  s     €r+   r9   zMimiModel.__init__«  s•  ø€ Ý‰Œ×Ò˜Ñ Ô Ð ØˆŒå" 6Ñ*Ô*ˆŒÝ#7¸Ñ#?Ô#?ˆÔ àˆŒØˆŒØÔ Ô 9Ò9Ð9Ý(ØØÔ"ØÔ"Ø¥ FÔ$=ÀÔ@QÑ$QÑ RÔ RÑRØØØ$Ý˜dœlÔBÑCÔCð	ñ 	ô 	ˆDŒOõ 0ØØÔ"ØÔ"Ø¥ FÔ$=ÀÔ@QÑ$QÑ RÔ RÑRØØØÔ-ðñ ô ˆDŒMõ $8¸Ñ#?Ô#?ˆÔ Ý" 6Ñ*Ô*ˆŒå9¸&ÑAÔAˆŒå!$¥T¤Y¨t¬{Ô/HÑ%IÔ%IÑ!JÔ!JˆÔØˆdÔ$Ñ$¨¬Ô(AÒAÐAÝÐFÑGÔGÐGð 	�ŠÑÔÐÐÐr*   Nrê  rÆ  Úpadding_maskre  r6   Úuse_streamingrŠ  rŒ   c                 óÂ  — |                       ||¬¦  «        }|                      |                     dd¦  «        |||¬¦  «        }	|r|	                     d¦  «        }nt	          |	¦  «        dk    r|	d         }|	d                              dd¦  «        }|                      ||¬¦  «        }| j                             ||¦  «        }
|
                     dd¦  «        }
|
||fS )z€
        Encodes the given input using the underlying VQVAE. The padding mask is required to compute the correct scale.
        rÏ   r   rm   )re  r|  rŠ  re  r   )r  r  r  Úgetr3   r  r  rµ  )r7   rê  rÆ  r  re  r6   r  rŠ  rÎ  Úencoder_outputsrÖ  s              r+   Ú_encode_framezMimiModel._encode_frameÖ  sô   € ð —\’\ ,¸m�\ÑLÔLˆ
ð ×2Ò2Ø× Ò   AÑ&Ô&Ø+Ø#Ø#ð	 3ñ 
ô 
ˆð ð 	1Ø-×1Ò1Ð2CÑDÔDˆOˆOÝ�Ñ!Ô! AÒ%Ð%Ø-¨aÔ0ˆOØ$ QÔ'×1Ò1°!°QÑ7Ô7ˆ
Ø—_’_ Z¸}�_ÑMÔMˆ
à”×%Ò% j°.ÑAÔAˆØ—’  1Ñ%Ô%ˆØ�o }Ð4Ð4r*   r¢   c                 ó¶   — |}| j         j        D ]/}| j                              |¦  «                             |¦  «        }Œ0| j                             |¦  «        }|S )zL
        Return the number of frames of the encoded audio waveform.
        )r  rã   rä   r¦   r  )r7   r¢   r¥   Ú
layer_names       r+   Úget_encoded_lengthzMimiModel.get_encoded_lengthù  sd   € ð %ˆð œ,Ô>ð 	eð 	eˆJØ œL×6Ò6°zÑBÔB×UÒUÐVcÑdÔdˆMˆMð œ×:Ò:¸=ÑIÔIˆàÐr*   ÚrightÚpadding_sidec                 ó”  — |                       |                     d¬¦  «        ¦  «        }t          j        |                     ¦   «         |j        ¬¦  «                             t          |¦  «        d¦  «        }||                     d¦  «        k     }| 	                    |j        ¦  «        }|dk    r|S | 
                    dg¬¦  «        S )zR
        Get the mask for the audio codes from the original padding mask.
        rM   rN   rŒ  r   r  )Údims)r  Úsumr%   r  rP   r?   r  r3   r*  r�   Úflip)r7   r  r  Úencoded_lengthsÚaudio_codes_masks        r+   Úget_audio_codes_maskzMimiModel.get_audio_codes_mask  sÅ   € ð ×1Ò1°,×2BÒ2BÀrÐ2BÑ2JÔ2JÑKÔKˆå œ<¨×(;Ò(;Ñ(=Ô(=ÀoÔF\Ð]Ñ]Ô]×dÒdÝ�Ñ Ô  "ñ
ô 
Ðð ,¨o×.GÒ.GÈÑ.JÔ.JÒJÐØ+×.Ò.¨|Ô/BÑCÔCÐà˜7Ò"Ð"Ø#Ð#à#×(Ò(¨r¨dÐ(Ñ3Ô3Ð3r*   r   c           	      ón  — |�|n| j         j        }|�|n| j         j        }|€| j         j        n|}|| j         j        k    r t	          d| j         j        › d|› d�¦  «        ‚|j        \  }}	}
|	dk     s|	dk    rt	          d|	› �¦  «        ‚|€&t          j        |¦  «                             ¦   «         }|�r8|�€5g g g }}}| j	        j
        D ]˜}|                     | j	                             |¦  «        j        ¦  «         |                     | j	                             |¦  «        j        ¦  «         |                     | j	                             |¦  «        j        ¦  «         Œ™|                     | j        j        ¦  «         |                     | j        j        ¦  «         |                     | j        j        ¦  «         t#          t%          | j	        j
        ¦  «        dz   |||¬¦  «        }|                      |||                     ¦   «         ||||¬	¦  «        \  }}}|s|||fS t)          |||¦  «        S )
aE  
        Encodes the input audio waveform into discrete codes.

        Args:
            input_values (`torch.Tensor` of shape `(batch_size, channels, sequence_length)`):
                Float values of the input audio waveform.
            padding_mask (`torch.Tensor` of shape `(batch_size, channels, sequence_length)`):
                Indicates which inputs are to be ignored due to padding, where elements are either 1 for *not masked* or 0
                for *masked*.
            num_quantizers (`int`, *optional*):
                Number of quantizers (i.e codebooks) to use. By default, all quantizers are used.
            encoder_past_key_values (`Cache`, *optional*):
                Pre-computed hidden-states (key and values in the self-attention blocks) that can be used to speed up sequential decoding of the encoder transformer.
                This typically consists in the `past_key_values` returned by the model at a previous stage of decoding, when `use_cache=True` or `config.use_cache=True`.

                The model will output the same cache format that is fed as input.

                If `past_key_values` are used, the user can optionally input only the last `audio_values` or `audio_codes (those that don't
                have their past key value states given to this model).
            return_dict (`bool`, *optional*):
                Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.

        Returns:
            `codebook` of shape `[batch_size, num_codebooks, frames]`, the discrete encoded codes for the input audio waveform.
        Nrâ  rã  rä  r   rm   z1Number of audio channels must be 1 or 2, but got )r.   r/   r0   r1   )re  r6   r  rŠ  )rz   rŠ  r  rÆ  r5   rB   r%   Ú	ones_likerª   r  rã   rß   rä   rl   re   rI   r  r-   r3   r  r[   )r7   rê  r  rÆ  r   r6   r  rŠ  rÉ  rø   r¢   r/   r0   r1   r  Úencoded_framess                   r+   rµ  zMimiModel.encode  s¢  € ðF &1Ð%<�k�kÀ$Ä+ÔBYˆØ)6Ð)B˜˜ÈÌÔHaˆà7EÐ7M˜œÔ3Ð3ÐSaˆà˜DœKÔ6Ò6Ð6Ýð wÐvzô  wBô  wQð  wð  wð  ftð  wð  wð  wñô ð ð %1Ô$6Ñ!ˆˆ8�\à�aŠ<ˆ<˜8 aš<˜<ÝÐ[ÐQYÐ[Ð[Ñ\Ô\Ð\àÐÝ œ?¨<Ñ8Ô8×=Ò=Ñ?Ô?ˆLàñ 	˜]Ñ2ØOQÐSUÐWYÐ7LÐ5ÐØ"œlÔBð að a�
Ø!×(Ò(¨¬×)CÒ)CÀJÑ)OÔ)OÔ)]Ñ^Ô^Ð^Ø&×-Ò-¨d¬l×.HÒ.HÈÑ.TÔ.TÔ.]Ñ^Ô^Ð^Ø%×,Ò,¨T¬\×-GÒ-GÈ
Ñ-SÔ-SÔ-_Ñ`Ô`Ð`Ð`ð ×$Ò$ T¤_Ô%BÑCÔCÐCØ"×)Ò)¨$¬/Ô*BÑCÔCÐCØ!×(Ò(¨¬Ô)DÑEÔEÐEå2Ý˜tœ|ÔCÑDÔDÀqÑHØ"3Ø'=Ø&;ð	ñ ô ˆMð BF×ASÒASØØØ×ÒÑÔØ3Ø'Ø'Ø#ð BTñ B
ô B
Ñ>ˆÐ/°ð ð 	àØ'Øðð õ ! Ð1HÈ-ÑXÔXÐXr*   rÖ  c                 óˆ  — | j                              |¦  «        }|                      |¦  «        }|                      |                     dd¦  «        ||¬¦  «        }|r|                     d¦  «        }nt          |¦  «        dk    r|d         }|d                              dd¦  «        }|                      |¦  «        }||fS )Nr   rm   ©re  rŠ  re  r   )r  r¹  r	  r  r  r  r3   r  )r7   rÖ  re  rŠ  rÎ  Údecoder_outputsr  s          r+   Ú_decode_framezMimiModel._decode_framet  sÏ   € ð ”^×*Ò*¨5Ñ1Ô1ˆ
à—]’] :Ñ.Ô.ˆ
Ø×2Ò2Ø× Ò   AÑ&Ô&¸ÐU`ð 3ñ 
ô 
ˆð ð 	1Ø-×1Ò1Ð2CÑDÔDˆOˆOÝ�Ñ!Ô! AÒ%Ð%Ø-¨aÔ0ˆOØ$ QÔ'×1Ò1°!°QÑ7Ô7ˆ
Ø—,’,˜zÑ*Ô*ˆØ˜Ð'Ð'r*   r   r    c                 óî   — |�|n| j         j        }|                      |||¬¦  «        \  }}|�3|j        d         |j        d         k     r|dd|j        d         …f         }|s||fS t	          ||¦  «        S )aÊ  
        Decodes the given frames into an output audio waveform.

        Note that the output might be a bit bigger than the input. In that case, any extra steps at the end can be
        trimmed.

        Args:
            audio_codes (`torch.LongTensor`  of shape `(batch_size, num_quantizers, codes_length)`, *optional*):
                Discrete code embeddings computed using `model.encode`.
            padding_mask (`torch.Tensor` of shape `(batch_size, channels, sequence_length)`):
                Indicates which inputs are to be ignored due to padding, where elements are either 1 for *not masked* or 0
                for *masked*.
            decoder_past_key_values (`Cache`, *optional*):
                Pre-computed hidden-states (key and values in the self-attention blocks) that can be used to speed up sequential decoding of the decoder transformer.
                This typically consists in the `past_key_values` returned by the model at a previous stage of decoding, when `use_cache=True` or `config.use_cache=True`.

                The model will output the same cache format that is fed as input.

                If `past_key_values` are used, the user can optionally input only the last `audio_values` or `audio_codes (those that don't
                have their past key value states given to this model).
            return_dict (`bool`, *optional*):
                Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.

        Nr(  rM   .)rz   rŠ  r*  rB   r]   )r7   r   r  r    rŠ  r   s         r+   r¹  zMimiModel.decodeˆ  s«   € ð> &1Ð%<�k�kÀ$Ä+ÔBYˆà04×0BÒ0BØÐ)@Èkð 1Cñ 1
ô 1
Ñ-ˆÐ-ð
 Ð#¨Ô(:¸2Ô(>ÀÔASÐTVÔAWÒ(WÐ(WØ'¨Ð-E¨|Ô/AÀ"Ô/EÐ-EÐ(EÔFˆLàð 	àØ'ðð õ ! Ð/FÑGÔGÐGr*   c                 óþ  — |�|n| j         j        }|€&t          j        |¦  «                             ¦   «         }|€U|                      |||||¬¦  «        }	|	d         }|r|	                     d¦  «        }nt          |	¦  «        dk    r|	d         }|                      ||||¬¦  «        }
|
d         }|r|
                     d¦  «        }nt          |
¦  «        dk    r|
d         }|s||||fS t          ||||¬¦  «        S )aº
  
        input_values (`torch.FloatTensor` of shape `(batch_size, channels, sequence_length)`, *optional*):
            Raw audio input converted to Float.
        padding_mask (`torch.Tensor` of shape `(batch_size, sequence_length)`, *optional*):
            Indicates which inputs are to be ignored due to padding, where elements are either 1 for *not masked* or 0
            for *masked*.
        num_quantizers (`int`, *optional*):
            Number of quantizers (i.e codebooks) to use. By default, all quantizers are used.
        audio_codes (`torch.LongTensor`  of shape `(batch_size, num_quantizers, codes_length)`, *optional*):
            Discrete code embeddings computed using `model.encode`.
        encoder_past_key_values (`Cache`, *optional*):
            Pre-computed hidden-states (key and values in the self-attention blocks) that can be used to speed up sequential decoding of the encoder transformer.
            This typically consists in the `past_key_values` returned by the model at a previous stage of decoding, when `use_cache=True` or `config.use_cache=True`.

            The model will output the same cache format that is fed as input.

            If `past_key_values` are used, the user can optionally input only the last `audio_values` or `audio_codes (those that don't
            have their past key value states given to this model).
        decoder_past_key_values (`Cache`, *optional*):
            Pre-computed hidden-states (key and values in the self-attention blocks) that can be used to speed up sequential decoding of the decoder transformer.
            This typically consists in the `past_key_values` returned by the model at a previous stage of decoding, when `use_cache=True` or `config.use_cache=True`.

            The model will output the same cache format that is fed as input.

            If `past_key_values` are used, the user can optionally input only the last `audio_values` or `audio_codes (those that don't
            have their past key value states given to this model).

        Examples:

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

        >>> dataset = load_dataset("hf-internal-testing/ashraq-esc50-1-dog-example")
        >>> audio_sample = dataset["train"]["audio"][0]["array"]

        >>> model_id = "kyutai/mimi"
        >>> model = MimiModel.from_pretrained(model_id)
        >>> feature_extractor = AutoFeatureExtractor.from_pretrained(model_id)

        >>> inputs = feature_extractor(raw_audio=audio_sample, return_tensors="pt")

        >>> outputs = model(**inputs)
        >>> audio_codes = outputs.audio_codes
        >>> audio_values = outputs.audio_values
        ```N)rŠ  r   re  r   )r   r   r   r    )
rz   rŠ  r%   r%  rª   rµ  r  r3   r¹  r   )r7   rê  r  rÆ  r   r   r    rŠ  rI  r  r)  r   s               r+   r©   zMimiModel.forward¸  sN  € ðt &1Ð%<�k�kÀ$Ä+ÔBYˆàÐÝ œ?¨<Ñ8Ô8×=Ò=Ñ?Ô?ˆLàÐØ"ŸkškØ˜l¨NÐ<SÐalð *ñ ô ˆOð *¨!Ô,ˆKØð =Ø*9×*=Ò*=Ð>OÑ*PÔ*PÐ'Ð'Ý�_Ñ%Ô%¨Ò)Ð)Ø*9¸!Ô*<Ð'àŸ+š+ k°<ÐAXÐfq˜+ÑrÔrˆØ& qÔ)ˆØð 	9Ø&5×&9Ò&9Ð:KÑ&LÔ&LÐ#Ð#Ý�Ñ!Ô! AÒ%Ð%Ø&5°aÔ&8Ð#àð 	aØ Ð/FÐH_Ð`Ð`åØ#Ø%Ø$;Ø$;ð	
ñ 
ô 
ð 	
r*   )NNNN)r  )NNNNNN)NNr#  )r!   r"   r#   r   r9   r%   rY   rV   r
   r-   rª   r¬   r  r&   r  rX   r#  r­   r[   rµ  r*  r]   r¹  r   r   r©   r®   r¯   s   @r+   r  r  ¥  s¡  ø€ € € € € ð)˜zð )ð )ð )ð )ð )ð )ð` )-Ø7;Ø%)Ø#'ð!5ð !5à”lð!5ð ð!5ð ð	!5ð
  ™ð!5ð .°Ñ4ð!5ð ˜d‘{ð!5ð ˜D‘[ð!5ð 
ˆuŒ|˜Uœ\¨DÑ0Ð0Ô	1ð!5ð !5ð !5ð !5ðF¨uÔ/?ð ÀEÔDTð ð ð ð ð4ð 4°´ð 4ÈSð 4ð 4ð 4ð 4ð( -1Ø'+Ø04Ø7;Ø%)Ø#'ðYYð YYà”lðYYð ”l TÑ)ðYYð  ™ð	YYð
 "'¨¡ðYYð .°Ñ4ðYYð ˜d‘{ðYYð ˜D‘[ðYYð 
ˆuŒ|˜Uœ\¨DÑ0Ð0Ô	1Ð4EÑ	EðYYð YYð YYð YYð| )-Ø#'ð	(ð (àŒ|ð(ð  ™ð(ð ˜D‘[ð	(ð
 
Œð(ð (ð (ð (ð. -1Ø04Ø#'ð.Hð .Hà”\ð.Hð ”l TÑ)ð.Hð "'¨¡ð	.Hð
 ˜D‘[ð.Hð 
ˆuŒ|˜Uœ\Ð)Ô	*Ð->Ñ	>ð.Hð .Hð .Hð .Hð` ð -1Ø%)Ø+/Ø04Ø04Ø#'ðW
ð W
à”lðW
ð ”l TÑ)ðW
ð ˜d™
ð	W
ð
 ”\ DÑ(ðW
ð "'¨¡ðW
ð "'¨¡ðW
ð ˜D‘[ðW
ð 
ˆuŒ|˜Uœ\Ð)Ô	*¨ZÑ	7ðW
ð W
ð W
ñ „^ðW
ð W
ð W
ð W
ð W
r*   r  )r   )r•   )Jr$   r¶   Úcollections.abcr   Údataclassesr   Útypingr   r%   r   Ú r   rð  Úactivationsr	   Úcache_utilsr
   r   Úmasking_utilsr   Úmodeling_layersr   Úmodeling_outputsr   Úmodeling_rope_utilsr   r   Úmodeling_utilsr   r   Úprocessing_utilsr   r€   r   r   r   r   Úutils.genericr   Úconfiguration_mimir   Ú
get_loggerr!   rq   r   r-   r[   r]   ÚModuler_   r±   r»   rÔ   rñ   rý   r(  r0  r2  rY   rV   rC  r­   rV  rX  rp  r�  rœ  r   r¼  rÅ  rÚ  rè  r  Ú__all__r)   r*   r+   ú<module>r>     s–  ðð Ð à €€€Ø $Ð $Ð $Ð $Ð $Ð $Ø !Ð !Ð !Ð !Ð !Ð !Ø Ð Ð Ð Ð Ð à €€€Ø Ð Ð Ð Ð Ð à &Ð &Ð &Ð &Ð &Ð &Ø !Ð !Ð !Ð !Ð !Ð !Ø .Ð .Ð .Ð .Ð .Ð .Ð .Ð .Ø >Ð >Ð >Ð >Ð >Ð >Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø 7Ð 7Ð 7Ð 7Ð 7Ð 7Ø KÐ KÐ KÐ KÐ KÐ KÐ KÐ KØ FÐ FÐ FÐ FÐ FÐ FÐ FÐ FØ &Ð &Ð &Ð &Ð &Ð &Ø MÐ MÐ MÐ MÐ MÐ MÐ MÐ MÐ MÐ MÐ MÐ MØ +Ð +Ð +Ð +Ð +Ð +Ø *Ð *Ð *Ð *Ð *Ð *ð 
ˆÔ	˜HÑ	%Ô	%€ð Ø
ð1ð 1ð 1ð 1ð 1�ñ 1ô 1ñ „ñ „ð1ð<[ð [ð [ð [ð [ñ [ô [ð [ð| Ø
ð8ð 8ð 8ð 8ð 8˜ñ 8ô 8ñ „ñ „ð8ð* Ø
ð1ð 1ð 1ð 1ð 1˜ñ 1ô 1ñ „ñ „ð1ð$Ið Ið Ið Ið I�”ñ Iô Ið IðX7ð 7ð 7ð 7ð 7˜"œ)ñ 7ô 7ð 7ðt'(ð '(ð '(ð '(ð '(�b”iñ '(ô '(ð '(ðT*ð *ð *ð *ð *�"”)ñ *ô *ð *ðZð ð ð ð �R”Yñ ô ð ð ><ð ><ð ><ð ><ð ><˜"œ)ñ ><ô ><ð ><ðD(ð (ð (ðð ð ð ð2ð ð ð ð ˆbŒiñ ô ð ð"	U˜Uœ\ð 	U°#ð 	U¸%¼,ð 	Uð 	Uð 	Uð 	Uð( ð%ð %ØŒIð%àŒ<ð%ð 
Œð%ð Œ<ð	%ð
 ”L 4Ñ'ð%ð ð%ð ð%ð Ð'Ô(ð%ð %ð %ð %ð2E)ð E)ð E)ð E)ð E)�B”Iñ E)ô E)ð E)ðP2ð 2ð 2ð 2ð 2Ð5ñ 2ô 2ð 2ðjR
ð R
ð R
ð R
ð R
˜2œ9ñ R
ô R
ð R
ðjð ð ð ð �"”)ñ ô ð ðB*ð *ð *ð *ð *˜BœIñ *ô *ð *ð\ð ð ð ð ˜RœYñ ô ð ð(3ð 3ð 3ð 3ð 3 "¤)ñ 3ô 3ð 3ðl5ð 5ð 5ð 5ð 5 r¤yñ 5ô 5ð 5ðp ð%*ð %*ð %*ð %*ð %*˜/ñ %*ô %*ñ „ð%*ðP €ððñ ô ð
f
ð f
ð f
ð f
ð f
Ð#ñ f
ô f
ñô ð
f
ðR Ð-Ð
.€€€r*   