Ë
    T^(h	@ ã                   óÎ  — d Z ddlZddlmZ ddlmZmZmZmZ ddl	Z	ddl
Z	ddl	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mZ ddl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#m$Z$ ddl%m&Z&  e«       rddlm'Z'  e#jP                  e)«      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j^                  «      Z0 G d„ dej^                  «      Z1 G d„ dej^                  «      Z2 G d„ d ej^                  «      Z3 G d!„ d"ej^                  «      Z4 G d#„ d$ej^                  «      Z5d%„ Z6dId&„Z7 G d'„ d(ej^                  «      Z8d)e	jr                  d*e:d+e	jr                  fd,„Z; G d-„ d.ej^                  «      Z< G d/„ d0e<«      Z= G d1„ d2e<«      Z>e<e=e>d3œZ? G d4„ d5ej^                  «      Z@ G d6„ d7ej^                  «      ZA G d8„ d9ej^                  «      ZB G d:„ d;ej^                  «      ZC G d<„ d=ej^                  «      ZD G d>„ d?ej^                  «      ZE G d@„ dAej^                  «      ZF G dB„ dCe«      ZGdDZHdEZI e!dFeH«       G dG„ dHeG«      «       ZJdHdCgZKy)JzPyTorch Mimi model.é    N)Ú	dataclass)ÚListÚOptionalÚTupleÚUnion)Únné   )ÚACT2FN)ÚCacheÚDynamicCacheÚSlidingWindowCacheÚStaticCache)ÚAttentionMaskConverter)Ú!flash_attn_supports_top_left_maskÚis_flash_attn_available)ÚBaseModelOutputWithPast)ÚROPE_INIT_FUNCTIONSÚdynamic_rope_update)ÚPreTrainedModel)ÚModelOutputÚadd_start_docstringsÚ%add_start_docstrings_to_model_forwardÚloggingÚreplace_return_docstringsé   )Ú
MimiConfig)Ú_flash_attention_forwardr   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eeej                     f      ed<   dZeeeeej                     f      ed<   y)Ú
MimiOutputaž  
    Args:
        audio_codes (`torch.LongTensor`  of shape `(batch_size, num_quantizers, codes_length)`, *optional*):
            Discret 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    r   ÚtorchÚ
LongTensorÚ__annotations__r!   ÚFloatTensorr"   r   r   r   r#   © ó    úd/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/mimi/modeling_mimi.pyr   r   4   s}   … ñð0 /3€K�˜%×*Ñ*Ñ+Ó2Ø04€L�(˜5×,Ñ,Ñ-Ó4ØOSÐ˜X e¨E°4¸×8IÑ8IÑ3JÐ,JÑ&KÑLÓSØOSÐ˜X e¨E°4¸×8IÑ8IÑ3JÐ,JÑ&KÑLÔSr-   r   c                   ór   — e Zd ZU dZdZeej                     ed<   dZ	ee
eeej                     f      ed<   y)ÚMimiEncoderOutputaY  
    Args:
        audio_codes (`torch.LongTensor`  of shape `(batch_size, num_quantizers, codes_length)`, *optional*):
            Discret 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).
    Nr    r"   )r$   r%   r&   r'   r    r   r(   r)   r*   r"   r   r   r   r+   r,   r-   r.   r0   r0   T   sC   … ñð /3€K�˜%×*Ñ*Ñ+Ó2ØOSÐ˜X e¨E°4¸×8IÑ8IÑ3JÐ,JÑ&KÑLÔSr-   r0   c                   ór   — e Zd ZU dZdZeej                     ed<   dZ	ee
eeej                     f      ed<   y)ÚMimiDecoderOutputaU  
    Args:
        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-   r.   r2   r2   h   sC   … ñð 15€L�(˜5×,Ñ,Ñ-Ó4ØOSÐ˜X e¨E°4¸×8IÑ8IÑ3JÐ,JÑ&KÑLÔSr-   r2   c                   óØ   ‡ — e Zd ZdZ	 	 	 	 	 ddedededededede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„ Zˆ xZS )Ú
MimiConv1dz;Conv1d with asymmetric or causal padding and normalization.Úin_channelsÚout_channelsÚkernel_sizeÚstrideÚdilationÚgroupsÚbiasc
           	      ó  •— t         ‰
| �  «        |j                  | _        |€|j                  n|| _        |dkD  r$|dkD  rt
        j                  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                  ¬«      }| j!                  d	|d
¬«       | j!                  d|d
¬«       | j!                  d||z
  d
¬«       | j"                  dz  | _        | j"                  | j$                  z
  | _        y )Nr   zNMimiConv1d has been initialized with stride > 1 and dilation > 1 (kernel_size=z stride=z, dilation=ú).)r9   r:   r;   r   ©Údtyper8   F©Ú
persistentr7   Úpadding_totalé   )ÚsuperÚ__init__Úuse_causal_convÚcausalÚpad_modeÚloggerÚwarningr   ÚConv1dÚconvr7   r(   Útensorr8   Úint64r9   Úregister_bufferrB   Úpadding_rightÚpadding_left)ÚselfÚconfigr5   r6   r7   r8   r9   r:   rH   r;   Ú	__class__s             €r.   rE   zMimiConv1d.__init__   so  ø€ ô 	‰ÑÔØ×,Ñ,ˆŒØ+3Ð+;˜ŸšÀˆŒð �AŠ:˜( Qš,Ü�N‰Nð!Ø!, ¨X°f°X¸[ÈÈ
ÐRTðVôô
 —I‘IØ˜ {°FÀXÐV\Ðcgô
ˆŒ	ð —i‘i×+Ñ+¨AÑ.ˆÜ—‘˜dŸi™i×.Ñ.¨qÑ1¼¿¹ÔEˆØ—9‘9×%Ñ% aÑ(ˆô —l‘l K°!¡O°xÑ#?À!Ñ#CÌ5Ï;É;ÔWˆà×Ñ˜X v¸%ÐÔ@Ø×Ñ˜]¨KÀEÐÔJØ×Ñ˜_¨k¸FÑ.BÈuÐÔUð "×/Ñ/°1Ñ4ˆÔØ ×.Ñ.°×1CÑ1CÑCˆÕr-   c                 óì   — t         j                  j                  }t        t         j                  j                  d«      r$t         j                  j                  j                  } || j
                  «       y ©NÚweight_norm©r   ÚutilsrW   ÚhasattrÚparametrizationsrL   ©rR   rW   s     r.   Úapply_weight_normzMimiConv1d.apply_weight_norm©   óF   € Ü—h‘h×*Ñ*ˆÜ”2—8‘8×,Ñ,¨mÔ<ÜŸ(™(×3Ñ3×?Ñ?ˆKá�D—I‘IÕr-   c                 óV   — t         j                  j                  | j                  «       y ©N©r   rY   Úremove_weight_normrL   ©rR   s    r.   rb   zMimiConv1d.remove_weight_norm°   ó   € Ü
�‰×#Ñ# D§I¡IÕ.r-   Úhidden_statesÚreturnc                 ó>  — |j                   d   }|| j                  z
  | j                  z   | j                  z  dz   }t	        j
                  |«      j                  t        j                  «      dz
  }|| j                  z  | j                  z   | j                  z
  }||z
  S )zSee `pad_for_conv1d`.éÿÿÿÿr   )Úshaper7   rB   r8   r(   ÚceilÚtorN   )rR   re   ÚlengthÚn_framesÚideal_lengths        r.   Ú_get_extra_padding_for_conv1dz(MimiConv1d._get_extra_padding_for_conv1d´   s�   € ð
 ×$Ñ$ RÑ(ˆØ˜T×-Ñ-Ñ-°×0BÑ0BÑBÀdÇkÁkÑQÐTUÑUˆÜ—:‘:˜hÓ'×*Ñ*¬5¯;©;Ó7¸!Ñ;ˆØ $§+¡+Ñ-°×0@Ñ0@Ñ@À4×CUÑCUÑUˆà˜fÑ$Ð$r-   ÚpaddingsÚmodeÚvaluec                 ól  — | j                   d   }|\  }}|dk(  s"t        j                  j                  | |||«      S t	        ||«      }d}||k  r*||z
  dz   }t        j                  j                  | d|f«      } t        j                  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.
        rh   Úreflectr   r   .N)ri   r   Ú
functionalÚpadÚmax)re   rp   rq   rr   rl   rQ   rP   Úmax_padÚ	extra_padÚpaddedÚends              r.   Ú_pad1dzMimiConv1d._pad1dÀ   sÃ   € ð ×$Ñ$ RÑ(ˆØ&.Ñ#ˆ�mØ�yÒ Ü—=‘=×$Ñ$ ]°H¸dÀEÓJÐJä�l MÓ2ˆØˆ	Ø�WÒØ &Ñ(¨1Ñ,ˆIÜŸM™M×-Ñ-¨m¸aÀ¸^ÓLˆMÜ—‘×"Ñ" =°(¸DÀ%ÓHˆØ�l‰l˜2Ñ Ñ*ˆØ�c˜4˜C˜4�iÑ Ð r-   c                 ó&  — | j                  |«      }| j                  r+| j                  || j                  |f| j                  ¬«      }n7| j                  || j
                  | j                  |z   f| j                  ¬«      }| j                  |«      }|S )N)rq   )ro   rG   r|   rB   rH   rQ   rP   rL   )rR   re   Úextra_paddings      r.   ÚforwardzMimiConv1d.forwardÔ   sŒ   € Ø×:Ñ:¸=ÓIˆà�;Š;à ŸK™K¨¸×8JÑ8JÈMÐ7ZÐae×anÑan˜KÓo‰Mà ŸK™KØ × 1Ñ 1°4×3EÑ3EÈÑ3UÐVÐ]a×]jÑ]jð (ó ˆMð Ÿ	™	 -Ó0ˆØÐr-   )r   r   r   NT)Úzeroç        )r$   r%   r&   r'   ÚintÚboolrE   r]   rb   r(   ÚTensorro   Ústaticmethodr   ÚstrÚfloatr|   r   Ú__classcell__©rT   s   @r.   r4   r4   |   sÖ   ø„ ÙEð ØØØØñ(Dð ð(Dð ð	(Dð
 ð(Dð ð(Dð ð(Dð ð(Dð õ(DòTò/ð
%à—|‘|ð
%ð 
�‰ó
%ð ñ!˜eŸl™lð !°e¸CÀ¸H±oð !ÈSð !Ðbgò !ó ð!ö$r-   r4   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.r5   r6   r7   r8   r:   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
  | _        y )N)r:   r;   ç      ð?zB`trim_right_ratio` != 1.0 only makes sense for causal convolutionsr   rC   )rD   rE   rF   rG   Útrim_right_ratior   ÚConvTranspose1drL   Ú
ValueErrorr7   r8   Úmathrj   rP   rQ   )
rR   rS   r5   r6   r7   r8   r:   r;   rB   rT   s
            €r.   rE   zMimiConvTranspose1d.__init__æ   sä   ø€ ô 	‰ÑÔØ×,Ñ,ˆŒØ &× 7Ñ 7ˆÔÜ×&Ñ& {°LÀ+ÈvÐ^dÐkoÔpˆŒ	à—’˜t×4Ñ4¸Ò;ÜÐaÓbÐbà—i‘i×+Ñ+¨AÑ.ˆØ—‘×!Ñ! !Ñ$ˆØ# fÑ,ˆð �;Š;ô "&§¡¨=¸4×;PÑ;PÑ+PÓ!QˆDÕð "/°!Ñ!3ˆDÔà)¨D×,>Ñ,>Ñ>ˆÕr-   c                 óì   — t         j                  j                  }t        t         j                  j                  d«      r$t         j                  j                  j                  } || j
                  «       y rV   rX   r\   s     r.   r]   z%MimiConvTranspose1d.apply_weight_norm
  r^   r-   c                 óV   — t         j                  j                  | j                  «       y r`   ra   rc   s    r.   rb   z&MimiConvTranspose1d.remove_weight_norm  rd   r-   c                 ó†   — | j                  |«      }|j                  d   | j                  z
  }|d| j                  |…f   }|S )Nrh   .)rL   ri   rP   rQ   )rR   re   r{   s      r.   r   zMimiConvTranspose1d.forward  sM   € ØŸ	™	 -Ó0ˆð ×!Ñ! "Ñ%¨×(:Ñ(:Ñ:ˆØ% c¨4×+<Ñ+<¸sÐ+BÐ&BÑCˆØÐr-   )r   r   T)
r$   r%   r&   r'   r‚   rE   r]   rb   r   rˆ   r‰   s   @r.   r‹   r‹   ã   sX   ø„ ÙNð ØØñ"?ð ð"?ð ð	"?ð
 ð"?ð ð"?ð õ"?òHò/ör-   r‹   c                   ó<   ‡ — e Zd ZdZdededee   fˆ fd„Zd„ Zˆ xZ	S )ÚMimiResnetBlockz;
    Residual block from SEANet model as used by Mimi.
    rS   ÚdimÚ	dilationsc           	      ó   •— t         ‰| �  «        |j                  df}t        |«      t        |«      k7  rt	        d«      ‚||j
                  z  }g }t        t        ||«      «      D ]R  \  }\  }}	|dk(  r|n|}
|t        |«      dz
  k(  r|n|}|t        j                  «       gz  }|t        ||
|||	¬«      gz  }ŒT t        j                  |«      | _        |j                  rt        |||d¬«      | _        y t        j                  «       | _        y )Nr   z7Number of kernel sizes should match number of dilationsr   )r9   )r7   )rD   rE   Úresidual_kernel_sizeÚlenr�   ÚcompressÚ	enumerateÚzipr   ÚELUr4   Ú
ModuleListÚblockÚuse_conv_shortcutÚshortcutÚIdentity)rR   rS   r—   r˜   Úkernel_sizesÚhiddenr¡   Úir7   r9   Úin_chsÚout_chsrT   s               €r.   rE   zMimiResnetBlock.__init__#  s  ø€ Ü‰ÑÔØ×3Ñ3°QÐ7ˆÜˆ|Ó¤ I£Ò.ÜÐVÓWÐWà˜Ÿ™Ñ'ˆØˆÜ*3´C¸ÀiÓ4PÓ*Qò 	[Ñ&ˆAÑ&�˜XØ šF‘S¨ˆFØ¤# lÓ"3°aÑ"7Ò7‘c¸VˆGØ”b—f‘f“h�ZÑˆEØ”j ¨°¸+ÐPXÔYÐZÑZ‰Eð		[ô
 —]‘] 5Ó)ˆŒ
à×#Ò#Ü& v¨s°CÀQÔGˆD�MäŸK™K›MˆD�Mr-   c                 ó`   — |}| j                   D ]
  } ||«      }Œ | j                  |«      |z   S r`   )r¡   r£   )rR   re   ÚresidualÚlayers       r.   r   zMimiResnetBlock.forward7  s:   € Ø ˆØ—Z‘Zò 	1ˆEÙ! -Ó0‰Mð	1ð �}‰}˜XÓ&¨Ñ6Ð6r-   )
r$   r%   r&   r'   r   r‚   r   rE   r   rˆ   r‰   s   @r.   r–   r–     s+   ø„ ñð*˜zð *°ð *ÀÀSÁ	õ *ö(7r-   r–   c                   ó.   ‡ — e Zd ZdZdefˆ fd„Zd„ Zˆ xZS )ÚMimiEncoderzSEANet encoder as used by Mimi.rS   c           	      ó~  •— t         ‰| �  «        t        ||j                  |j                  |j
                  «      g}d}t        |j                  «      D ]‚  }||j                  z  }t        |j                  «      D ]"  }|t        |||j                  |z  dg«      gz  }Œ$ |t        j                  «       gz  }|t        |||dz  |dz  |¬«      gz  }|dz  }Œ„ |t        j                  «       gz  }|t        |||j                  z  |j                  |j                  «      gz  }t        j                   |«      | _        y )Nr   rC   ©r7   r8   )rD   rE   r4   Úaudio_channelsÚnum_filtersr7   ÚreversedÚupsampling_ratiosÚrangeÚnum_residual_layersr–   Údilation_growth_rater   rŸ   Úhidden_sizeÚlast_kernel_sizer    Úlayers)rR   rS   ÚmodelÚscalingÚratioÚcurrent_scaleÚjrT   s          €r.   rE   zMimiEncoder.__init__B  s@  ø€ Ü‰ÑÔÜ˜F F×$9Ñ$9¸6×;MÑ;MÈv×OaÑOaÓbÐcˆØˆô ˜f×6Ñ6Ó7ò 	ˆEØ# f×&8Ñ&8Ñ8ˆMä˜6×5Ñ5Ó6ò g�Øœ/¨&°-À&×B]ÑB]Ð_`ÑB`ÐbcÐAdÓeÐfÑf‘ðgð ”b—f‘f“h�ZÑˆEØ”j ¨¸ÈÑ8IÐW\Ð_`ÑW`ÐinÔoÐpÑpˆEØ�q‰L‰Gð	ð 	”"—&‘&“(�ÑˆØ”*˜V W¨v×/AÑ/AÑ%AÀ6×CUÑCUÐW]×WnÑWnÓoÐpÑpˆä—m‘m EÓ*ˆ�r-   c                 ó8   — | j                   D ]
  } ||«      }Œ |S r`   ©rº   ©rR   re   r¬   s      r.   r   zMimiEncoder.forwardX  ó%   € Ø—[‘[ò 	1ˆEÙ! -Ó0‰Mð	1àÐr-   ©r$   r%   r&   r'   r   rE   r   rˆ   r‰   s   @r.   r®   r®   ?  s   ø„ Ù)ð+˜zõ +ö,r-   r®   c                   óB   ‡ — 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://arxiv.org/pdf/2103.17239.pdf).
    This rescales diagonally the residual outputs close to 0, with a learnt scale.
    c                 ó´   •— t         ‰| �  «        |j                  }|j                  }t	        j
                  t        j                  |f|d¬«      «      | _        y )NT)Úrequires_grad)	rD   rE   r¸   Úlayer_scale_initial_scaler   Ú	Parameterr(   ÚfullÚscale)rR   rS   ÚchannelsÚinitial_scalerT   s       €r.   rE   zMimiLayerScale.__init__c  sD   ø€ Ü‰ÑÔØ×%Ñ%ˆØ×8Ñ8ˆÜ—\‘\¤%§*¡*¨h¨[¸-ÐW[Ô"\Ó]ˆ�
r-   Úxc                 ó    — | j                   |z  S r`   )rÌ   )rR   rÏ   s     r.   r   zMimiLayerScale.forwardi  s   € Ø�z‰z˜A‰~Ðr-   )	r$   r%   r&   r'   rE   r(   r„   r   rˆ   r‰   s   @r.   rÆ   rÆ   ^  s   ø„ ñô^ð˜Ÿ™÷ r-   rÆ   c                   ó^   ‡ — e Zd Zddefˆ fd„Z ej                  «       ed„ «       «       Zˆ xZ	S )ÚMimiRotaryEmbeddingrS   c                 óê  •— t         ‰| �  «        t        |d«      rG|j                  �;|j                  j	                  d|j                  j	                  d«      «      | _        nd| _        |j                  | _        |j                  | _        || _	        t        | j
                     | _        | j                  | j                  |«      \  }| _        | j                  d|d¬«       | j                  | _        y )NÚrope_scalingÚ	rope_typeÚtypeÚdefaultÚinv_freqFr@   )rD   rE   rZ   rÔ   ÚgetrÕ   Úmax_position_embeddingsÚmax_seq_len_cachedÚoriginal_max_seq_lenrS   r   Úrope_init_fnÚattention_scalingrO   rØ   Úoriginal_inv_freq)rR   rS   ÚdevicerØ   rT   s       €r.   rE   zMimiRotaryEmbedding.__init__o  sÉ   ø€ Ü‰ÑÔä�6˜>Ô*¨v×/BÑ/BÐ/NØ#×0Ñ0×4Ñ4°[À&×BUÑBU×BYÑBYÐZ`ÓBaÓbˆD�Nà&ˆDŒNØ"(×"@Ñ"@ˆÔØ$*×$BÑ$BˆÔ!àˆŒÜ/°·±Ñ?ˆÔà+/×+<Ñ+<¸T¿[¹[È&Ó+QÑ(ˆ�$Ô(Ø×Ñ˜Z¨¸eÐÔDØ!%§¡ˆÕr-   c                 ób  — | j                   d d d …d f   j                  «       j                  |j                  d   dd«      j	                  |j
                  «      }|d d …d d d …f   j                  «       }t        |j
                  j                  t        «      r/|j
                  j                  dk7  r|j
                  j                  nd}t        j                  |d¬«      5  |j                  «       |j                  «       z  j                  dd«      }t        j                  ||fd¬	«      }|j                  «       | j                  z  }|j                  «       | j                  z  }	d d d «       j	                  |j                   ¬
«      	j	                  |j                   ¬
«      fS # 1 sw Y   ŒAxY w)Nr   rh   r   ÚmpsÚcpuF)Údevice_typeÚenabledrC   ©r—   r>   )rØ   r‡   Úexpandri   rk   rà   Ú
isinstancerÖ   r†   r(   ÚautocastÚ	transposeÚcatÚcosrÞ   Úsinr?   )
rR   rÏ   Úposition_idsÚinv_freq_expandedÚposition_ids_expandedrä   ÚfreqsÚembrì   rí   s
             r.   r   zMimiRotaryEmbedding.forward€  sV  € ð !ŸM™M¨$²°4¨-Ñ8×>Ñ>Ó@×GÑGÈ×HZÑHZÐ[\ÑH]Ð_aÐcdÓe×hÑhÐij×iqÑiqÓrÐØ ,ªQ°²a¨ZÑ 8× >Ñ >Ó @Ðä'1°!·(±(·-±-ÄÔ'EÈ!Ï(É(Ï-É-Ð[`ÒJ`�a—h‘h—m’mÐfkˆÜ�^‰^¨¸UÔCñ 	5Ø&×,Ñ,Ó.Ð1F×1LÑ1LÓ1NÑN×YÑYÐZ[Ð]^Ó_ˆEÜ—)‘)˜U E˜N°Ô3ˆCØ—'‘'“)˜d×4Ñ4Ñ4ˆCØ—'‘'“)˜d×4Ñ4Ñ4ˆC÷		5ð �v‰v˜AŸG™GˆvÓ$ c§f¡f°1·7±7 fÓ&;Ð;Ð;÷	5ð 	5ús   Ã BF%Æ%F.r`   )
r$   r%   r&   r   rE   r(   Úno_gradr   r   rˆ   r‰   s   @r.   rÒ   rÒ   n  s3   ø„ ñ/˜zõ /ð" €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..Nrh   rC   ræ   )ri   r(   rë   )rÏ   Úx1Úx2s      r.   Úrotate_halfr÷   ‘  sZ   € à	
ˆ3Ð"�!—'‘'˜"‘+ Ñ"Ð"Ð"Ñ	#€BØ	
ˆ3�—‘˜‘˜qÑ Ñ"Ð"Ñ	#€BÜ�9‰9�r�c˜2�Y BÔ'Ð'r-   c                 óž   — |j                  |«      }|j                  |«      }| |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.
        position_ids (`torch.Tensor`, *optional*):
            Deprecated and unused.
        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í   rî   Úunsqueeze_dimÚq_embedÚk_embeds           r.   Úapply_rotary_pos_embrÿ   ™  sY   € ð( �-‰-˜Ó
&€CØ
�-‰-˜Ó
&€CØ�3‰wœ; q›>¨CÑ/Ñ0€GØ�3‰wœ; q›>¨CÑ/Ñ0€GØ�GÐÐr-   c                   óV   ‡ — 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¬«      | _        y )NF©r;   )rD   rE   rS   r
   Ú
hidden_actÚactivation_fnr   ÚLinearr¸   Úintermediate_sizeÚfc1Úfc2©rR   rS   rT   s     €r.   rE   zMimiMLP.__init__µ  sj   ø€ Ü‰ÑÔØˆŒÜ# F×$5Ñ$5Ñ6ˆÔÜ—9‘9˜V×/Ñ/°×1IÑ1IÐPUÔVˆŒÜ—9‘9˜V×5Ñ5°v×7IÑ7IÐPUÔVˆ�r-   re   rf   c                 ól   — | j                  |«      }| j                  |«      }| j                  |«      }|S r`   )r  r  r	  )rR   re   s     r.   r   zMimiMLP.forward½  s4   € ØŸ™ Ó/ˆØ×*Ñ*¨=Ó9ˆØŸ™ Ó/ˆØÐr-   )r$   r%   r&   rE   r(   r„   r   rˆ   r‰   s   @r.   r  r  ´  s$   ø„ ôWð U§\¡\ð °e·l±l÷ r-   r  re   Ún_reprf   c                 óª   — | j                   \  }}}}|dk(  r| S | dd…dd…ddd…dd…f   j                  |||||«      } | j                  |||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)ri   rç   Úreshape)re   r  ÚbatchÚnum_key_value_headsÚslenÚhead_dims         r.   Ú	repeat_kvr  Å  so   € ð
 2?×1DÑ1DÑ.€EÐ  hØ�‚zØÐØ!¢!¢Q¨ªa²Ð"2Ñ3×:Ñ:¸5ÐBUÐW\Ð^bÐdlÓm€MØ× Ñ  Ð(;¸eÑ(CÀTÈ8ÓTÐTr-   c                   ó,  ‡ — e Zd ZdZddedee   fˆ fd„Z	 	 	 	 	 	 ddej                  deej                     deej                     dee   d	ed
edeej                     deej                  eej                     eeej                        f   fd„Zˆ xZS )ÚMimiAttentionz=Multi-headed attention from 'Attention Is All You Need' paperrS   Ú	layer_idxc                 ó(  •— t         ‰| �  «        || _        || _        |€-t        j                  d| j                  j                  › d�«       |j                  | _        |j                  | _	        |j                  | _        |j                  | _        |j                  | _        | j                  | j                  z  | _        |j                  | _        |j                   | _        d| _        dt%        j&                  |j                  «      z  | _        | j                  | j                  z  dk7  r&t+        d| j                  › d| j                  › d�«      ‚t-        j.                  | j                  | j                  | j                  z  |j0                  ¬	«      | _        t-        j.                  | j                  | j                  | j                  z  |j0                  ¬	«      | _        t-        j.                  | j                  | j                  | j                  z  |j0                  ¬	«      | _        t-        j.                  | j                  | j                  z  | j                  |j0                  ¬	«      | _        t;        |«      | _        |j>                  | _        y )
NzInstantiating z¹ without passing a `layer_idx` is not recommended and will lead to errors during the forward call if caching is used. Please make sure to provide a `layer_idx` when creating this class.Tr   r   z?hidden_size must be divisible by num_heads (got `hidden_size`: z and `num_heads`: r=   r  ) rD   rE   rS   r  rI   Úwarning_oncerT   r$   Úattention_dropoutr¸   Únum_attention_headsÚ	num_headsr  r  Únum_key_value_groupsrÚ   Ú
rope_thetaÚ	is_causalr‘   Úsqrtr¼   r�   r   r  Úattention_biasÚq_projÚk_projÚv_projÚo_projrÒ   Ú
rotary_embÚsliding_window©rR   rS   r  rT   s      €r.   rE   zMimiAttention.__init__Ö  s   ø€ Ü‰ÑÔØˆŒØ"ˆŒØÐÜ×ÑØ  §¡×!8Ñ!8Ð 9ð :,ð ,ôð "(×!9Ñ!9ˆÔØ!×-Ñ-ˆÔØ×3Ñ3ˆŒØŸ™ˆŒØ#)×#=Ñ#=ˆÔ Ø$(§N¡N°d×6NÑ6NÑ$NˆÔ!Ø'-×'EÑ'EˆÔ$Ø ×+Ñ+ˆŒØˆŒØœ4Ÿ9™9 V§_¡_Ó5Ñ5ˆŒà×Ñ˜dŸn™nÑ,°Ò1ÜØQÐRV×RbÑRbÐQcØ$ T§^¡^Ð$4°Bð8óð ô
 —i‘i × 0Ñ 0°$·.±.À4Ç=Á=Ñ2PÐW]×WlÑWlÔmˆŒÜ—i‘i × 0Ñ 0°$×2JÑ2JÈTÏ]É]Ñ2ZÐag×avÑavÔwˆŒÜ—i‘i × 0Ñ 0°$×2JÑ2JÈTÏ]É]Ñ2ZÐag×avÑavÔwˆŒÜ—i‘i §¡°·±Ñ >À×@PÑ@PÐW]×WlÑWlÔmˆŒÜ-¨fÓ5ˆŒØ$×3Ñ3ˆÕr-   re   Úattention_maskrî   Úpast_key_valueÚoutput_attentionsÚ	use_cacheÚcache_positionrf   c                 ó  — |j                  «       \  }}	}
| j                  |«      }| j                  |«      }| j                  |«      }|j	                  ||	| j
                  | j                  «      j                  dd«      }|j	                  ||	| j                  | j                  «      j                  dd«      }|j	                  ||	| j                  | j                  «      j                  dd«      }| j                  ||«      \  }}t        ||||«      \  }}|�'|||dœ}|j                  ||| j                  |«      \  }}t        || j                  «      }t        || j                  «      }t        j                   ||j                  dd«      «      | j"                  z  }|�#|d d …d d …d d …d |j$                  d   …f   }||z   }t&        j(                  j+                  |dt        j,                  ¬«      j/                  |j0                  «      }t&        j(                  j3                  || j4                  | j6                  ¬«      }t        j                   ||«      }|j                  «       || j
                  |	| j                  fk7  r7t9        d	|| j
                  |	| j                  f› d
|j                  «       › �«      ‚|j                  dd«      j;                  «       }|j	                  ||	d«      }| j=                  |«      }|sd }|||fS )Nr   rC   ©rí   rì   r,  r	   éþÿÿÿrh   )r—   r?   )ÚpÚtrainingz `attn_output` should be of size z	, but is )Úsizer!  r"  r#  Úviewr  r  rê   r  r%  rÿ   Úupdater  r  r  r(   Úmatmulr¼   ri   r   ru   ÚsoftmaxÚfloat32rk   r?   Údropoutr  r1  r�   Ú
contiguousr$  )rR   re   r(  rî   r)  r*  r+  r,  ÚbszÚq_lenÚ_Úquery_statesÚ
key_statesÚvalue_statesrì   rí   Úcache_kwargsÚattn_weightsÚcausal_maskÚattn_outputs                       r.   r   zMimiAttention.forwardù  sÐ  € ð &×*Ñ*Ó,‰ˆˆU�Aà—{‘{ =Ó1ˆØ—[‘[ Ó/ˆ
Ø—{‘{ =Ó1ˆà#×(Ñ(¨¨e°T·^±^ÀTÇ]Á]ÓS×]Ñ]Ð^_ÐabÓcˆØ—_‘_ S¨%°×1IÑ1IÈ4Ï=É=ÓY×cÑcÐdeÐghÓiˆ
Ø#×(Ñ(¨¨e°T×5MÑ5MÈtÏ}É}Ó]×gÑgÐhiÐklÓmˆà—?‘? <°Ó>‰ˆˆSÜ#7¸ÀjÐRUÐWZÓ#[Ñ ˆ�jàÐ%à#&¨sÀnÑUˆLØ'5×'<Ñ'<¸ZÈÐW[×WeÑWeÐgsÓ'tÑ$ˆJ˜ä˜z¨4×+DÑ+DÓEˆ
Ü  ¨t×/HÑ/HÓIˆä—|‘| L°*×2FÑ2FÀqÈ!Ó2LÓMÐPT×P\ÑP\Ñ\ˆàÐ%Ø(ªªAªqÐ2H°J×4DÑ4DÀRÑ4HÐ2HÐ)HÑIˆKØ'¨+Ñ5ˆLô —}‘}×,Ñ,¨\¸rÌÏÉÐ,ÓW×ZÑZÐ[g×[mÑ[mÓnˆÜ—}‘}×,Ñ,¨\¸T×=SÑ=SÐ^b×^kÑ^kÐ,ÓlˆÜ—l‘l <°Ó>ˆà×ÑÓ # t§~¡~°u¸d¿m¹mÐ!LÒLÜØ2°C¸¿¹ÈÐPT×P]ÑP]Ð3^Ð2_ð `Ø×$Ñ$Ó&Ð'ð)óð ð
 "×+Ñ+¨A¨qÓ1×<Ñ<Ó>ˆà!×&Ñ& s¨E°2Ó6ˆØ—k‘k +Ó.ˆá ØˆLà˜L¨.Ð8Ð8r-   r`   ©NNNFFN)r$   r%   r&   r'   r   r   r‚   rE   r(   r„   r)   r   rƒ   r   r   rˆ   r‰   s   @r.   r  r  Ó  sÓ   ø„ ÙGñ!4˜zð !4°h¸s±mõ !4ðL 26Ø37Ø*.Ø"'ØØ59ñ89à—|‘|ð89ð ! §¡Ñ.ð89ð ˜u×/Ñ/Ñ0ð	89ð
 ! ™ð89ð  ð89ð ð89ð ! ×!1Ñ!1Ñ2ð89ð 
ˆu�|‰|˜X e§l¡lÑ3°X¸eÀEÇLÁLÑ>QÑ5RÐRÑ	S÷89r-   r  c                   ó  ‡ — e Zd ZdZˆ fd„Z	 	 	 	 	 	 ddej                  deej                     deej                     dee	   de
de
d	eej                     d
eej                  eej                     eeej                        f   fd„Zˆ xZS )ÚMimiFlashAttention2aD  
    Mimi flash attention module. This module inherits from `MimiAttention` as the weights of the module stays
    untouched. The only required change would be on the forward pass where it needs to correctly call the public API of
    flash attention and deal with padding tokens in case the input contains any of them.
    c                 óB   •— t        ‰| �  |i |¤Ž t        «       | _        y r`   )rD   rE   r   Ú_flash_attn_uses_top_left_mask)rR   ÚargsÚkwargsrT   s      €r.   rE   zMimiFlashAttention2.__init__=  s#   ø€ Ü‰Ñ˜$Ð) &Ò)ô
 /PÓ.QˆÕ+r-   re   r(  rî   r)  r*  r+  r,  rf   c                 óø  — t        |t        «      rt        d«      ‚d}|j                  «       \  }}	}
| j	                  |«      }| j                  |«      }| j                  |«      }|j                  ||	| j                  | j                  «      j                  dd«      }|j                  ||	| j                  | j                  «      j                  dd«      }|j                  ||	| j                  | j                  «      j                  dd«      }| j                  ||«      \  }}t        ||||«      \  }}|�'|||dœ}|j                  ||| j                  |«      \  }}|j                  dd«      }|j                  dd«      }|j                  dd«      }| j                   r| j"                  nd}|j$                  }|t&        j(                  k(  rÂt'        j*                  «       rt'        j,                  «       }nMt/        | j0                  d«      r| j0                  j2                  }n | j                  j4                  j$                  }t6        j9                  d|› d	�«       |j;                  |«      }|j;                  |«      }|j;                  |«      }t=        |||||	||t?        | d
d «      | j@                  | jB                  ¬«
      }|jE                  ||	d«      jG                  «       }| jI                  |«      }|sd }||fS )NzÈ`static` cache implementation is not compatible with `attn_implementation==flash_attention_2` make sure to use `sdpa` in the mean time, and open an issue at https://github.com/huggingface/transformersFr   rC   r.  r�   Ú_pre_quantization_dtypez¾The input hidden states seems to be silently casted in float32, this might be related to the fact you have upcasted embedding or layer norm layers in float32. We will cast back the input in ú.r&  )rî   r8  r&  r  Úuse_top_left_maskrh   )%rè   r   r�   r2  r!  r"  r#  r3  r  r  rê   r  r%  rÿ   r4  r  r1  r  r?   r(   r7  Úis_autocast_enabledÚget_autocast_gpu_dtyperZ   rS   rL  ÚweightrI   r  rk   r   Úgetattrr  rH  r  r9  r$  )rR   re   r(  rî   r)  r*  r+  r,  r:  r;  r<  r=  r>  r?  rì   rí   r@  Údropout_rateÚinput_dtypeÚtarget_dtyperC  rA  s                         r.   r   zMimiFlashAttention2.forwardE  sÌ  € ô �n¤kÔ2Üð}óð ð
 "Ðà%×*Ñ*Ó,‰ˆˆU�Aà—{‘{ =Ó1ˆØ—[‘[ Ó/ˆ
Ø—{‘{ =Ó1ˆð
 $×(Ñ(¨¨e°T·^±^ÀTÇ]Á]ÓS×]Ñ]Ð^_ÐabÓcˆØ—_‘_ S¨%°×1IÑ1IÈ4Ï=É=ÓY×cÑcÐdeÐghÓiˆ
Ø#×(Ñ(¨¨e°T×5MÑ5MÈtÏ}É}Ó]×gÑgÐhiÐklÓmˆà—?‘? <°Ó>‰ˆˆSÜ#7¸ÀjÐRUÐWZÓ#[Ñ ˆ�jàÐ%à#&¨sÀnÑUˆLØ'5×'<Ñ'<¸ZÈÐW[×WeÑWeÐgsÓ'tÑ$ˆJ˜ð $×-Ñ-¨a°Ó3ˆØ×)Ñ)¨!¨QÓ/ˆ
Ø#×-Ñ-¨a°Ó3ˆà15·²�t×-Ò-ÀCˆð #×(Ñ(ˆØœ%Ÿ-™-Ò'Ü×(Ñ(Ô*Ü$×;Ñ;Ó=‘ä˜Ÿ™Ð&?Ô@Ø#Ÿ{™{×BÑB‘à#Ÿ{™{×1Ñ1×7Ñ7�ä×Ñðà �> ð$ôð (Ÿ?™?¨<Ó8ˆLØ#Ÿ™ |Ó4ˆJØ'Ÿ?™?¨<Ó8ˆLä.ØØØØØØ%Ø Ü" 4Ð)9¸4Ó@Ø—n‘nØ"×AÑAô
ˆð "×)Ñ)¨#¨u°bÓ9×DÑDÓFˆØ—k‘k +Ó.ˆá ØˆLà˜L¨.Ð8Ð8r-   rD  )r$   r%   r&   r'   rE   r(   r„   r   r)   r   rƒ   r   r   rˆ   r‰   s   @r.   rF  rF  6  sÎ   ø„ ñôRð 6:Ø37Ø*.Ø"'ØØ59ñ\9à—|‘|ð\9ð ! ×!1Ñ!1Ñ2ð\9ð ˜u×/Ñ/Ñ0ð	\9ð
 ! ™ð\9ð  ð\9ð ð\9ð ! ×!1Ñ!1Ñ2ð\9ð 
ˆu�|‰|˜X e§l¡lÑ3°X¸eÀEÇLÁLÑ>QÑ5RÐRÑ	S÷\9r-   rF  c                   ó  ‡ — e Zd ZdZ	 	 	 	 	 	 ddej
                  deej
                     deej                     dee   de	de	deej                     d	e
ej
                  eej
                     ee
ej
                        f   fˆ fd
„Zˆ xZS )ÚMimiSdpaAttentionzö
    Mimi attention module using torch.nn.functional.scaled_dot_product_attention. This module inherits from
    `MimiAttention` as the weights of the module stays untouched. The only changes are on the forward pass to adapt to
    SDPA API.
    re   r(  rî   r)  r*  r+  r,  rf   c           	      óB  •— |r+t         j                  d«       t        ‰| �  |||||||¬«      S |j	                  «       \  }	}
}| j                  |«      }| j                  |«      }| j                  |«      }|j                  |	|
| j                  | j                  «      j                  dd«      }|j                  |	|
| j                  | j                  «      j                  dd«      }|j                  |	|
| j                  | j                  «      j                  dd«      }| j                  ||«      \  }}t        ||||«      \  }}|�'|||dœ}|j                  ||| j                   |«      \  }}t#        || j$                  «      }t#        || j$                  «      }|}|�|d d …d d …d d …d |j&                  d   …f   }|j(                  j*                  dk(  r2|�0|j-                  «       }|j-                  «       }|j-                  «       }|€|
dkD  rdnd	}t.        j0                  j2                  j5                  ||||| j6                  r| j8                  nd
|¬«      }|j                  dd«      j-                  «       }|j                  |	|
d«      }| j;                  |«      }|d |fS )Na…  MimiModel is using MimiSdpaAttention, but `torch.nn.functional.scaled_dot_product_attention` does not support `output_attentions=True`. Falling back to the manual attention implementation, but specifying the manual implementation will be required from Transformers version v5.0.0 onwards. This warning can be removed using the argument `attn_implementation="eager"` when loading the model.©re   r(  rî   r)  r*  r+  r,  r   rC   r.  r/  ÚcudaTFr�   )Ú	attn_maskÚ	dropout_pr  rh   )rI   r  rD   r   r2  r!  r"  r#  r3  r  r  rê   r  r%  rÿ   r4  r  r  r  ri   rà   rÖ   r9  r(   r   ru   Úscaled_dot_product_attentionr1  r  r$  )rR   re   r(  rî   r)  r*  r+  r,  rJ  r:  r;  r<  r=  r>  r?  rì   rí   r@  rB  r  rC  rT   s                        €r.   r   zMimiSdpaAttention.forward®  sª  ø€ ñ ä×Ñð[ôô ‘7‘?Ø+Ø-Ø)Ø-Ø"3Ø#Ø-ð #ó ð ð &×*Ñ*Ó,‰ˆˆU�Aà—{‘{ =Ó1ˆØ—[‘[ Ó/ˆ
Ø—{‘{ =Ó1ˆà#×(Ñ(¨¨e°T·^±^ÀTÇ]Á]ÓS×]Ñ]Ð^_ÐabÓcˆØ—_‘_ S¨%°×1IÑ1IÈ4Ï=É=ÓY×cÑcÐdeÐghÓiˆ
Ø#×(Ñ(¨¨e°T×5MÑ5MÈtÏ}É}Ó]×gÑgÐhiÐklÓmˆà—?‘? <°Ó>‰ˆˆSÜ#7¸ÀjÐRUÐWZÓ#[Ñ ˆ�jàÐ%à#&¨sÀnÑUˆLØ'5×'<Ñ'<¸ZÈÐW[×WeÑWeÐgsÓ'tÑ$ˆJ˜ä˜z¨4×+DÑ+DÓEˆ
Ü  ¨t×/HÑ/HÓIˆà$ˆØÐ%Ø%¢aªªAÐ/E°×1AÑ1AÀ"Ñ1EÐ/EÐ&EÑFˆKð ×Ñ×#Ñ# vÒ-°+Ð2IØ'×2Ñ2Ó4ˆLØ#×.Ñ.Ó0ˆJØ'×2Ñ2Ó4ˆLð (Ð/°E¸A²I‘DÀ5ˆ	ä—h‘h×)Ñ)×FÑFØØØØ!Ø04·²�d×,Ò,À3Øð Gó 
ˆð "×+Ñ+¨A¨qÓ1×<Ñ<Ó>ˆØ!×&Ñ& s¨E°2Ó6ˆà—k‘k +Ó.ˆà˜D .Ð0Ð0r-   rD  )r$   r%   r&   r'   r(   r„   r   r)   r   rƒ   r   r   rˆ   r‰   s   @r.   rW  rW  ¦  sÌ   ø„ ñð 26Ø37Ø*.Ø"'ØØ59ñM1à—|‘|ðM1ð ! §¡Ñ.ðM1ð ˜u×/Ñ/Ñ0ð	M1ð
 ! ™ðM1ð  ðM1ð ðM1ð ! ×!1Ñ!1Ñ2ðM1ð 
ˆu�|‰|˜X e§l¡lÑ3°X¸eÀEÇLÁLÑ>QÑ5RÐRÑ	S÷M1ñ M1r-   rW  )ÚeagerÚflash_attention_2Úsdpac                   ó(  ‡ — e Zd Zdedefˆ fd„Z	 	 	 	 	 	 ddej                  deej                     deej                     dee
   dee   d	ee   d
eej                     deej                  eeej                  ej                  f      f   fd„Zˆ xZS )ÚMimiTransformerLayerrS   r  c                 ó¢  •— t         ‰| �  «        |j                  | _        t        |j                     ||¬«      | _        t        |«      | _        t        j                  |j                  |j                  ¬«      | _        t        j                  |j                  |j                  ¬«      | _        t        |«      | _        t        |«      | _        y )N)rS   r  )Úeps)rD   rE   r¸   ÚMIMI_ATTENTION_CLASSESÚ_attn_implementationÚ	self_attnr  Úmlpr   Ú	LayerNormÚnorm_epsÚinput_layernormÚpost_attention_layernormrÆ   Úself_attn_layer_scaleÚmlp_layer_scaler'  s      €r.   rE   zMimiTransformerLayer.__init__  s–   ø€ Ü‰ÑÔØ!×-Ñ-ˆÔä/°×0KÑ0KÑLÐTZÐfoÔpˆŒä˜6“?ˆŒÜ!Ÿ|™|¨F×,>Ñ,>ÀFÇOÁOÔTˆÔÜ(*¯©°V×5GÑ5GÈVÏ_É_Ô(]ˆÔ%Ü%3°FÓ%;ˆÔ"Ü-¨fÓ5ˆÕr-   re   r(  rî   r)  r*  r+  r,  rf   c                 ó&  — |}	| j                  |«      } | j                  d|||||||dœ|¤Ž\  }}
}|	| j                  |«      z   }|}	| j                  |«      }| j	                  |«      }|	| j                  |«      z   }|f}|r||
fz  }|r||fz  }|S )a  
        Args:
            hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)`
            attention_mask (`torch.FloatTensor`, *optional*):
                attention mask of size `(batch_size, sequence_length)` if flash attention is used or `(batch_size, 1,
                query_sequence_length, key_sequence_length)` if default attention is used.
            output_attentions (`bool`, *optional*):
                Whether or not to return the attentions tensors of all attention layers. See `attentions` under
                returned tensors for more detail.
            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`).
            past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states
            cache_position (`torch.LongTensor` of shape `(sequence_length)`, *optional*):
                Indices depicting the position of the input sequence tokens in the sequence
            kwargs (`dict`, *optional*):
                Arbitrary kwargs to be ignored, used for FSDP and other methods that injects code
                into the model
        rY  r,   )rk  rg  rm  rl  rh  rn  )rR   re   r(  rî   r)  r*  r+  r,  rJ  r«   Úself_attn_weightsÚpresent_key_valueÚoutputss                r.   r   zMimiTransformerLayer.forward  sÙ   € ð< !ˆà×,Ñ,¨]Ó;ˆð ?M¸d¿n¹nð 	?
Ø'Ø)Ø%Ø)Ø/ØØ)ñ	?
ð ñ	?
Ñ;ˆÐ(Ð*;ð ! 4×#=Ñ#=¸mÓ#LÑLˆð !ˆØ×5Ñ5°mÓDˆØŸ™ Ó/ˆØ  4×#7Ñ#7¸Ó#FÑFˆà Ð"ˆáØÐ)Ð+Ñ+ˆGáØÐ)Ð+Ñ+ˆGàˆr-   rD  )r$   r%   r&   r   r‚   rE   r(   r„   r   r)   r   rƒ   r   r+   r   rˆ   r‰   s   @r.   rb  rb    s×   ø„ ð
6˜zð 
6°cõ 
6ð 26Ø37Ø*.Ø,1Ø$)Ø59ñ=à—|‘|ð=ð ! §¡Ñ.ð=ð ˜u×/Ñ/Ñ0ð	=ð
 ! ™ð=ð $ D™>ð=ð ˜D‘>ð=ð ! ×!1Ñ!1Ñ2ð=ð 
ˆu× Ñ  (¨5°×1BÑ1BÀE×DUÑDUÐ1UÑ+VÑ"WÐWÑ	X÷=r-   rb  c                   ó  ‡ — e Zd ZdZdefˆ fd„Z	 	 	 	 	 	 	 	 	 ddeej                     deej                     deej                     dee
eeej                     f      dee   d	ee   d
ee   dee   deej                     de
eef   fd„Z	 ddej                  dej                  dej                  ded	ef
d„Zedej                  dededej*                  dej,                  dej                  dededefd„«       Zˆ xZS )ÚMimiTransformerModelz�
    Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`MimiTransformerLayer`]

    Args:
        config: MimiConfig
    rS   c           	      óô   •— t         ‰| �  «        t        j                  t	        |j
                  «      D �cg c]  }t        ||«      ‘Œ c}«      | _        |j                  | _        d| _	        || _
        y c c}w )NF)rD   rE   r   r    rµ   Únum_hidden_layersrb  rº   rf  Úgradient_checkpointingrS   r'  s      €r.   rE   zMimiTransformerModel.__init__Z  sd   ø€ Ü‰ÑÔä—m‘mÜFKÈF×LdÑLdÓFeÖf¸Ô! &¨)Õ4Òfó
ˆŒð %+×$?Ñ$?ˆÔ!à&+ˆÔ#Øˆ�ùò gs   ¶A5re   r(  rî   Úpast_key_valuesr+  r*  Úoutput_hidden_statesÚreturn_dictr,  rf   c
                 ó2  — |�|n| j                   j                  }|�|n| j                   j                  }|�|n| j                   j                  }|�|n| j                   j                  }| j
                  r%| j                  r|rt        j                  d«       d}|rGt        |t        «      s7|€t        «       }n*t        j                  |«      }t        j                  d«       |	€F|�|j                  «       nd}
t        j                  |
|
|j                   d   z   |j"                  ¬«      }	|€|	j%                  d«      }d}|�| j'                  |||	||«      }|rdnd}|rdnd}d}| j(                  D ]p  }|r||fz  }| j
                  r/| j                  r#| j+                  |j,                  |||||||	«      }n ||||||||	¬	«      }|d   }|r	||rd
nd   }|sŒh||d   fz  }Œr |r||fz  }|r|nd}|st/        d„ ||||fD «       «      S t1        ||||¬«      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://arxiv.org/abs/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` or `tuple(tuple(torch.FloatTensor))`, *optional*):
                Pre-computed hidden-states (key and values in the self-attention blocks and in the cross-attention
                blocks) that can be used to speed up sequential decoding. 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`.

                Two formats are allowed:
                - a [`~cache_utils.Cache`] instance;
                - Tuple of `tuple(torch.FloatTensor)` of length `config.n_layers`, with each tuple having 2 tensors of
                shape `(batch_size, num_heads, sequence_length, embed_size_per_head)`). This is also known as the legacy
                cache format.

                The model will output the same cache format that is fed as input. If no `past_key_values` are passed, the
                legacy cache format will be returned.

                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`.FzÿWe detected that you are passing `past_key_values` as a tuple of tuples. This is deprecated and will be removed in v4.47. Please convert your cache or use an appropriate `Cache` class (https://huggingface.co/docs/transformers/kv_cache#legacy-cache-format)r   r   ©rà   r,   )r(  rî   r)  r*  r+  r,  rC   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wr`   r,   )Ú.0Úvs     r.   ú	<genexpr>z/MimiTransformerModel.forward.<locals>.<genexpr>  s   è ø€ Òt˜qÐfgÑfsœÑtùs   ‚Š)Úlast_hidden_staterx  re   Ú
attentions)rS   r*  ry  r+  Úuse_return_dictrw  r1  rI   r  rè   r   r   Úfrom_legacy_cacheÚget_seq_lengthr(   Úarangeri   rà   rù   Ú_update_causal_maskrº   Ú_gradient_checkpointing_funcÚ__call__Útupler   )rR   re   r(  rî   rx  r+  r*  ry  rz  r,  Úpast_seen_tokensrB  Úall_hidden_statesÚall_self_attnsÚnext_decoder_cacheÚdecoder_layerÚlayer_outputsÚ
next_caches                     r.   r   zMimiTransformerModel.forwarde  s‚  € ðL 2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð "+Ð!6‘I¸D¿K¹K×<QÑ<Qˆ	à%0Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà×&Ò&¨4¯=ª=¹YÜ×ÑØjôð ˆIáœZ¨¼Ô?ØÐ&Ü".£.‘ä".×"@Ñ"@ÀÓ"Q�Ü×#Ñ#ð^ôð Ð!ØCRÐC^˜×=Ñ=Ô?ÐdeÐÜ"Ÿ\™\Ø Ð"2°]×5HÑ5HÈÑ5KÑ"KÐTa×ThÑThôˆNð ÐØ)×3Ñ3°AÓ6ˆLàˆØÐ%Ø×2Ñ2Ø ¨~¸ÐPaóˆKñ
 #7™B¸DÐÙ0™°dˆØ!Ðà!Ÿ[™[ò  	6ˆMÙ#Ø! mÐ%5Ñ5Ð!à×*Ò*¨t¯}ª}Ø $× AÑ AØ!×*Ñ*Ø!ØØ Ø#Ø%ØØ"ó	!‘ñ !.Ø!Ø#.Ø!-Ø#2Ø&7Ø'Ø#1ô!�ð *¨!Ñ,ˆMáØ%2Ñ8I±1ÈqÑ%QÐ"â Ø =°Ñ#3Ð"5Ñ5‘ðA 	6ñF  Ø -Ð!1Ñ1Ðá+4Ñ'¸$ˆ
áÜÑt ]°JÐ@QÐSaÐ$bÔtÓtÐtä&Ø+Ø&Ø+Ø%ô	
ð 	
r-   Úinput_tensorc                 ó  — | j                   j                  dk(  rS|�H|�F|d d …df   j                  «       j                  «       |j	                  «       d   k7  }|rt        d«      ‚|�d|v r|S y |�|j                  «       nd}t        |t        «      }t        |t        «      }	| j                   j                  dk(  r?|s=|	s;|s9t        j                  |||| j                   j                  | j                  ¬«      ry |j                  |j                  }}
t!        j"                  |
«      j$                  }|j&                  d   }|	s|r|j)                  «       }n1t        |t         j*                  «      r|j&                  d   n||z   dz   }| j-                  ||||
|||j&                  d   | j                   |¬	«	      }| j                   j                  dk(  r2|�0|j                  j.                  d
v r|st        j0                  ||«      }|S )Nr_  rh   r   zéYou are attempting to perform batched generation with padding_side='right' this may lead to unexpected behaviour for Flash Attention version of Mimi. Make sure to  call `tokenizer.padding_side  = 'left'` before tokenizing the input. r�   r`  )Úinputs_embedsÚpast_key_values_lengthr&  Úis_trainingr   )Úsequence_lengthÚtarget_lengthr?   rà   r,  Ú
batch_sizerS   rx  )rZ  Úxpu)rS   rf  ÚsumÚitemr2  r�   r…  rè   r   r   r   Ú_ignore_causal_mask_sdpar&  r1  r?   rà   r(   ÚfinfoÚminri   Úget_max_cache_shaper„   Ú5_prepare_4d_causal_attention_mask_with_cache_positionrÖ   Ú_unmask_unattended)rR   r(  r’  r,  rx  r*  Úis_padding_rightr‹  Úusing_static_cacheÚusing_sliding_window_cacher?   rà   Ú	min_dtyper—  r˜  rB  s                   r.   r‡  z(MimiTransformerModel._update_causal_mask  s  € ð �;‰;×+Ñ+Ð/BÒBØÐ)¨oÐ.IØ#1²!°R°%Ñ#8×#<Ñ#<Ó#>×#CÑ#CÓ#EÈ×IZÑIZÓI\Ð]^ÑI_Ñ#_Ð Ù#Ü$ðaóð ð
 Ð)¨c°^Ñ.CØ%Ð%Øð
 @OÐ?Z˜?×9Ñ9Ô;Ð`aÐÜ'¨¼ÓEÐÜ%/°ÔASÓ%TÐ"ð �K‰K×,Ñ,°Ò6Ù'Ñ+EÙ%ä%×>Ñ>ØØ*Ø'7Ø#Ÿ{™{×9Ñ9Ø ŸM™Mõð à$×*Ñ*¨L×,?Ñ,?ˆvˆÜ—K‘K Ó&×*Ñ*ˆ	Ø&×,Ñ,¨QÑ/ˆá%Ñ);Ø+×?Ñ?ÓA‰Mô
 ˜n¬e¯l©lÔ;ð ×$Ñ$ RÒ(à%¨Ñ7¸!Ñ;ð ð ×PÑPØØ+Ø'ØØØ)Ø#×)Ñ)¨!Ñ,Ø—;‘;Ø+ð Qó 

ˆð �K‰K×,Ñ,°Ò6ØÐ*Ø×%Ñ%×*Ñ*¨oÑ=Ù%ô
 1×CÑCÀKÐQZÓ[ˆKàÐr-   r—  r˜  r?   rà   r™  c	                 óp  — | �| j                  «       dk(  r| }	|	S t        j                  |«      j                  }
t        j                  ||f|
||¬«      }	t        j
                  ||¬«      |j                  dd«      kD  }|j                  �]t        |t        «      r||kD  rHt        j
                  ||¬«      |j                  dd«      |j                  z
  k  }|j                  |«       |	|z  }	|	dddd…dd…f   j                  |ddd«      }	| �©|	j                  «       }	| j                  d   |kD  r| dd…d|…f   } | j                  d   }|	dd…dd…dd…d|…f   | dd…dddd…f   j                  |	j                  «      z   }|dk(  }|	dd…dd…dd…d|…f   j!                  ||
«      |	dd…dd…dd…d|…f<   |	S )aS  
        Creates a causal 4D mask of shape `(batch_size, 1, query_length, key_value_length)` from a 2D mask of shape
        `(batch_size, key_value_length)`, or if the input `attention_mask` is already 4D, do nothing.

        Args:
            attention_mask (`torch.Tensor`):
                A 2D attention mask of shape `(batch_size, key_value_length)` or a 4D attention mask of shape `(batch_size, 1, query_length, key_value_length)`.
            sequence_length (`int`):
                The sequence length being processed.
            target_length (`int`):
                The target length: when generating with static cache, the mask should be as long as the static cache, to account for the 0 padding, the part of the cache that is not filled yet.
            dtype (`torch.dtype`):
                The dtype to use for the 4D attention mask.
            device (`torch.device`):
                The device to place the 4D attention mask on.
            cache_position (`torch.Tensor`):
                Indices depicting the position of the input sequence tokens in the sequence.
            batch_size (`torch.Tensor`):
                Batch size.
            config (`MimiConfig`):
                The model's configuration class
            past_key_values (`Cache`):
                The cache class that is being used currently to generate
        Né   )Ú
fill_valuer?   rà   r|  rh   r   r   )r—   r(   rž  rŸ  rË   r†  r  r&  rè   r   Úbitwise_or_rç   Úcloneri   rk   rà   Úmasked_fill)r(  r—  r˜  r?   rà   r,  r™  rS   rx  rB  r¦  Údiagonal_attend_maskÚsliding_attend_maskÚmask_lengthÚpadding_masks                  r.   r¡  zJMimiTransformerModel._prepare_4d_causal_attention_mask_with_cache_position^  sò  € ðJ Ð%¨.×*<Ñ*<Ó*>À!Ò*Cà(ˆKð: Ðô7 Ÿ™ EÓ*×.Ñ.ˆIÜŸ*™*Ø  -Ð0¸YÈeÐ\bôˆKô $)§<¡<°ÀfÔ#MÐP^×PfÑPfÐgiÐklÓPmÑ#mÐ Ø×$Ñ$Ð0ô " /Ô3EÔFÈ/Ð\iÒJiÜ*/¯,©,°}ÈVÔ*TØ&×.Ñ.¨r°1Ó5¸×8MÑ8MÑMñ+Ð'ð )×4Ñ4Ð5HÔIØÐ/Ñ/ˆKØ% d¨D²!²QÐ&6Ñ7×>Ñ>¸zÈ1ÈbÐRTÓUˆKØÐ)Ø)×/Ñ/Ó1�Ø!×'Ñ'¨Ñ+¨mÒ;Ø%3²A°~¸°~Ð4EÑ%F�NØ,×2Ñ2°2Ñ6�Ø*ª1ªa²°L°[°LÐ+@ÑAÀNÒSTÐVZÐ\`ÒbcÐScÑDd×DgÑDgØ×&Ñ&óEñ  �ð  ,¨qÑ0�Ø5@ÂÂAÂqÈ,È;È,ÐAVÑ5W×5cÑ5cØ  )ó6�šAšq¢! \ k \Ð1Ñ2ð Ðr-   )	NNNNNNNNN)F)r$   r%   r&   r'   r   rE   r   r(   r)   r„   r   r   r   r+   rƒ   r   r   r   r‡  r…   r‚   r?   rà   r¡  rˆ   r‰   s   @r.   rt  rt  R  sÕ  ø„ ñð	˜zõ 	ð 59Ø15Ø37ØKOØ$(Ø,0Ø/3Ø&*Ø59ñc
à × 0Ñ 0Ñ1ðc
ð ! §¡Ñ.ðc
ð ˜u×/Ñ/Ñ0ð	c
ð
 " %¨¨t°E×4EÑ4EÑ/FÐ(FÑ"GÑHðc
ð ˜D‘>ðc
ð $ D™>ðc
ð ' t™nðc
ð ˜d‘^ðc
ð ! ×!1Ñ!1Ñ2ðc
ð 
ˆuÐ-Ð-Ñ	.óc
ðX #(ñQàŸ™ðQð —l‘lðQð Ÿ™ð	Qð
 ðQð  óQðf ðBØŸ™ðBàðBð ðBð �{‰{ð	Bð
 —‘ðBð Ÿ™ðBð ðBð ðBð òBó ôBr-   rt  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.rS   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$                  |«      | _        y )NrC   r°   r   )rD   rE   r‚   r›   r´   r4   r¸   r²   r7   r   rŸ   r‹   rµ   r¶   r–   r·   r±   r¹   r    rº   )rR   rS   r¼   r»   r½   r¾   r¿   rT   s          €r.   rE   zMimiDecoder.__init__¨  s[  ø€ Ü‰ÑÔÜ�aœ3˜v×7Ñ7Ó8Ñ8Ó9ˆÜ˜F F×$6Ñ$6¸À&×BTÑBTÑ8TÐV\×VhÑVhÓiÐjˆð ×-Ñ-ò 
	ˆEØ# f×&8Ñ&8Ñ8ˆMà”b—f‘f“h�ZÑˆEØÜ# F¨M¸=ÈAÑ;MÐ[`ÐcdÑ[dÐmrÔsðñ ˆEô ˜6×5Ñ5Ó6ò l�Øœ/¨&°-À1Ñ2DÀv×GbÑGbÐdeÑGeÐghÐFiÓjÐkÑk‘ðlà˜‰M‰Gð
	ð 	”"—&‘&“(�ÑˆØ”*˜V V×%7Ñ%7¸×9NÑ9NÐPV×PgÑPgÓhÐiÑiˆÜ—m‘m EÓ*ˆ�r-   c                 ó8   — | j                   D ]
  } ||«      }Œ |S r`   rÁ   rÂ   s      r.   r   zMimiDecoder.forwardÀ  rÃ   r-   rÄ   r‰   s   @r.   r²  r²  ¥  s   ø„ Ù)ð+˜zõ +ö0r-   r²  c                   ój   ‡ — 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.rS   Úepsilonc                 ó¢  •— t         ‰| �  «        t        j                  |j                  |j
                  «      }|j                  | _        | j                  dt        j                  dgt        j                  ¬«      «       | j                  dt        j                  |j                  «      «       | j                  d|«       d | _
        || _        y )NÚinitializedTr>   Úcluster_usageÚ	embed_sum)rD   rE   r(   ÚzerosÚcodebook_sizeÚcodebook_dimrO   rM   r7  ÚonesÚ_embedr·  )rR   rS   r·  ÚembedrT   s       €r.   rE   zMimiEuclideanCodebook.__init__É  s–   ø€ Ü‰ÑÔÜ—‘˜F×0Ñ0°&×2EÑ2EÓFˆà#×1Ñ1ˆÔà×Ñ˜]¬E¯L©L¸$¸ÄuÇ}Á}Ô,UÔVØ×Ñ˜_¬e¯j©j¸×9MÑ9MÓ.NÔOØ×Ñ˜[¨%Ô0ØˆŒØˆ�r-   rf   c                 ó°   — | j                   €?| j                  | j                  j                  | j                  ¬«      d d …d f   z  | _         | j                   S )N)rŸ  )rÀ  r»  rº  Úclampr·  rc   s    r.   rÁ  zMimiEuclideanCodebook.embedÕ  sJ   € à�;‰;ÐØŸ.™.¨4×+=Ñ+=×+CÑ+CÈÏÉÐ+CÓ+UÒVWÐY]ÐV]Ñ+^Ñ^ˆDŒKØ�{‰{Ðr-   c                 ó€   — t        j                  |d    | j                  d    d¬«      d   }|j                  d¬«      }|S )NrC   )r0  r   rh   ræ   )r(   ÚcdistrÁ  Úargmin)rR   re   ÚdistsÚ	embed_inds       r.   ÚquantizezMimiEuclideanCodebook.quantizeÛ  s?   € ô —‘˜M¨$Ñ/°·±¸DÑ1AÀQÔGÈÑJˆØ—L‘L R�LÓ(ˆ	ØÐr-   c                 ó�   — |j                   }|j                  d|d   f«      }| j                  |«      } |j                  |d d Ž }|S )Nrh   )ri   r  rÉ  r3  )rR   re   ri   rÈ  s       r.   ÚencodezMimiEuclideanCodebook.encodeã  sO   € Ø×#Ñ#ˆà%×-Ñ-¨r°5¸±9¨oÓ>ˆà—M‘M -Ó0ˆ	à"�I—N‘N E¨#¨2 JÐ/ˆ	ØÐr-   c                 óZ   — t         j                  j                  || j                  «      }|S r`   )r   ru   Ú	embeddingrÁ  ©rR   rÈ  rÉ  s      r.   ÚdecodezMimiEuclideanCodebook.decodeî  s!   € Ü—=‘=×*Ñ*¨9°d·j±jÓAˆØˆr-   )gñhãˆµøä>)r$   r%   r&   r'   r   r‡   rE   Úpropertyr(   r„   rÁ  rÉ  rË  rÏ  rˆ   r‰   s   @r.   r¶  r¶  Æ  sG   ø„ Ù+ñ
˜zð 
°Eõ 
ð ð�u—|‘|ò ó ðò
òö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.
    rS   c                 óB   •— t         ‰| �  «        t        |«      | _        y r`   )rD   rE   r¶  Úcodebookr
  s     €r.   rE   zMimiVectorQuantization.__init__ù  s   ø€ Ü‰ÑÔÜ-¨fÓ5ˆ�r-   c                 ób   — |j                  ddd«      }| j                  j                  |«      }|S ©Nr   rC   r   )ÚpermuterÔ  rË  )rR   re   Úembed_ins      r.   rË  zMimiVectorQuantization.encodeý  s/   € Ø%×-Ñ-¨a°°AÓ6ˆØ—=‘=×'Ñ'¨Ó6ˆØˆr-   c                 ób   — | j                   j                  |«      }|j                  ddd«      }|S rÖ  )rÔ  rÏ  r×  rÎ  s      r.   rÏ  zMimiVectorQuantization.decode  s/   € Ø—=‘=×'Ñ'¨	Ó2ˆØ×#Ñ# A q¨!Ó,ˆØˆr-   )	r$   r%   r&   r'   r   rE   rË  rÏ  rˆ   r‰   s   @r.   rÒ  rÒ  ô  s   ø„ ñð6˜zõ 6òö
r-   rÒ  c                   ó°   ‡ — e Zd ZdZd
dedee   fˆ fd„Zd
dej                  dee   dej                  fd„Z
dej                  dej                  fd	„Zˆ xZS )ÚMimiResidualVectorQuantizerzResidual Vector Quantizer.rS   Únum_quantizersc                 ób  •— t         ‰| �  «        |j                  | _        |j                  | _        |�|n|j                  | _        t        j                  t        | j                  «      D �cg c]  }t        |«      ‘Œ c}«      | _	        d | _
        d | _        |j                  |j                  k7  ryt        j
                  j                  |j                  |j                  dd¬«      | _
        t        j
                  j                  |j                  |j                  dd¬«      | _        y y c c}w )Nr   Fr  )rD   rE   r½  Ú
frame_raterÜ  r   r    rµ   rÒ  rº   Ú
input_projÚoutput_projÚ$vector_quantization_hidden_dimensionr¸   r(   rK   )rR   rS   rÜ  r<  rT   s       €r.   rE   z$MimiResidualVectorQuantizer.__init__  sþ   ø€ Ü‰ÑÔØ#×1Ñ1ˆÔØ ×+Ñ+ˆŒØ0>Ð0J™nÐPV×PeÑPeˆÔÜ—m‘mÌUÐSW×SfÑSfÓMgÖ$hÈÔ%;¸FÕ%CÒ$hÓiˆŒàˆŒØˆÔØ×6Ñ6¸&×:LÑ:LÒLÜ#Ÿh™hŸo™oØ×"Ñ" F×$OÑ$OÐQRÐY^ð .ó ˆDŒOô  %Ÿx™xŸ™Ø×;Ñ;¸V×=OÑ=OÐQRÐY^ð  /ó  ˆDÕð	 Mùò	 %is   Á-D,Ú
embeddingsrf   c                 ó*  — | j                   �| j                  |«      }|�|n| j                  }|}g }| j                  d| D ]:  }|j                  |«      }|j	                  |«      }||z
  }|j                  |«       Œ< t        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Ï  Úappendr(   Ústack)	rR   râ  rÜ  r«   Úall_indicesr¬   ÚindicesÚ	quantizedÚout_indicess	            r.   rË  z"MimiResidualVectorQuantizer.encode  sœ   € ð
 �?‰?Ð&ØŸ™¨Ó4ˆJà+9Ð+E™È4×K^ÑK^ˆàˆØˆØ—[‘[  .Ð1ò 	(ˆEØ—l‘l 8Ó,ˆGØŸ™ WÓ-ˆIØ )Ñ+ˆHØ×Ñ˜wÕ'ð		(ô
 —k‘k +Ó.ˆØÐr-   Úcodesc                 ó  — t        j                  d|j                  ¬«      }|j                  dd«      }t	        |«      D ]*  \  }}| j
                  |   }|j                  |«      }||z   }Œ, | j                  �| j                  |«      }|S )zJDecode the given codes of shape [B, K, T] to the quantized representation.r�   r|  r   r   )r(   rM   rà   rê   r�   rº   rÏ  rà  )rR   rë  Úquantized_outr§   rè  r¬   ré  s          r.   rÏ  z"MimiResidualVectorQuantizer.decode0  s‡   € äŸ™ S°·±Ô>ˆØ—‘  1Ó%ˆÜ# EÓ*ò 	6‰JˆAˆwØ—K‘K ‘NˆEØŸ™ WÓ-ˆIØ)¨IÑ5‰Mð	6ð
 ×ÑÐ'Ø ×,Ñ,¨]Ó;ˆMØÐr-   r`   )r$   r%   r&   r'   r   r   r‚   rE   r(   r„   rË  rÏ  rˆ   r‰   s   @r.   rÛ  rÛ    sa   ø„ Ù$ñ˜zð ¸8ÀC¹=õ ñ" §¡ð ¸xÈ¹}ð ÐX]×XdÑXdó ð(˜EŸL™Lð ¨U¯\©\÷ r-   rÛ  c                   ó¤   ‡ — e Zd ZdZdefˆ fd„Zd
dej                  dee	   dej                  fd„Z
dej                  dej                  fd	„Zˆ xZS )Ú MimiSplitResidualVectorQuantizerz Split Residual Vector Quantizer.rS   c                 óR  •— t         ‰| �  «        |j                  | _        |j                  | _        |j                  | _        |j                  | _        |j                  |j                  z
  | _        t        || j                  «      | _	        t        || j                  «      | _
        y r`   )rD   rE   r½  rÞ  rÜ  Úmax_num_quantizersÚnum_semantic_quantizersÚnum_acoustic_quantizersrÛ  Ú"semantic_residual_vector_quantizerÚ"acoustic_residual_vector_quantizerr
  s     €r.   rE   z)MimiSplitResidualVectorQuantizer.__init__A  sŠ   ø€ Ü‰ÑÔØ#×1Ñ1ˆÔØ ×+Ñ+ˆŒØ"(×"7Ñ"7ˆÔà'-×'EÑ'EˆÔ$Ø'-×'<Ñ'<¸v×?]Ñ?]Ñ']ˆÔ$ä2MÈfÐVZ×VrÑVrÓ2sˆÔ/Ü2MÈfÐVZ×VrÑVrÓ2sˆÕ/r-   râ  rÜ  rf   c                 ó¬  — |€| j                   n|}|| j                   kD  rt        d| j                   › d|› d�«      ‚|| j                  k  rt        d| j                  › d|› d�«      ‚| j                  j	                  |«      }|| j                  kD  rC| j
                  j	                  ||| j                  z
  ¬«      }t        j                  ||gd¬«      }|S )rä  úcThe number of quantizers (i.e codebooks) asked should be lower than the total number of quantizers ú, but is currently rM  zgThe number of quantizers (i.e codebooks) asked should be higher than the number of semantic quantizers )rÜ  r   ræ   )rñ  r�   rò  rô  rË  rõ  r(   rë   )rR   râ  rÜ  rë  Úacoustic_codess        r.   rË  z'MimiSplitResidualVectorQuantizer.encodeM  s8  € ð 5CÐ4J˜×0Ò0ÐP^ˆà˜D×3Ñ3Ò3ÜØuÐvz÷  wNñ  wNð  vOð  Obð  cqð  brð  rsð  tóð ð ˜D×8Ñ8Ò8ÜØyÐz~÷  {Wñ  {Wð  zXð  Xkð  lzð  k{ð  {|ð  }óð ð
 ×7Ñ7×>Ñ>¸zÓJˆà˜D×8Ñ8Ò8Ø!×DÑD×KÑKØ¨>¸D×<XÑ<XÑ+Xð Ló ˆNô —I‘I˜u nÐ5¸1Ô=ˆEàˆr-   rë  c                 óü   — | j                   j                  |dd…d| j                  …f   «      }|j                  d   | j                  kD  r1|| j                  j                  |dd…| j                  d…f   «      z  }|S )z7Decode the given codes to the quantized representation.Nr   )rô  rÏ  rò  ri   rõ  )rR   rë  rí  s      r.   rÏ  z'MimiSplitResidualVectorQuantizer.decodej  s   € ð ×?Ñ?×FÑFÀuÊQÐPnÐRV×RnÑRnÐPnÐMnÑGoÓpˆð �;‰;�q‰>˜D×8Ñ8Ò8Ø˜T×DÑD×KÑKÈEÒRSÐUY×UqÑUqÑUsÐRsÑLtÓuÑuˆMØÐr-   r`   )r$   r%   r&   r'   r   rE   r(   r„   r   r‡   rË  rÏ  rˆ   r‰   s   @r.   rï  rï  >  sX   ø„ Ù*ð
t˜zõ 
tñ §¡ð ¸xÈ¹ð ÐZ_×ZfÑZfó ð:	˜EŸL™Lð 	¨U¯\©\÷ 	r-   rï  c                   ó@   — e Zd ZdZeZdZdZdZdgZ	dZ
dZdZdZdZd„ Zy)	ÚMimiPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    ÚmimiÚinput_valuesTÚMimiDecoderLayerrx  c                 óŽ  — t        |t        j                  «      rm|j                  j                  j                  d| j                  j                  ¬«       |j                  �%|j                  j                  j                  «        yyt        |t        j                  t        j                  f«      rJ|j                  j                  j                  «        |j                  j                  j                  d«       yt        |t        j                  «      r t        j                  j                  |j                  «       |j                  �jt!        j"                  |j$                  |j&                  |j(                  d   z  z  «      }t        j                  j+                  |j                  | |¬«       yyt        |t        j,                  «      rz|j                  j                  j                  d| j                  j                  ¬«       |j.                  �2|j                  j                  |j.                     j                  «        yyt        |t        j0                  «      rb|j3                  «       D ]N  \  }}d|v r t        j                  j5                  |«       Œ*d|v sŒ/t        j                  j7                  |d«       ŒP yy)	zInitialize the weightsr�   )ÚmeanÚstdNr�   r   )ÚaÚbrQ  r;   )rè   r   r  rQ  ÚdataÚnormal_rS   Úinitializer_ranger;   Úzero_ri  Ú	GroupNormÚfill_rK   ÚinitÚkaiming_normal_r‘   r  r:   r5   r7   Úuniform_Ú	EmbeddingÚpadding_idxÚLSTMÚnamed_parametersÚxavier_uniform_Ú	constant_)rR   Úmodulerû   ÚnameÚparams        r.   Ú_init_weightsz!MimiPreTrainedModel._init_weightsˆ  sî  € ä�fœbŸi™iÔ(Ø�M‰M×Ñ×&Ñ&¨C°T·[±[×5RÑ5RÐ&ÔSØ�{‰{Ð&Ø—‘× Ñ ×&Ñ&Õ(ð 'ä˜¤§¡¬r¯|©|Ð <Ô=Ø�K‰K×Ñ×"Ñ"Ô$Ø�M‰M×Ñ×$Ñ$ SÕ)Ü˜¤§	¡	Ô*Ü�G‰G×#Ñ# F§M¡MÔ2Ø�{‰{Ð&Ü—I‘I˜fŸm™m¨v×/AÑ/AÀF×DVÑDVÐWXÑDYÑ/YÑZÓ[�Ü—‘× Ñ  §¡°°°aÐ Õ8ð 'ô ˜¤§¡Ô-Ø�M‰M×Ñ×&Ñ&¨C°T·[±[×5RÑ5RÐ&ÔSØ×!Ñ!Ð-Ø—‘×"Ñ" 6×#5Ñ#5Ñ6×<Ñ<Õ>ð .ä˜¤§¡Ô(Ø%×6Ñ6Ó8ò 2‘��eØ˜tÑ#Ü—G‘G×+Ñ+¨EÕ2Ø˜t’^Ü—G‘G×%Ñ% e¨SÕ1ñ	2ð )r-   N)r$   r%   r&   r'   r   Úconfig_classÚbase_model_prefixÚmain_input_nameÚsupports_gradient_checkpointingÚ_no_split_modulesÚ_skip_keys_device_placementÚ_supports_flash_attn_2Ú_supports_sdpaÚ_supports_cache_classÚ_supports_static_cacher  r,   r-   r.   rü  rü  v  sJ   „ ñð
 €LØÐØ$€OØ&*Ð#Ø+Ð,ÐØ"3ÐØ!ÐØ€NØ ÐØ!Ðó2r-   rü  aI  
    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 ([`MimiConfig`]):
            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.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*):
            Discret 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).
        return_dict (`bool`, *optional*):
            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
z"The Mimi neural audio codec model.c                   ó   ‡ — e Zd Zdefˆ fd„Zd„ Zd„ Z	 	 ddej                  de	de	de
eeeej                     f      d	e
e   d
eej                  e
ej                     f   fd„Z	 	 	 	 ddej                  de
ej                     de
e   de
eeeej                     f      d	e
e   d
eeej                  e
ej                     f   ef   fd„Z	 	 ddej                  de
eeeej                     f      d	e
e   d
ej                  fd„Z	 	 	 ddej                  de
ej                     de
eeeej                     f      d	e
e   d
eeej                  ej                  f   ef   f
d„Z ee«       eee¬«      	 	 	 	 	 	 ddej                  de
ej                     de
e	   de
ej                     de
eeeej                     f      de
eeeej                     f      d	e
e   d
eeej                  ej                  f   ef   fd„«       «       Zˆ xZS )Ú	MimiModelrS   c           
      ó\  •— t         ‰| �  |«       || _        t        |«      | _        t        |«      | _        d | _        d | _        |j                  |j                  k7  r¦t        ||j                  |j                  dt        |j                  |j                  z  «      z  ddd¬«      | _        t        ||j                  |j                  dt        |j                  |j                  z  «      z  dd|j                  ¬«      | _        t        |«      | _        t#        |«      | _        t'        |«      | _        t        t+        j,                  | j                  j.                  «      «      | _        d| j0                  z  | j                  j.                  k7  rt3        d«      ‚| j5                  «        y )NrC   FÚ	replicate)r7   r8   r;   rH   )r7   r8   r;   r:   z'The codebook_size must be a power of 2.)rD   rE   rS   r®   Úencoderrt  Úencoder_transformerÚ
downsampleÚupsamplerÞ  Úencodec_frame_rater4   r¸   r‚   r‹   Úupsample_groupsÚdecoder_transformerr²  Údecoderrï  Ú	quantizerr‘   Úlog2r½  Úbits_per_codebookr�   Ú	post_initr
  s     €r.   rE   zMimiModel.__init__Ø  sf  ø€ Ü‰Ñ˜Ô ØˆŒä" 6Ó*ˆŒÜ#7¸Ó#?ˆÔ àˆŒØˆŒØ×Ñ × 9Ñ 9Ò9Ü(ØØ×"Ñ"Ø×"Ñ"Ø¤ F×$=Ñ$=À×@QÑ@QÑ$QÓ RÑRØØØ$ôˆDŒOô 0ØØ×"Ñ"Ø×"Ñ"Ø¤ F×$=Ñ$=À×@QÑ@QÑ$QÓ RÑRØØØ×-Ñ-ôˆDŒMô $8¸Ó#?ˆÔ Ü" 6Ó*ˆŒä9¸&ÓAˆŒä!$¤T§Y¡Y¨t¯{©{×/HÑ/HÓ%IÓ!JˆÔØˆd×$Ñ$Ñ$¨¯©×(AÑ(AÒAÜÐFÓGÐGð 	�‰Õr-   c                 ó   — | j                   S r`   )r&  rc   s    r.   Úget_encoderzMimiModel.get_encoder  ó   € Ø�|‰|Ðr-   c                 ó   — | j                   S r`   )r-  rc   s    r.   Úget_decoderzMimiModel.get_decoder  r4  r-   rþ  rÜ  r°  rx  rz  rf   c                 ój  — | j                  |«      }| j                  |j                  dd«      ||¬«      }|r|j                  d«      }nt	        |«      dkD  r|d   }|d   j                  dd«      }| j                  |«      }| j                  j                  ||«      }|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   rC   ©rx  rz  rx  r   )r&  r'  rê   rÙ   r›   r(  r.  rË  )	rR   rþ  rÜ  r°  rx  rz  râ  Úencoder_outputsrë  s	            r.   Ú_encode_framezMimiModel._encode_frame  s¿   € ð —\‘\ ,Ó/ˆ
Ø×2Ñ2Ø× Ñ   AÓ&¸ÐU`ð 3ó 
ˆñ Ø-×1Ñ1Ð2CÓD‰OÜ�Ó! AÒ%Ø-¨aÑ0ˆOØ$ QÑ'×1Ñ1°!°QÓ7ˆ
Ø—_‘_ ZÓ0ˆ
à—‘×%Ñ% j°.ÓAˆØ—‘  1Ó%ˆØ�oÐ%Ð%r-   r"   c                 óô  — |�|n| j                   j                  }|€| j                   j                  n|}|| j                   j                  kD  r&t        d| j                   j                  › d|› d�«      ‚|j                  \  }}}|dk  s|dkD  rt        d|› �«      ‚|€#t        j                  |«      j                  «       }| j                  |||j                  «       ||¬«      \  }	}|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.
        r÷  rø  rM  r   rC   z1Number of audio channels must be 1 or 2, but got r8  )
rS   rz  rÜ  r�   ri   r(   Ú	ones_likerƒ   r:  r0   )
rR   rþ  r°  rÜ  r"   rz  r<  rÍ   Úinput_lengthÚencoded_framess
             r.   rË  zMimiModel.encode"  sH  € ðB &1Ð%<‘kÀ$Ç+Á+×BYÑBYˆà7EÐ7M˜Ÿ™×3Ò3ÐSaˆà˜DŸK™K×6Ñ6Ò6ÜØuÐvz÷  wBñ  wB÷  wQñ  wQð  vRð  Reð  ftð  euð  uvð  wóð ð %1×$6Ñ$6Ñ!ˆˆ8�\à�aŠ<˜8 aš<ÜÐPÐQYÐPZÐ[Ó\Ð\àÐÜ Ÿ?™?¨<Ó8×=Ñ=Ó?ˆLà26×2DÑ2DØØØ×ÑÓØ3Ø#ð 3Eó 3
Ñ/ˆÐ/ñ àØ'ðð ô
 ! Ð1HÓIÐIr-   rë  c                 óD  — | j                   j                  |«      }| j                  |«      }| j                  |j	                  dd«      ||¬«      }|r|j                  d«      }nt        |«      dkD  r|d   }|d   j	                  dd«      }| j                  |«      }||fS )Nr   rC   r8  rx  r   )r.  rÏ  r)  r,  rê   rÙ   r›   r-  )rR   rë  rx  rz  râ  Údecoder_outputsrr  s          r.   Ú_decode_framezMimiModel._decode_framed  s­   € ð —^‘^×*Ñ*¨5Ó1ˆ
à—]‘] :Ó.ˆ
Ø×2Ñ2Ø× Ñ   AÓ&¸ÐU`ð 3ó 
ˆñ Ø-×1Ñ1Ð2CÓD‰OÜ�Ó! AÒ%Ø-¨aÑ0ˆOØ$ QÑ'×1Ñ1°!°QÓ7ˆ
Ø—,‘,˜zÓ*ˆØ˜Ð'Ð'r-   r    r#   c                 óö   — |�|n| j                   j                  }| j                  |||¬«      \  }}|�5|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*):
                Discret 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.

        Nr8  rh   .)rS   rz  rA  ri   r2   )rR   r    r°  r#   rz  r!   s         r.   rÏ  zMimiModel.decodex  s§   € ð> &1Ð%<‘kÀ$Ç+Á+×BYÑBYˆà04×0BÑ0BØÐ)@Èkð 1Có 1
Ñ-ˆÐ-ð
 Ð#¨×(:Ñ(:¸2Ñ(>À×ASÑASÐTVÑAWÒ(WØ'¨Ð-E¨|×/AÑ/AÀ"Ñ/EÐ-EÐ(EÑFˆLáàØ'ðð ô ! Ð/FÓGÐGr-   )Úoutput_typer  c                 ó¸  — |�|n| j                   j                  }|€#t        j                  |«      j	                  «       }|€B| j                  |||||¬«      }|d   }|r|j                  d«      }nt        |«      dkD  r|d   }| j                  ||||¬«      }	|	d   }
|r|	j                  d«      }nt        |	«      dkD  r|	d   }|s||
||fS t        ||
||¬«      S )aÔ  
        Returns:

        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
        ```)rz  r   rx  r   )r    r!   r"   r#   )
rS   rz  r(   r<  rƒ   rË  rÙ   r›   rÏ  r   )rR   rþ  r°  rÜ  r    r"   r#   rz  r9  r@  r!   s              r.   r   zMimiModel.forward¨  s  € ðD &1Ð%<‘kÀ$Ç+Á+×BYÑBYˆàÐÜ Ÿ?™?¨<Ó8×=Ñ=Ó?ˆLàÐØ"Ÿk™kØ˜l¨NÐ<SÐalð *ó ˆOð *¨!Ñ,ˆKÙØ*9×*=Ñ*=Ð>OÓ*PÑ'Ü�_Ó%¨Ò)Ø*9¸!Ñ*<Ð'àŸ+™+ k°<ÐAXÐfq˜+ÓrˆØ& qÑ)ˆÙØ&5×&9Ñ&9Ð:KÓ&LÑ#Ü�Ó! AÒ%Ø&5°aÑ&8Ð#áØ Ð/FÐH_Ð`Ð`äØ#Ø%Ø$;Ø$;ô	
ð 	
r-   )NN)NNNN)NNN)NNNNNN)r$   r%   r&   r   rE   r3  r6  r(   r„   r‚   r   r   r   r   r+   rƒ   r   r:  r‡   r0   rË  rA  r2   rÏ  r   ÚMIMI_INPUTS_DOCSTRINGr   r   Ú_CONFIG_FOR_DOCr   rˆ   r‰   s   @r.   r#  r#  Ó  s.  ø„ ð
(˜zõ (òTòð LPØ&*ñ&à—l‘lð&ð ð&ð ð	&ð
 " %¨¨t°E×4EÑ4EÑ/FÐ(FÑ"GÑHð&ð ˜d‘^ð&ð 
ˆu�|‰|˜X e§l¡lÑ3Ð3Ñ	4ó&ð: 04Ø*.ØSWØ&*ñ@Jà—l‘lð@Jð ˜uŸ|™|Ñ,ð@Jð ! ™ð	@Jð
 "*¨%°°t¸E×<MÑ<MÑ7NÐ0NÑ*OÑ!Pð@Jð ˜d‘^ð@Jð 
ˆu�U—\‘\ 8¨E¯L©LÑ#9Ð9Ñ:Ð<MÐMÑ	Nó@JðJ LPØ&*ñ	(à�|‰|ð(ð " %¨¨t°E×4EÑ4EÑ/FÐ(FÑ"GÑHð(ð ˜d‘^ð	(ð
 
�‰ó(ð. 04ØSWØ&*ñ.Hà—\‘\ð.Hð ˜uŸ|™|Ñ,ð.Hð "*¨%°°t¸E×<MÑ<MÑ7NÐ0NÑ*OÑ!Pð	.Hð
 ˜d‘^ð.Hð 
ˆu�U—\‘\ 5§<¡<Ð/Ñ0Ð2CÐCÑ	Dó.Hñ` +Ð+@ÓAÙ¨:ÀOÔTð 04Ø(,Ø.2ØSWØSWØ&*ñ>
à—l‘lð>
ð ˜uŸ|™|Ñ,ð>
ð ! ™ð	>
ð
 ˜eŸl™lÑ+ð>
ð "*¨%°°t¸E×<MÑ<MÑ7NÐ0NÑ*OÑ!Pð>
ð "*¨%°°t¸E×<MÑ<MÑ7NÐ0NÑ*OÑ!Pð>
ð ˜d‘^ð>
ð 
ˆu�U—\‘\ 5§<¡<Ð/Ñ0°*Ð<Ñ	=ò>
ó Uó Bô>
r-   r#  )Nr   )Lr'   r‘   Údataclassesr   Útypingr   r   r   r   r(   Útorch.utils.checkpointr   Úactivationsr
   Úcache_utilsr   r   r   r   Úmodeling_attn_mask_utilsr   Úmodeling_flash_attention_utilsr   r   Úmodeling_outputsr   Úmodeling_rope_utilsr   r   Úmodeling_utilsr   rY   r   r   r   r   r   Úconfiguration_mimir   r   Ú
get_loggerr$   rI   rF  r   r0   r2   ÚModuler4   r‹   r–   r®   rÆ   rÒ   r÷   rÿ   r  r„   r‚   r  r  rF  rW  re  rb  rt  r²  r¶  rÒ  rÛ  rï  rü  ÚMIMI_START_DOCSTRINGrE  r#  Ú__all__r,   r-   r.   ú<module>rV     s~  ðñ ã Ý !ß /Ó /ã Û Ý å !ß OÓ OÝ >ß hÝ 7ß KÝ -÷õ õ +ñ ÔÝJà	ˆ×	Ñ	˜HÓ	%€ð €ð ôT�ó Tó ðTð> ôT˜ó Tó ðTð& ôT˜ó Tó ðTô&d�—‘ô dôN7˜"Ÿ)™)ô 7ôv7�b—i‘iô 7ôB�"—)‘)ô ô>�R—Y‘Yô ô <˜"Ÿ)™)ô <òF(óô6ˆb�i‰iô ð"	U˜UŸ\™\ð 	U°#ð 	U¸%¿,¹,ó 	Uô^9�B—I‘Iô ^9ôFk9˜-ô k9ô`U1˜ô U1ðr Ø,ØñÐ ôJ˜2Ÿ9™9ô JôZP˜2Ÿ9™9ô Pôf
�"—)‘)ô ôB*˜BŸI™Iô *ô\˜RŸY™Yô ô(3 "§)¡)ô 3ôl5 r§y¡yô 5ôp)2˜/ô )2ðXÐ ð"Ð ñ@ Ø(ØóôQ
Ð#ó Q
ó	ðQ
ðh Ð-Ð
.�r-   