Skip to main content

eredu_codec/
mimi.rs

1//! Mimi neural audio tokenizer support.
2//!
3//! Mimi is the neural audio codec used by Moshi-family realtime models. This
4//! module implements backend-neutral checkpoint parameters, the split residual
5//! vector quantizer, and the non-streaming SEANet/transformer encoder and
6//! decoder used to map between PCM and Mimi codebook tokens.
7
8use eredu_nn::{
9    AttentionMask, Index, LayerNorm, Linear, LinearSpec, PadMode, Parameter, ParameterId,
10    ParameterMetadata, ParameterSpec, ParameterVisitor, ParameterVisitorMut, Parameterized, Rope,
11    Tensor,
12};
13use std::collections::HashMap;
14
15use crate::{AudioTokenizer, AudioTokenizerConfig, Error};
16
17const EPSILON: f32 = 1e-5;
18
19fn parameter_name(prefix: &str, field: &str) -> String {
20    if prefix.is_empty() {
21        field.to_owned()
22    } else {
23        format!("{prefix}.{field}")
24    }
25}
26
27fn parameter_spec(id: &str) -> ParameterSpec {
28    ParameterSpec::trainable(id).expect("Mimi parameter identities are non-empty")
29}
30
31fn unloaded_parameter<T: Tensor>(
32    shape: &[i32],
33    context: &T::Context,
34) -> Result<Parameter<T>, eredu_nn::Error> {
35    Parameter::unloaded(parameter_spec("value"), shape, context)
36}
37
38fn unloaded_linear<T: Tensor>(
39    input: i32,
40    output: i32,
41    bias: bool,
42    context: &T::Context,
43) -> Result<Linear<T>, eredu_nn::Error> {
44    Linear::unloaded(
45        LinearSpec {
46            input,
47            output,
48            weight: parameter_spec("weight"),
49            bias: bias.then(|| parameter_spec("bias")),
50            format: eredu_nn::LinearFormatSpec::unscaled(eredu_checkpoint::LinearFormat::Dense)
51                .unwrap(),
52        },
53        context,
54    )
55}
56
57fn unloaded_layer_norm<T: Tensor>(
58    dimensions: i32,
59    epsilon: f32,
60    context: &T::Context,
61) -> Result<LayerNorm<T>, eredu_nn::Error> {
62    LayerNorm::unloaded(
63        dimensions,
64        epsilon,
65        Some(parameter_spec("weight")),
66        Some(parameter_spec("bias")),
67        context,
68    )
69}
70
71/// Mimi resampling strategy.
72#[derive(Debug, Clone, Copy, Eq, PartialEq)]
73pub enum ResampleMethod {
74    /// Learned convolutional resampling.
75    Conv,
76}
77
78/// Mimi codec configuration.
79#[derive(Debug, Clone)]
80pub struct Config {
81    /// Audio channels.
82    pub channels: i32,
83    /// PCM sample rate.
84    pub sample_rate: f64,
85    /// Codec frame rate.
86    pub frame_rate: f64,
87    /// Whether the original training path renormalized audio.
88    pub renormalize: bool,
89    /// Latent resampling method.
90    pub resample_method: ResampleMethod,
91    /// Active residual codebooks.
92    pub num_codebooks: i32,
93    /// Total codebooks available in the released checkpoint.
94    pub total_codebooks: i32,
95    /// Codebook cardinality.
96    pub bins: i32,
97    /// Codebook embedding dimension.
98    pub quantizer_dim: i32,
99    /// Model latent dimension.
100    pub latent_dim: i32,
101}
102
103impl Config {
104    /// Released Mimi v0.1 defaults, with a caller-selected active codebook count.
105    pub fn v0_1(num_codebooks: Option<i32>) -> Self {
106        Self {
107            channels: 1,
108            sample_rate: 24_000.0,
109            frame_rate: 12.5,
110            renormalize: true,
111            resample_method: ResampleMethod::Conv,
112            num_codebooks: num_codebooks.unwrap_or(16),
113            total_codebooks: 32,
114            bins: 2_048,
115            quantizer_dim: 256,
116            latent_dim: 512,
117        }
118    }
119
120    fn validate(&self) -> Result<(), Error> {
121        if self.channels <= 0
122            || self.sample_rate <= 0.0
123            || self.frame_rate <= 0.0
124            || self.num_codebooks <= 0
125            || self.num_codebooks > self.total_codebooks
126            || self.bins <= 0
127            || self.quantizer_dim <= 0
128            || self.latent_dim <= 0
129        {
130            return Err(Error::InvalidShape(format!(
131                "invalid Mimi config: channels={}, sample_rate={}, frame_rate={}, num_codebooks={}, total_codebooks={}, bins={}, quantizer_dim={}, latent_dim={}",
132                self.channels,
133                self.sample_rate,
134                self.frame_rate,
135                self.num_codebooks,
136                self.total_codebooks,
137                self.bins,
138                self.quantizer_dim,
139                self.latent_dim
140            )));
141        }
142        Ok(())
143    }
144}
145
146/// Mimi audio tokenizer.
147#[derive(Debug, Clone)]
148pub struct Mimi<T: Tensor> {
149    /// Split residual vector quantizer.
150    pub quantizer: SplitResidualVectorQuantizer<T>,
151    encoder: SeaNetEncoder<T>,
152    encoder_transformer: MimiTransformer<T>,
153    downsample: StreamableConv1d<T>,
154    upsample: StreamableConvTranspose1d<T>,
155    decoder_transformer: MimiTransformer<T>,
156    decoder: SeaNetDecoder<T>,
157    config: Config,
158}
159
160impl<T: Tensor> Mimi<T> {
161    /// Creates an unloaded Mimi tokenizer from config.
162    pub fn new(config: Config, context: &T::Context) -> Result<Self, Error> {
163        config.validate()?;
164        Ok(Self {
165            quantizer: SplitResidualVectorQuantizer::unloaded(&config, context)?,
166            encoder: SeaNetEncoder::unloaded(context)?,
167            encoder_transformer: MimiTransformer::unloaded(context)?,
168            downsample: StreamableConv1d::unloaded_with_pad_mode(
169                config.latent_dim,
170                config.latent_dim,
171                4,
172                2,
173                false,
174                PadMode::Edge,
175                context,
176            )?,
177            upsample: StreamableConvTranspose1d::unloaded(
178                config.latent_dim,
179                config.latent_dim,
180                4,
181                2,
182                config.latent_dim,
183                false,
184                context,
185            )?,
186            decoder_transformer: MimiTransformer::unloaded(context)?,
187            decoder: SeaNetDecoder::unloaded(context)?,
188            config,
189        })
190    }
191
192    /// Returns the Mimi configuration.
193    pub fn mimi_config(&self) -> &Config {
194        &self.config
195    }
196
197    /// Strictly replaces every model parameter with backend-native checkpoint
198    /// tensors using Eredu's stable parameter names.
199    ///
200    /// Backend integrations remain responsible for reading their preferred
201    /// checkpoint format and converting weight layouts. Mimi owns parameter
202    /// completeness and shape validation, so those integrations do not need
203    /// to reproduce the architecture or its parameter tree.
204    pub fn load_parameters(
205        &mut self,
206        parameters: impl IntoIterator<Item = (String, T)>,
207    ) -> Result<(), Error> {
208        let mut parameters = parameters.into_iter().collect::<HashMap<_, _>>();
209        let mut missing = Vec::new();
210        let mut mismatch = None;
211        self.visit_mimi_parameters("", &mut |metadata, parameter| {
212            let key = metadata.id.as_str();
213            match parameters.get(key) {
214                None => missing.push(key.to_owned()),
215                Some(value) if value.shape() != parameter.shape() => {
216                    mismatch = Some(format!(
217                        "Mimi checkpoint tensor {key} has shape {:?}, expected {:?}",
218                        value.shape(),
219                        parameter.shape()
220                    ));
221                }
222                Some(_) => {}
223            }
224        });
225        if let Some(mismatch) = mismatch {
226            return Err(Error::InvalidShape(mismatch));
227        }
228        if !missing.is_empty() {
229            missing.sort();
230            return Err(Error::InvalidShape(format!(
231                "Mimi checkpoint is missing {} model tensors: {}",
232                missing.len(),
233                missing.join(", ")
234            )));
235        }
236        self.visit_mimi_parameters_mut("", &mut |metadata, parameter| {
237            *parameter = parameters
238                .remove(metadata.id.as_str())
239                .expect("checkpoint presence was validated before parameter update");
240        });
241        Ok(())
242    }
243
244    /// Encodes latent frames shaped `[batch, 512, frames]` into Mimi tokens.
245    pub fn encode_latent(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
246        self.quantizer.encode(latent, context)
247    }
248
249    /// Encodes PCM shaped `[batch, 1, samples]` into Mimi tokens `[batch, codebooks, frames]`.
250    pub fn encode(&mut self, pcm: &T, context: &T::Context) -> Result<T, Error> {
251        let latent = self.encoder.forward(pcm, context)?;
252        let latent = self.encoder_transformer.forward(&latent, context)?;
253        let latent = self.downsample.forward(&latent, context)?;
254        self.quantizer.encode(&latent, context)
255    }
256
257    /// Resets state used by [`Mimi::encode_step`].
258    pub fn reset_encode_state(&mut self) {
259        self.encoder.reset_state();
260        self.encoder_transformer.reset_state();
261        self.downsample.reset_state();
262    }
263
264    /// Encodes one PCM frame into the next Mimi token frame.
265    ///
266    /// Accepts PCM shaped `[batch, 1, samples]`. Returns `None` until the
267    /// streaming encoder has enough samples to emit a complete codec frame.
268    pub fn encode_step(&mut self, pcm: &T, context: &T::Context) -> Result<Option<T>, Error> {
269        let latent = match self.encoder.step(pcm, context)? {
270            Some(latent) => latent,
271            None => return Ok(None),
272        };
273        let latent = self.encoder_transformer.step(&latent, context)?;
274        let latent = match self.downsample.step(&latent, context)? {
275            Some(latent) => latent,
276            None => return Ok(None),
277        };
278        Ok(Some(
279            self.quantizer
280                .encode(&latent, context)?
281                .squeeze_axes(&[2], context)?,
282        ))
283    }
284
285    /// Decodes Mimi tokens shaped `[batch, codebooks, frames]` into latent frames.
286    pub fn decode_latent(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
287        self.quantizer.decode(codes, context)
288    }
289
290    /// Decodes Mimi tokens shaped `[batch, codebooks, frames]` into PCM `[batch, 1, samples]`.
291    pub fn decode(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
292        let latent = self.quantizer.decode(codes, context)?;
293        let latent = self.upsample.forward(&latent, context)?;
294        let latent = self.decoder_transformer.forward(&latent, context)?;
295        self.decoder.forward(&latent, context)
296    }
297
298    /// Resets state used by [`Mimi::decode_step`].
299    pub fn reset_decode_state(&mut self) {
300        self.upsample.reset_state();
301        self.decoder_transformer.reset_state();
302        self.decoder.reset_state();
303    }
304
305    /// Decodes one Mimi token frame into the next PCM chunk.
306    ///
307    /// Accepts codes shaped `[batch, codebooks]` or `[batch, codebooks, 1]`.
308    pub fn decode_step(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
309        let codes = match codes.shape() {
310            [_, _] => codes.expand_dims(2, context)?,
311            [_, _, 1] => codes.clone(),
312            _ => {
313                return Err(Error::InvalidShape(format!(
314                    "Mimi decode_step expects [batch, codebooks] or [batch, codebooks, 1], got {:?}",
315                    codes.shape()
316                )));
317            }
318        };
319        let latent = self.quantizer.decode(&codes, context)?;
320        let latent = self.upsample.step(&latent, context)?;
321        let latent = self.decoder_transformer.step(&latent, context)?;
322        self.decoder.step(&latent, context)
323    }
324}
325
326/// Backend-independent layout conversion required by a Mimi checkpoint tensor.
327///
328/// Backends use this metadata to convert framework-native checkpoint layouts
329/// into the layouts expected by [`Mimi::load_parameters`].
330#[derive(Debug, Clone, Copy, Eq, PartialEq)]
331pub enum CheckpointTensorLayout {
332    /// The checkpoint and model use the same tensor layout.
333    Identity,
334    /// A rank-three checkpoint tensor must be transposed along these axes.
335    Transpose3d([i32; 3]),
336}
337
338/// Backend-independent loading plan for one tensor in a released Mimi
339/// checkpoint.
340#[derive(Debug, Clone, Eq, PartialEq)]
341pub struct CheckpointTensorPlan {
342    /// Stable [`Mimi`] parameter identity.
343    pub parameter: String,
344    /// Checkpoint-to-model layout conversion required before loading.
345    pub layout: CheckpointTensorLayout,
346}
347
348/// Maps a released Mimi checkpoint tensor name to its stable model parameter
349/// and required layout conversion.
350///
351/// Returns `None` for checkpoint entries that are not model parameters owned
352/// by [`Mimi`]. Reading a checkpoint and performing the planned conversion are
353/// backend responsibilities.
354pub fn checkpoint_tensor_plan(key: &str) -> Option<CheckpointTensorPlan> {
355    let parameter = transform_decoder_key(key)?;
356    let layout = if parameter.ends_with(".weight") && is_conv_weight_key(&parameter) {
357        // PyTorch ConvTranspose1d stores non-depthwise weights as
358        // [input, output/groups, kernel]. All other Mimi convolutions,
359        // including its depthwise top-level upsampler, retain their leading
360        // channel dimension.
361        let axes = if parameter.contains(".upsample.") {
362            [1, 2, 0]
363        } else {
364            [0, 2, 1]
365        };
366        CheckpointTensorLayout::Transpose3d(axes)
367    } else {
368        CheckpointTensorLayout::Identity
369    };
370    Some(CheckpointTensorPlan { parameter, layout })
371}
372
373fn transform_decoder_key(key: &str) -> Option<String> {
374    if key.starts_with("quantizer.") {
375        return Some(key.to_string());
376    }
377    if key == "downsample.conv.conv.conv.weight" {
378        return Some("downsample.weight".to_string());
379    }
380    if let Some(key) = key.strip_prefix("encoder_transformer.transformer.") {
381        let key = key
382            .replace(".self_attn.in_proj_weight", ".self_attn.in_proj.weight")
383            .replace(".linear1.", ".mlp.linear1.")
384            .replace(".linear2.", ".mlp.linear2.");
385        return Some(format!("encoder_transformer.{key}"));
386    }
387    if let Some(key) = key.strip_prefix("encoder.model.") {
388        return transform_seanet_encoder_key(key);
389    }
390    if key == "upsample.convtr.convtr.convtr.weight" {
391        return Some("upsample.weight".to_string());
392    }
393    if let Some(key) = key.strip_prefix("decoder_transformer.transformer.") {
394        let key = key
395            .replace(".self_attn.in_proj_weight", ".self_attn.in_proj.weight")
396            .replace(".linear1.", ".mlp.linear1.")
397            .replace(".linear2.", ".mlp.linear2.");
398        return Some(format!("decoder_transformer.{key}"));
399    }
400    if let Some(key) = key.strip_prefix("decoder.model.") {
401        return transform_seanet_decoder_key(key);
402    }
403    None
404}
405
406fn transform_seanet_encoder_key(key: &str) -> Option<String> {
407    let (source, target) = [
408        ("0.conv.conv.", "encoder.init_conv1d."),
409        (
410            "1.block.1.conv.conv.",
411            "encoder.layers.0.residuals.0.block.0.",
412        ),
413        (
414            "1.block.3.conv.conv.",
415            "encoder.layers.0.residuals.0.block.1.",
416        ),
417        ("3.conv.conv.", "encoder.layers.0.downsample."),
418        (
419            "4.block.1.conv.conv.",
420            "encoder.layers.1.residuals.0.block.0.",
421        ),
422        (
423            "4.block.3.conv.conv.",
424            "encoder.layers.1.residuals.0.block.1.",
425        ),
426        ("6.conv.conv.", "encoder.layers.1.downsample."),
427        (
428            "7.block.1.conv.conv.",
429            "encoder.layers.2.residuals.0.block.0.",
430        ),
431        (
432            "7.block.3.conv.conv.",
433            "encoder.layers.2.residuals.0.block.1.",
434        ),
435        ("9.conv.conv.", "encoder.layers.2.downsample."),
436        (
437            "10.block.1.conv.conv.",
438            "encoder.layers.3.residuals.0.block.0.",
439        ),
440        (
441            "10.block.3.conv.conv.",
442            "encoder.layers.3.residuals.0.block.1.",
443        ),
444        ("12.conv.conv.", "encoder.layers.3.downsample."),
445        ("14.conv.conv.", "encoder.final_conv1d."),
446    ]
447    .into_iter()
448    .find(|(source, _)| key.starts_with(source))?;
449    Some(format!("{target}{}", &key[source.len()..]))
450}
451
452fn transform_seanet_decoder_key(key: &str) -> Option<String> {
453    let (source, target) = [
454        ("0.conv.conv.", "decoder.init_conv1d."),
455        ("2.convtr.convtr.", "decoder.layers.0.upsample."),
456        (
457            "3.block.1.conv.conv.",
458            "decoder.layers.0.residuals.0.block.0.",
459        ),
460        (
461            "3.block.3.conv.conv.",
462            "decoder.layers.0.residuals.0.block.1.",
463        ),
464        ("5.convtr.convtr.", "decoder.layers.1.upsample."),
465        (
466            "6.block.1.conv.conv.",
467            "decoder.layers.1.residuals.0.block.0.",
468        ),
469        (
470            "6.block.3.conv.conv.",
471            "decoder.layers.1.residuals.0.block.1.",
472        ),
473        ("8.convtr.convtr.", "decoder.layers.2.upsample."),
474        (
475            "9.block.1.conv.conv.",
476            "decoder.layers.2.residuals.0.block.0.",
477        ),
478        (
479            "9.block.3.conv.conv.",
480            "decoder.layers.2.residuals.0.block.1.",
481        ),
482        ("11.convtr.convtr.", "decoder.layers.3.upsample."),
483        (
484            "12.block.1.conv.conv.",
485            "decoder.layers.3.residuals.0.block.0.",
486        ),
487        (
488            "12.block.3.conv.conv.",
489            "decoder.layers.3.residuals.0.block.1.",
490        ),
491        ("14.conv.conv.", "decoder.final_conv1d."),
492    ]
493    .into_iter()
494    .find(|(source, _)| key.starts_with(source))?;
495    Some(format!("{target}{}", &key[source.len()..]))
496}
497
498fn is_conv_weight_key(key: &str) -> bool {
499    key.starts_with("upsample.")
500        || key.starts_with("downsample.")
501        || key.contains(".upsample.")
502        || key.contains(".downsample.")
503        || key.contains(".init_conv1d.")
504        || key.contains(".final_conv1d.")
505        || key.contains(".block.")
506}
507
508impl<T: Tensor> AudioTokenizer for Mimi<T> {
509    type Tensor = T;
510
511    fn config(&self) -> AudioTokenizerConfig {
512        AudioTokenizerConfig {
513            sample_rate: self.config.sample_rate,
514            frame_rate: self.config.frame_rate,
515            channels: self.config.channels,
516            codebooks: self.config.num_codebooks,
517            cardinality: self.config.bins,
518        }
519    }
520
521    fn encode(&mut self, pcm: &T, context: &T::Context) -> Result<T, Error> {
522        self.encode(pcm, context)
523    }
524
525    fn decode(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
526        self.decode(codes, context)
527    }
528}
529
530#[derive(Debug, Clone)]
531struct SeaNetEncoder<T: Tensor> {
532    init_conv1d: StreamableConv1d<T>,
533    layers: Vec<EncoderLayer<T>>,
534    final_conv1d: StreamableConv1d<T>,
535}
536
537impl<T: Tensor> SeaNetEncoder<T> {
538    fn unloaded(context: &T::Context) -> Result<Self, Error> {
539        let ratios = [4, 5, 6, 8];
540        let mut channels = 64;
541        let mut layers = Vec::with_capacity(ratios.len());
542        for ratio in ratios {
543            layers.push(EncoderLayer::unloaded(
544                channels,
545                channels * 2,
546                ratio,
547                context,
548            )?);
549            channels *= 2;
550        }
551        Ok(Self {
552            init_conv1d: StreamableConv1d::unloaded(1, 64, 7, 1, context)?,
553            layers,
554            final_conv1d: StreamableConv1d::unloaded(1024, 512, 3, 1, context)?,
555        })
556    }
557
558    fn forward(&mut self, pcm: &T, context: &T::Context) -> Result<T, Error> {
559        validate_pcm(pcm)?;
560        let mut x = self.init_conv1d.forward(pcm, context)?;
561        for layer in &mut self.layers {
562            x = layer.forward(&x, context)?;
563        }
564        self.final_conv1d
565            .forward(&T::elu(&x, 1.0, context)?, context)
566    }
567
568    fn reset_state(&mut self) {
569        self.init_conv1d.reset_state();
570        for layer in &mut self.layers {
571            layer.reset_state();
572        }
573        self.final_conv1d.reset_state();
574    }
575
576    fn step(&mut self, pcm: &T, context: &T::Context) -> Result<Option<T>, Error> {
577        validate_pcm(pcm)?;
578        let mut x = match self.init_conv1d.step(pcm, context)? {
579            Some(x) => x,
580            None => return Ok(None),
581        };
582        for layer in &mut self.layers {
583            x = match layer.step(&x, context)? {
584                Some(x) => x,
585                None => return Ok(None),
586            };
587        }
588        self.final_conv1d.step(&T::elu(&x, 1.0, context)?, context)
589    }
590}
591
592#[derive(Debug, Clone)]
593struct EncoderLayer<T: Tensor> {
594    residuals: Vec<SeaNetResnetBlock<T>>,
595    downsample: StreamableConv1d<T>,
596}
597
598impl<T: Tensor> EncoderLayer<T> {
599    fn unloaded(
600        in_channels: i32,
601        out_channels: i32,
602        ratio: i32,
603        context: &T::Context,
604    ) -> Result<Self, Error> {
605        Ok(Self {
606            residuals: vec![SeaNetResnetBlock::unloaded(in_channels, context)?],
607            downsample: StreamableConv1d::unloaded(
608                in_channels,
609                out_channels,
610                ratio * 2,
611                ratio,
612                context,
613            )?,
614        })
615    }
616
617    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
618        let mut x = x.clone();
619        for residual in &mut self.residuals {
620            x = residual.forward(&x, context)?;
621        }
622        self.downsample.forward(&T::elu(&x, 1.0, context)?, context)
623    }
624
625    fn reset_state(&mut self) {
626        for residual in &mut self.residuals {
627            residual.reset_state();
628        }
629        self.downsample.reset_state();
630    }
631
632    fn step(&mut self, x: &T, context: &T::Context) -> Result<Option<T>, Error> {
633        let mut x = x.clone();
634        for residual in &mut self.residuals {
635            x = residual.step(&x, context)?;
636        }
637        self.downsample.step(&T::elu(&x, 1.0, context)?, context)
638    }
639}
640
641#[derive(Debug, Clone)]
642struct MimiTransformer<T: Tensor> {
643    layers: Vec<MimiTransformerLayer<T>>,
644}
645
646impl<T: Tensor> MimiTransformer<T> {
647    fn unloaded(context: &T::Context) -> Result<Self, Error> {
648        Ok(Self {
649            layers: (0..8)
650                .map(|_| MimiTransformerLayer::unloaded(context))
651                .collect::<Result<Vec<_>, _>>()?,
652        })
653    }
654
655    fn forward(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
656        let mut x = latent.swap_axes(1, 2, context)?;
657        for layer in &mut self.layers {
658            x = layer.forward(&x, context)?;
659        }
660        Ok(x.swap_axes(1, 2, context)?)
661    }
662
663    fn reset_state(&mut self) {
664        for layer in &mut self.layers {
665            layer.reset_state();
666        }
667    }
668
669    fn step(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
670        let mut x = latent.swap_axes(1, 2, context)?;
671        for layer in &mut self.layers {
672            x = layer.step(&x, context)?;
673        }
674        Ok(x.swap_axes(1, 2, context)?)
675    }
676}
677
678#[derive(Debug, Clone)]
679struct MimiTransformerLayer<T: Tensor> {
680    norm1: LayerNorm<T>,
681    norm2: LayerNorm<T>,
682    self_attn: MimiSelfAttention<T>,
683    mlp: MimiMlp<T>,
684    layer_scale_1: LayerScale<T>,
685    layer_scale_2: LayerScale<T>,
686}
687
688impl<T: Tensor> MimiTransformerLayer<T> {
689    fn unloaded(context: &T::Context) -> Result<Self, Error> {
690        Ok(Self {
691            norm1: unloaded_layer_norm(512, 1e-5, context)?,
692            norm2: unloaded_layer_norm(512, 1e-5, context)?,
693            self_attn: MimiSelfAttention::unloaded(context)?,
694            mlp: MimiMlp::unloaded(context)?,
695            layer_scale_1: LayerScale::unloaded(512, context)?,
696            layer_scale_2: LayerScale::unloaded(512, context)?,
697        })
698    }
699
700    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
701        let normed = self.norm1.forward(x, context)?;
702        let attended = self
703            .self_attn
704            .forward(&normed, context)?
705            .multiply(self.layer_scale_1.scale.as_ref(), context)?;
706        let x = x.add(&attended, context)?;
707        let normed = self.norm2.forward(&x, context)?;
708        let mlp = self
709            .mlp
710            .forward(&normed, context)?
711            .multiply(self.layer_scale_2.scale.as_ref(), context)?;
712        Ok(x.add(&mlp, context)?)
713    }
714
715    fn reset_state(&mut self) {
716        self.self_attn.reset_state();
717    }
718
719    fn step(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
720        let normed = self.norm1.forward(x, context)?;
721        let attended = self
722            .self_attn
723            .step(&normed, context)?
724            .multiply(self.layer_scale_1.scale.as_ref(), context)?;
725        let x = x.add(&attended, context)?;
726        let normed = self.norm2.forward(&x, context)?;
727        let mlp = self
728            .mlp
729            .forward(&normed, context)?
730            .multiply(self.layer_scale_2.scale.as_ref(), context)?;
731        Ok(x.add(&mlp, context)?)
732    }
733}
734
735#[derive(Debug, Clone)]
736struct LayerScale<T: Tensor> {
737    scale: Parameter<T>,
738}
739
740impl<T: Tensor> LayerScale<T> {
741    fn unloaded(dim: i32, context: &T::Context) -> Result<Self, Error> {
742        Ok(Self {
743            scale: unloaded_parameter(&[dim], context)?,
744        })
745    }
746}
747
748#[derive(Debug, Clone)]
749struct MimiMlp<T: Tensor> {
750    linear1: Linear<T>,
751    linear2: Linear<T>,
752}
753
754impl<T: Tensor> MimiMlp<T> {
755    fn unloaded(context: &T::Context) -> Result<Self, Error> {
756        Ok(Self {
757            linear1: unloaded_linear(512, 2048, false, context)?,
758            linear2: unloaded_linear(2048, 512, false, context)?,
759        })
760    }
761
762    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
763        let x = self.linear1.forward(x, context)?;
764        let x = T::gelu(&x, context)?;
765        Ok(self.linear2.forward(&x, context)?)
766    }
767}
768
769#[derive(Debug, Clone)]
770struct MimiSelfAttention<T: Tensor> {
771    in_proj: Linear<T>,
772    out_proj: Linear<T>,
773    rope: Rope,
774    num_heads: i32,
775    head_dim: i32,
776    scale: f32,
777    context: i32,
778    key_cache: Option<T>,
779    value_cache: Option<T>,
780}
781
782impl<T: Tensor> MimiSelfAttention<T> {
783    fn unloaded(context: &T::Context) -> Result<Self, Error> {
784        let head_dim = 64;
785        Ok(Self {
786            in_proj: unloaded_linear(512, 1536, false, context)?,
787            out_proj: unloaded_linear(512, 512, false, context)?,
788            rope: Rope::new(head_dim, true, 10_000.0, 1.0),
789            num_heads: 8,
790            head_dim,
791            scale: (head_dim as f32).sqrt().recip(),
792            context: 250,
793            key_cache: None,
794            value_cache: None,
795        })
796    }
797
798    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
799        let shape = x.shape();
800        if shape.len() != 3 || shape[2] != 512 {
801            return Err(Error::InvalidShape(format!(
802                "Mimi decoder transformer expects [batch, frames, 512], got {:?}",
803                x.shape()
804            )));
805        }
806        let (batch, seq, dim) = (shape[0], shape[1], shape[2]);
807        let qkv = self
808            .in_proj
809            .forward(x, context)?
810            .reshape(&[batch, seq, 3, self.num_heads, self.head_dim], context)?;
811        let mut q = qkv
812            .index(
813                &[
814                    Index::Full,
815                    Index::Full,
816                    Index::At(0),
817                    Index::Full,
818                    Index::Full,
819                ],
820                context,
821            )?
822            .transpose_axes(&[0, 2, 1, 3], context)?;
823        let mut k = qkv
824            .index(
825                &[
826                    Index::Full,
827                    Index::Full,
828                    Index::At(1),
829                    Index::Full,
830                    Index::Full,
831                ],
832                context,
833            )?
834            .transpose_axes(&[0, 2, 1, 3], context)?;
835        let v = qkv
836            .index(
837                &[
838                    Index::Full,
839                    Index::Full,
840                    Index::At(2),
841                    Index::Full,
842                    Index::Full,
843                ],
844                context,
845            )?
846            .transpose_axes(&[0, 2, 1, 3], context)?;
847        q = self.rope.forward(&q, 0, context)?;
848        k = self.rope.forward(&k, 0, context)?;
849        let attended = T::scaled_dot_product_attention(
850            &q,
851            &k,
852            &v,
853            self.scale,
854            AttentionMask::Causal,
855            context,
856        )?
857        .transpose_axes(&[0, 2, 1, 3], context)?
858        .reshape(&[batch, seq, dim], context)?;
859        Ok(self.out_proj.forward(&attended, context)?)
860    }
861
862    fn reset_state(&mut self) {
863        self.key_cache = None;
864        self.value_cache = None;
865    }
866
867    fn step(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
868        let shape = x.shape();
869        if shape.len() != 3 || shape[2] != 512 {
870            return Err(Error::InvalidShape(format!(
871                "Mimi decoder transformer step expects [batch, frames, 512], got {:?}",
872                x.shape()
873            )));
874        }
875        let (batch, seq, dim) = (shape[0], shape[1], shape[2]);
876        let prev_len = self
877            .key_cache
878            .as_ref()
879            .map(|cache| cache.dim(2))
880            .unwrap_or(0);
881        let qkv = self
882            .in_proj
883            .forward(x, context)?
884            .reshape(&[batch, seq, 3, self.num_heads, self.head_dim], context)?;
885        let mut q = qkv
886            .index(
887                &[
888                    Index::Full,
889                    Index::Full,
890                    Index::At(0),
891                    Index::Full,
892                    Index::Full,
893                ],
894                context,
895            )?
896            .transpose_axes(&[0, 2, 1, 3], context)?;
897        let mut k = qkv
898            .index(
899                &[
900                    Index::Full,
901                    Index::Full,
902                    Index::At(1),
903                    Index::Full,
904                    Index::Full,
905                ],
906                context,
907            )?
908            .transpose_axes(&[0, 2, 1, 3], context)?;
909        let v = qkv
910            .index(
911                &[
912                    Index::Full,
913                    Index::Full,
914                    Index::At(2),
915                    Index::Full,
916                    Index::Full,
917                ],
918                context,
919            )?
920            .transpose_axes(&[0, 2, 1, 3], context)?;
921        q = self.rope.forward(&q, prev_len, context)?;
922        k = self.rope.forward(&k, prev_len, context)?;
923
924        let mut keys = match self.key_cache.take() {
925            Some(prev) => T::concatenate(&[prev, k], 2, context)?,
926            None => k,
927        };
928        let mut values = match self.value_cache.take() {
929            Some(prev) => T::concatenate(&[prev, v], 2, context)?,
930            None => v,
931        };
932        let key_len = keys.dim(2);
933        if key_len > self.context + seq {
934            let start = key_len - (self.context + seq);
935            keys = keys.index(
936                &[
937                    Index::Full,
938                    Index::Full,
939                    Index::Range(start, key_len),
940                    Index::Full,
941                ],
942                context,
943            )?;
944            values = values.index(
945                &[
946                    Index::Full,
947                    Index::Full,
948                    Index::Range(start, key_len),
949                    Index::Full,
950                ],
951                context,
952            )?;
953        }
954        let retained_prev_len = keys.dim(2) - seq;
955        let mask =
956            streaming_attention_mask::<T>(batch, seq, retained_prev_len, self.context, context)?;
957        let attended = T::scaled_dot_product_attention(
958            &q,
959            &keys,
960            &values,
961            self.scale,
962            AttentionMask::Tensor(&mask),
963            context,
964        )?
965        .transpose_axes(&[0, 2, 1, 3], context)?
966        .reshape(&[batch, seq, dim], context)?;
967        self.key_cache = Some(keys);
968        self.value_cache = Some(values);
969        Ok(self.out_proj.forward(&attended, context)?)
970    }
971}
972
973fn streaming_attention_mask<T: Tensor>(
974    batch: i32,
975    query_len: i32,
976    prev_len: i32,
977    attention_context: i32,
978    execution: &T::Context,
979) -> Result<T, Error> {
980    let key_len = prev_len + query_len;
981    let mut mask = Vec::with_capacity((batch * query_len * key_len) as usize);
982    for _ in 0..batch {
983        for q in 0..query_len {
984            let q_pos = prev_len + q;
985            for k in 0..key_len {
986                if k <= q_pos && q_pos <= k + attention_context {
987                    mask.push(0.0f32);
988                } else {
989                    mask.push(f32::NEG_INFINITY);
990                }
991            }
992        }
993    }
994    Ok(T::from_f32_slice(
995        &mask,
996        &[batch, 1, query_len, key_len],
997        execution,
998    )?)
999}
1000
1001#[derive(Debug, Clone)]
1002struct SeaNetDecoder<T: Tensor> {
1003    init_conv1d: StreamableConv1d<T>,
1004    layers: Vec<DecoderLayer<T>>,
1005    final_conv1d: StreamableConv1d<T>,
1006}
1007
1008impl<T: Tensor> SeaNetDecoder<T> {
1009    fn unloaded(context: &T::Context) -> Result<Self, Error> {
1010        let ratios = [8, 6, 5, 4];
1011        let mut channels = 1024;
1012        let mut layers = Vec::with_capacity(ratios.len());
1013        for ratio in ratios {
1014            let out_channels = channels / 2;
1015            layers.push(DecoderLayer::unloaded(
1016                channels,
1017                out_channels,
1018                ratio,
1019                context,
1020            )?);
1021            channels = out_channels;
1022        }
1023        Ok(Self {
1024            init_conv1d: StreamableConv1d::unloaded(512, 1024, 7, 1, context)?,
1025            layers,
1026            final_conv1d: StreamableConv1d::unloaded(64, 1, 3, 1, context)?,
1027        })
1028    }
1029
1030    fn forward(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
1031        let mut x = self.init_conv1d.forward(latent, context)?;
1032        for layer in &mut self.layers {
1033            x = layer.forward(&T::elu(&x, 1.0, context)?, context)?;
1034        }
1035        self.final_conv1d
1036            .forward(&T::elu(&x, 1.0, context)?, context)
1037    }
1038
1039    fn reset_state(&mut self) {
1040        self.init_conv1d.reset_state();
1041        for layer in &mut self.layers {
1042            layer.reset_state();
1043        }
1044        self.final_conv1d.reset_state();
1045    }
1046
1047    fn step(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
1048        let mut x = self.init_conv1d.step(latent, context)?.ok_or_else(|| {
1049            Error::InvalidShape("Mimi decoder init conv produced no streaming output".into())
1050        })?;
1051        for layer in &mut self.layers {
1052            x = layer.step(&T::elu(&x, 1.0, context)?, context)?;
1053        }
1054        self.final_conv1d
1055            .step(&T::elu(&x, 1.0, context)?, context)?
1056            .ok_or_else(|| Error::InvalidShape("Mimi decoder final conv produced no output".into()))
1057    }
1058}
1059
1060#[derive(Debug, Clone)]
1061struct DecoderLayer<T: Tensor> {
1062    upsample: StreamableConvTranspose1d<T>,
1063    residuals: Vec<SeaNetResnetBlock<T>>,
1064}
1065
1066impl<T: Tensor> DecoderLayer<T> {
1067    fn unloaded(
1068        in_channels: i32,
1069        out_channels: i32,
1070        ratio: i32,
1071        context: &T::Context,
1072    ) -> Result<Self, Error> {
1073        Ok(Self {
1074            upsample: StreamableConvTranspose1d::unloaded(
1075                in_channels,
1076                out_channels,
1077                ratio * 2,
1078                ratio,
1079                1,
1080                true,
1081                context,
1082            )?,
1083            residuals: vec![SeaNetResnetBlock::unloaded(out_channels, context)?],
1084        })
1085    }
1086
1087    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1088        let mut x = self.upsample.forward(x, context)?;
1089        for residual in &mut self.residuals {
1090            x = residual.forward(&x, context)?;
1091        }
1092        Ok(x)
1093    }
1094
1095    fn reset_state(&mut self) {
1096        self.upsample.reset_state();
1097        for residual in &mut self.residuals {
1098            residual.reset_state();
1099        }
1100    }
1101
1102    fn step(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1103        let mut x = self.upsample.step(x, context)?;
1104        for residual in &mut self.residuals {
1105            x = residual.step(&x, context)?;
1106        }
1107        Ok(x)
1108    }
1109}
1110
1111#[derive(Debug, Clone)]
1112struct SeaNetResnetBlock<T: Tensor> {
1113    block: Vec<StreamableConv1d<T>>,
1114}
1115
1116impl<T: Tensor> SeaNetResnetBlock<T> {
1117    fn unloaded(channels: i32, context: &T::Context) -> Result<Self, Error> {
1118        Ok(Self {
1119            block: vec![
1120                StreamableConv1d::unloaded(channels, channels / 2, 3, 1, context)?,
1121                StreamableConv1d::unloaded(channels / 2, channels, 1, 1, context)?,
1122            ],
1123        })
1124    }
1125
1126    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1127        let mut y = x.clone();
1128        for conv in &mut self.block {
1129            y = conv.forward(&T::elu(&y, 1.0, context)?, context)?;
1130        }
1131        Ok(y.add(x, context)?)
1132    }
1133
1134    fn reset_state(&mut self) {
1135        for conv in &mut self.block {
1136            conv.reset_state();
1137        }
1138    }
1139
1140    fn step(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1141        let mut y = x.clone();
1142        for conv in &mut self.block {
1143            y = conv
1144                .step(&T::elu(&y, 1.0, context)?, context)?
1145                .ok_or_else(|| {
1146                    Error::InvalidShape("Mimi residual conv produced no output".into())
1147                })?;
1148        }
1149        Ok(y.add(x, context)?)
1150    }
1151}
1152
1153#[derive(Debug, Clone)]
1154struct StreamableConv1d<T: Tensor> {
1155    weight: Parameter<T>,
1156    bias: Option<Parameter<T>>,
1157    stride: i32,
1158    dilation: i32,
1159    groups: i32,
1160    pad_mode: PadMode,
1161    state_prev_xs: Option<T>,
1162    left_pad_applied: bool,
1163}
1164
1165impl<T: Tensor> StreamableConv1d<T> {
1166    fn unloaded(
1167        in_channels: i32,
1168        out_channels: i32,
1169        kernel_size: i32,
1170        stride: i32,
1171        context: &T::Context,
1172    ) -> Result<Self, Error> {
1173        Self::unloaded_with_pad_mode(
1174            in_channels,
1175            out_channels,
1176            kernel_size,
1177            stride,
1178            true,
1179            PadMode::Constant,
1180            context,
1181        )
1182    }
1183
1184    fn unloaded_with_pad_mode(
1185        in_channels: i32,
1186        out_channels: i32,
1187        kernel_size: i32,
1188        stride: i32,
1189        bias: bool,
1190        pad_mode: PadMode,
1191        context: &T::Context,
1192    ) -> Result<Self, Error> {
1193        Ok(Self {
1194            weight: unloaded_parameter(&[out_channels, kernel_size, in_channels], context)?,
1195            bias: bias
1196                .then(|| unloaded_parameter(&[out_channels], context))
1197                .transpose()?,
1198            stride,
1199            dilation: 1,
1200            groups: 1,
1201            pad_mode,
1202            state_prev_xs: None,
1203            left_pad_applied: false,
1204        })
1205    }
1206
1207    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1208        let kernel_size = self.weight.as_ref().dim(1);
1209        let effective_kernel = (kernel_size - 1) * self.dilation + 1;
1210        let padding_total = effective_kernel - self.stride;
1211        let extra_padding =
1212            extra_padding_for_conv1d(x.dim(2), effective_kernel, self.stride, padding_total);
1213        let x = pad_bct(x, padding_total, extra_padding, self.pad_mode, context)?;
1214        let x = x.swap_axes(1, 2, context)?;
1215        let mut y = T::conv1d(
1216            &x,
1217            self.weight.as_ref(),
1218            self.stride,
1219            0,
1220            self.dilation,
1221            self.groups,
1222            context,
1223        )?;
1224        if let Some(bias) = &self.bias {
1225            y = y.add(bias.as_ref(), context)?;
1226        }
1227        Ok(y.swap_axes(1, 2, context)?)
1228    }
1229
1230    fn reset_state(&mut self) {
1231        self.state_prev_xs = None;
1232        self.left_pad_applied = false;
1233    }
1234
1235    fn step(&mut self, x: &T, context: &T::Context) -> Result<Option<T>, Error> {
1236        let kernel_size = self.weight.as_ref().dim(1);
1237        let effective_kernel = (kernel_size - 1) * self.dilation + 1;
1238        let padding_total = effective_kernel - self.stride;
1239        let x = if self.left_pad_applied {
1240            x.clone()
1241        } else {
1242            self.left_pad_applied = true;
1243            pad_bct(x, padding_total, 0, self.pad_mode, context)?
1244        };
1245        let x = match self.state_prev_xs.take() {
1246            Some(prev) => T::concatenate(&[prev, x], 2, context)?,
1247            None => x,
1248        };
1249        let seq_len = x.dim(2);
1250        let num_frames = (seq_len + self.stride).saturating_sub(effective_kernel) / self.stride;
1251        if num_frames <= 0 {
1252            self.state_prev_xs = Some(x);
1253            return Ok(None);
1254        }
1255        let offset = num_frames * self.stride;
1256        self.state_prev_xs = Some(x.index(
1257            &[Index::Full, Index::Full, Index::Range(offset, seq_len)],
1258            context,
1259        )?);
1260        let in_len = (num_frames - 1) * self.stride + effective_kernel;
1261        let x = x.index(
1262            &[Index::Full, Index::Full, Index::Range(0, in_len)],
1263            context,
1264        )?;
1265        self.forward_unpadded(&x, context).map(Some)
1266    }
1267
1268    fn forward_unpadded(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1269        let x = x.swap_axes(1, 2, context)?;
1270        let mut y = T::conv1d(
1271            &x,
1272            self.weight.as_ref(),
1273            self.stride,
1274            0,
1275            self.dilation,
1276            self.groups,
1277            context,
1278        )?;
1279        if let Some(bias) = &self.bias {
1280            y = y.add(bias.as_ref(), context)?;
1281        }
1282        Ok(y.swap_axes(1, 2, context)?)
1283    }
1284}
1285
1286#[derive(Debug, Clone)]
1287struct StreamableConvTranspose1d<T: Tensor> {
1288    weight: Parameter<T>,
1289    bias: Option<Parameter<T>>,
1290    kernel_size: i32,
1291    stride: i32,
1292    groups: i32,
1293    state_prev_ys: Option<T>,
1294}
1295
1296impl<T: Tensor> StreamableConvTranspose1d<T> {
1297    fn unloaded(
1298        in_channels: i32,
1299        out_channels: i32,
1300        kernel_size: i32,
1301        stride: i32,
1302        groups: i32,
1303        bias: bool,
1304        context: &T::Context,
1305    ) -> Result<Self, Error> {
1306        Ok(Self {
1307            weight: unloaded_parameter(
1308                &[out_channels, kernel_size, in_channels / groups],
1309                context,
1310            )?,
1311            bias: bias
1312                .then(|| unloaded_parameter(&[out_channels], context))
1313                .transpose()?,
1314            kernel_size,
1315            stride,
1316            groups,
1317            state_prev_ys: None,
1318        })
1319    }
1320
1321    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1322        let y = self.forward_untrimmed(x, context)?;
1323        let padding_total = self.kernel_size.saturating_sub(self.stride);
1324        unpad_bct(&y, 0, padding_total, context)
1325    }
1326
1327    fn reset_state(&mut self) {
1328        self.state_prev_ys = None;
1329    }
1330
1331    fn step(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1332        let y = self.forward_untrimmed(x, context)?;
1333        let out_len = y.dim(2);
1334        let y = match self.state_prev_ys.take() {
1335            None => y,
1336            Some(prev) => {
1337                let prev_len = prev.dim(2);
1338                let prev = match &self.bias {
1339                    None => prev,
1340                    Some(bias) => prev.subtract(
1341                        &bias
1342                            .as_ref()
1343                            .reshape(&[1, bias.as_ref().dim(0), 1], context)?,
1344                        context,
1345                    )?,
1346                };
1347                let y1 = y
1348                    .index(
1349                        &[Index::Full, Index::Full, Index::Range(0, prev_len)],
1350                        context,
1351                    )?
1352                    .add(&prev, context)?;
1353                let y2 = y.index(
1354                    &[Index::Full, Index::Full, Index::Range(prev_len, out_len)],
1355                    context,
1356                )?;
1357                T::concatenate(&[y1, y2], 2, context)?
1358            }
1359        };
1360        let invalid_steps = self.kernel_size - self.stride;
1361        let split = out_len - invalid_steps;
1362        let out = y.index(&[Index::Full, Index::Full, Index::Range(0, split)], context)?;
1363        self.state_prev_ys = Some(y.index(
1364            &[Index::Full, Index::Full, Index::Range(split, out_len)],
1365            context,
1366        )?);
1367        Ok(out)
1368    }
1369
1370    fn forward_untrimmed(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1371        let x = x.swap_axes(1, 2, context)?;
1372        let mut y = T::conv_transpose1d(
1373            &x,
1374            self.weight.as_ref(),
1375            self.stride,
1376            0,
1377            1,
1378            0,
1379            self.groups,
1380            context,
1381        )?;
1382        if let Some(bias) = &self.bias {
1383            y = y.add(bias.as_ref(), context)?;
1384        }
1385        Ok(y.swap_axes(1, 2, context)?)
1386    }
1387}
1388
1389fn extra_padding_for_conv1d(len: i32, kernel_size: i32, stride: i32, padding_total: i32) -> i32 {
1390    let n_frames = (len + padding_total - kernel_size) as f64 / stride as f64 + 1.0;
1391    let ideal_len = ((n_frames.ceil() as i32 - 1) * stride + kernel_size) - padding_total;
1392    ideal_len.saturating_sub(len)
1393}
1394
1395fn pad_bct<T: Tensor>(
1396    x: &T,
1397    left: i32,
1398    right: i32,
1399    mode: PadMode,
1400    context: &T::Context,
1401) -> Result<T, Error> {
1402    Ok(T::pad(x, &[(0, 0), (0, 0), (left, right)], mode, context)?)
1403}
1404
1405fn unpad_bct<T: Tensor>(x: &T, left: i32, right: i32, context: &T::Context) -> Result<T, Error> {
1406    let len = x.dim(2);
1407    if len < left + right {
1408        return Err(Error::InvalidShape(format!(
1409            "cannot unpad Mimi tensor of length {len} by {left}+{right}"
1410        )));
1411    }
1412    Ok(x.index(
1413        &[Index::Full, Index::Full, Index::Range(left, len - right)],
1414        context,
1415    )?)
1416}
1417
1418/// Split residual vector quantizer used by Mimi.
1419#[derive(Debug, Clone)]
1420pub struct SplitResidualVectorQuantizer<T: Tensor> {
1421    /// First semantic codebook branch.
1422    pub rvq_first: ResidualVectorQuantizer<T>,
1423    /// Remaining acoustic codebook branch.
1424    pub rvq_rest: ResidualVectorQuantizer<T>,
1425    n_q: i32,
1426}
1427
1428impl<T: Tensor> SplitResidualVectorQuantizer<T> {
1429    fn unloaded(config: &Config, context: &T::Context) -> Result<Self, Error> {
1430        Ok(Self {
1431            rvq_first: ResidualVectorQuantizer::unloaded(
1432                config.latent_dim,
1433                config.quantizer_dim,
1434                1,
1435                config.bins,
1436                context,
1437            )?,
1438            rvq_rest: ResidualVectorQuantizer::unloaded(
1439                config.latent_dim,
1440                config.quantizer_dim,
1441                config.num_codebooks - 1,
1442                config.bins,
1443                context,
1444            )?,
1445            n_q: config.num_codebooks,
1446        })
1447    }
1448
1449    /// Encodes latent frames shaped `[batch, 512, frames]`.
1450    pub fn encode(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
1451        validate_latent(latent)?;
1452        let first = self.rvq_first.encode(latent, context)?;
1453        if self.n_q == 1 {
1454            Ok(first)
1455        } else {
1456            let rest = self.rvq_rest.encode(latent, context)?;
1457            Ok(T::concatenate(&[first, rest], 1, context)?)
1458        }
1459    }
1460
1461    /// Decodes tokens shaped `[batch, codebooks, frames]`.
1462    pub fn decode(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
1463        validate_codes(codes, self.n_q)?;
1464        let first_codes = codes.index(&[Index::Full, Index::Range(0, 1), Index::Full], context)?;
1465        let mut quantized = self.rvq_first.decode(&first_codes, context)?;
1466        if codes.dim(1) > 1 {
1467            let rest_codes = codes.index(
1468                &[Index::Full, Index::Range(1, codes.dim(1)), Index::Full],
1469                context,
1470            )?;
1471            quantized = quantized.add(&self.rvq_rest.decode(&rest_codes, context)?, context)?;
1472        }
1473        Ok(quantized)
1474    }
1475}
1476
1477/// Residual vector quantizer branch.
1478#[derive(Debug, Clone)]
1479pub struct ResidualVectorQuantizer<T: Tensor> {
1480    /// Input projection from Mimi latent dimension into codebook dimension.
1481    pub input_proj: Conv1x1NoBias<T>,
1482    /// Output projection from codebook dimension back into Mimi latent dimension.
1483    pub output_proj: Conv1x1NoBias<T>,
1484    /// Residual codebook layers.
1485    pub vq: ResidualVectorQuantization<T>,
1486}
1487
1488impl<T: Tensor> ResidualVectorQuantizer<T> {
1489    fn unloaded(
1490        latent_dim: i32,
1491        codebook_dim: i32,
1492        layers: i32,
1493        bins: i32,
1494        context: &T::Context,
1495    ) -> Result<Self, Error> {
1496        Ok(Self {
1497            input_proj: Conv1x1NoBias::unloaded(latent_dim, codebook_dim, context)?,
1498            output_proj: Conv1x1NoBias::unloaded(codebook_dim, latent_dim, context)?,
1499            vq: ResidualVectorQuantization::unloaded(layers, codebook_dim, bins, context)?,
1500        })
1501    }
1502
1503    fn encode(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
1504        self.vq
1505            .encode(&self.input_proj.forward(latent, context)?, context)
1506    }
1507
1508    fn decode(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
1509        self.output_proj
1510            .forward(&self.vq.decode(codes, context)?, context)
1511    }
1512}
1513
1514/// Residual vector quantization layers.
1515#[derive(Debug, Clone)]
1516pub struct ResidualVectorQuantization<T: Tensor> {
1517    /// Ordered residual quantization layers.
1518    pub layers: Vec<VectorQuantization<T>>,
1519}
1520
1521impl<T: Tensor> ResidualVectorQuantization<T> {
1522    fn unloaded(layers: i32, dim: i32, bins: i32, context: &T::Context) -> Result<Self, Error> {
1523        Ok(Self {
1524            layers: (0..layers)
1525                .map(|_| VectorQuantization::unloaded(dim, bins, context))
1526                .collect::<Result<Vec<_>, _>>()?,
1527        })
1528    }
1529
1530    fn encode(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
1531        if self.layers.is_empty() {
1532            return Err(Error::InvalidShape("Mimi RVQ has no layers".into()));
1533        }
1534        let mut residual = latent.clone();
1535        let mut codes = Vec::with_capacity(self.layers.len());
1536        for layer in &mut self.layers {
1537            let indices = layer.encode(&residual, context)?;
1538            let quantized = layer.decode_one(&indices, context)?;
1539            residual = residual.subtract(&quantized, context)?;
1540            codes.push(indices);
1541        }
1542        Ok(T::stack(&codes, 1, context)?)
1543    }
1544
1545    fn decode(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
1546        if codes.dim(1) != self.layers.len() as i32 {
1547            return Err(Error::InvalidShape(format!(
1548                "Mimi RVQ expected {} codebooks, got {:?}",
1549                self.layers.len(),
1550                codes.shape()
1551            )));
1552        }
1553        let mut out: Option<T> = None;
1554        for (index, layer) in self.layers.iter_mut().enumerate() {
1555            let code = codes.index(
1556                &[Index::Full, Index::At(index as i32), Index::Full],
1557                context,
1558            )?;
1559            let quantized = layer.decode_one(&code, context)?;
1560            out = Some(match out {
1561                None => quantized,
1562                Some(prev) => prev.add(&quantized, context)?,
1563            });
1564        }
1565        out.ok_or_else(|| Error::InvalidShape("Mimi RVQ has no layers".into()))
1566    }
1567}
1568
1569/// Single vector-quantization layer.
1570#[derive(Debug, Clone)]
1571pub struct VectorQuantization<T: Tensor> {
1572    /// Euclidean codebook.
1573    pub _codebook: EuclideanCodebook<T>,
1574}
1575
1576impl<T: Tensor> VectorQuantization<T> {
1577    fn unloaded(dim: i32, bins: i32, context: &T::Context) -> Result<Self, Error> {
1578        Ok(Self {
1579            _codebook: EuclideanCodebook::unloaded(dim, bins, context)?,
1580        })
1581    }
1582
1583    fn encode(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
1584        let latent = latent.swap_axes(1, 2, context)?;
1585        self._codebook.encode(&latent, context)
1586    }
1587
1588    fn decode_one(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
1589        self._codebook
1590            .decode(codes, context)?
1591            .swap_axes(1, 2, context)
1592            .map_err(Into::into)
1593    }
1594}
1595
1596/// Euclidean codebook backed by EMA cluster statistics.
1597#[derive(Debug, Clone)]
1598pub struct EuclideanCodebook<T: Tensor> {
1599    /// Checkpoint initialization flag.
1600    pub _initialized: Parameter<T>,
1601    /// EMA cluster usage.
1602    pub cluster_usage: Parameter<T>,
1603    /// EMA embedding sum.
1604    pub embedding_sum: Parameter<T>,
1605}
1606
1607impl<T: Tensor> EuclideanCodebook<T> {
1608    fn unloaded(dim: i32, bins: i32, context: &T::Context) -> Result<Self, Error> {
1609        Ok(Self {
1610            _initialized: unloaded_parameter(&[1], context)?,
1611            cluster_usage: unloaded_parameter(&[bins], context)?,
1612            embedding_sum: unloaded_parameter(&[bins, dim], context)?,
1613        })
1614    }
1615
1616    fn embedding(&self, context: &T::Context) -> Result<T, Error> {
1617        let usage = self
1618            .cluster_usage
1619            .as_ref()
1620            .maximum_scalar(EPSILON, context)?
1621            .expand_dims(1, context)?;
1622        Ok(self.embedding_sum.as_ref().divide(&usage, context)?)
1623    }
1624
1625    fn encode(&self, latent_btd: &T, context: &T::Context) -> Result<T, Error> {
1626        if latent_btd.shape().len() != 3 {
1627            return Err(Error::InvalidShape(format!(
1628                "Mimi codebook encode expects [batch, frames, dim], got {:?}",
1629                latent_btd.shape()
1630            )));
1631        }
1632        let batch = latent_btd.dim(0);
1633        let frames = latent_btd.dim(1);
1634        let dim = latent_btd.dim(2);
1635        let flat = latent_btd.reshape(&[batch * frames, dim], context)?;
1636        let embedding = self.embedding(context)?;
1637        let x2 = T::sum_axis(&flat.square(context)?, -1, true, context)?;
1638        let e2 = T::sum_axis(&embedding.square(context)?, -1, false, context)?
1639            .expand_dims(0, context)?;
1640        let dot = T::matmul(&flat, &embedding.transpose(context)?, context)?;
1641        let dists = x2
1642            .add(&e2, context)?
1643            .subtract(&dot.multiply_scalar(2.0, context)?, context)?;
1644        Ok(T::argmin_axis(&dists, -1, false, context)?.reshape(&[batch, frames], context)?)
1645    }
1646
1647    fn decode(&self, codes: &T, context: &T::Context) -> Result<T, Error> {
1648        if codes.shape().len() != 2 {
1649            return Err(Error::InvalidShape(format!(
1650                "Mimi codebook decode expects [batch, frames], got {:?}",
1651                codes.shape()
1652            )));
1653        }
1654        let batch = codes.dim(0);
1655        let frames = codes.dim(1);
1656        let embedding = self.embedding(context)?;
1657        let flat = codes.reshape(&[batch * frames], context)?;
1658        Ok(embedding
1659            .take_axis(&flat, 0, context)?
1660            .reshape(&[batch, frames, embedding.dim(1)], context)?)
1661    }
1662}
1663
1664/// Bias-free 1x1 convolution over `[batch, channels, frames]` tensors.
1665#[derive(Debug, Clone)]
1666pub struct Conv1x1NoBias<T: Tensor> {
1667    /// Weight shaped `[out_channels, in_channels, 1]`.
1668    pub weight: Parameter<T>,
1669}
1670
1671impl<T: Tensor> Conv1x1NoBias<T> {
1672    fn unloaded(in_channels: i32, out_channels: i32, context: &T::Context) -> Result<Self, Error> {
1673        Ok(Self {
1674            weight: unloaded_parameter(&[out_channels, in_channels, 1], context)?,
1675        })
1676    }
1677
1678    fn forward(&self, latent: &T, context: &T::Context) -> Result<T, Error> {
1679        if latent.shape().len() != 3 {
1680            return Err(Error::InvalidShape(format!(
1681                "Mimi 1x1 projection expects [batch, channels, frames], got {:?}",
1682                latent.shape()
1683            )));
1684        }
1685        let x = latent.swap_axes(1, 2, context)?;
1686        let weight = self.weight.as_ref().squeeze_axes(&[-1], context)?;
1687        Ok(T::matmul(&x, &weight.transpose(context)?, context)?.swap_axes(1, 2, context)?)
1688    }
1689}
1690
1691trait MimiModuleParameters<T: Tensor> {
1692    fn visit_mimi_parameters<'a>(
1693        &'a self,
1694        prefix: &str,
1695        visitor: &mut dyn FnMut(ParameterMetadata, &'a T),
1696    );
1697    fn visit_mimi_parameters_mut<'a>(
1698        &'a mut self,
1699        prefix: &str,
1700        visitor: &mut dyn FnMut(ParameterMetadata, &'a mut T),
1701    );
1702    fn set_mimi_trainable(&mut self, trainable: bool);
1703}
1704
1705struct PrefixVisitor<'a, F: ?Sized> {
1706    prefix: &'a str,
1707    visitor: &'a mut F,
1708    exact: bool,
1709}
1710
1711impl<'a, 'value, T, F: ?Sized> ParameterVisitor<'value, T> for PrefixVisitor<'a, F>
1712where
1713    T: 'value,
1714    F: FnMut(ParameterMetadata, &'value T),
1715{
1716    fn visit(&mut self, mut metadata: ParameterMetadata, value: &'value T) {
1717        let id = if self.exact {
1718            self.prefix.to_owned()
1719        } else {
1720            parameter_name(self.prefix, metadata.id.as_str())
1721        };
1722        metadata.id = ParameterId::new(id).expect("Mimi parameter identities are non-empty");
1723        (self.visitor)(metadata, value);
1724    }
1725}
1726
1727struct PrefixVisitorMut<'a, F: ?Sized> {
1728    prefix: &'a str,
1729    visitor: &'a mut F,
1730    exact: bool,
1731}
1732
1733impl<'a, 'value, T, F: ?Sized> ParameterVisitorMut<'value, T> for PrefixVisitorMut<'a, F>
1734where
1735    T: 'value,
1736    F: FnMut(ParameterMetadata, &'value mut T),
1737{
1738    fn visit_mut(&mut self, mut metadata: ParameterMetadata, value: &'value mut T) {
1739        let id = if self.exact {
1740            self.prefix.to_owned()
1741        } else {
1742            parameter_name(self.prefix, metadata.id.as_str())
1743        };
1744        metadata.id = ParameterId::new(id).expect("Mimi parameter identities are non-empty");
1745        (self.visitor)(metadata, value);
1746    }
1747}
1748
1749impl<T: Tensor> MimiModuleParameters<T> for Parameter<T> {
1750    fn visit_mimi_parameters<'a>(
1751        &'a self,
1752        prefix: &str,
1753        visitor: &mut dyn FnMut(ParameterMetadata, &'a T),
1754    ) {
1755        self.visit_parameters(&mut PrefixVisitor {
1756            prefix,
1757            visitor,
1758            exact: true,
1759        });
1760    }
1761
1762    fn visit_mimi_parameters_mut<'a>(
1763        &'a mut self,
1764        prefix: &str,
1765        visitor: &mut dyn FnMut(ParameterMetadata, &'a mut T),
1766    ) {
1767        self.visit_parameters_mut(&mut PrefixVisitorMut {
1768            prefix,
1769            visitor,
1770            exact: true,
1771        });
1772    }
1773
1774    fn set_mimi_trainable(&mut self, trainable: bool) {
1775        self.set_trainable(trainable);
1776    }
1777}
1778
1779macro_rules! structured_leaf_parameters {
1780    ($type:ty) => {
1781        impl<T: Tensor> MimiModuleParameters<T> for $type {
1782            fn visit_mimi_parameters<'a>(
1783                &'a self,
1784                prefix: &str,
1785                visitor: &mut dyn FnMut(ParameterMetadata, &'a T),
1786            ) {
1787                self.visit_parameters(&mut PrefixVisitor {
1788                    prefix,
1789                    visitor,
1790                    exact: false,
1791                });
1792            }
1793
1794            fn visit_mimi_parameters_mut<'a>(
1795                &'a mut self,
1796                prefix: &str,
1797                visitor: &mut dyn FnMut(ParameterMetadata, &'a mut T),
1798            ) {
1799                self.visit_parameters_mut(&mut PrefixVisitorMut {
1800                    prefix,
1801                    visitor,
1802                    exact: false,
1803                });
1804            }
1805
1806            fn set_mimi_trainable(&mut self, trainable: bool) {
1807                self.set_trainable(trainable);
1808            }
1809        }
1810    };
1811}
1812
1813structured_leaf_parameters!(Linear<T>);
1814structured_leaf_parameters!(LayerNorm<T>);
1815
1816impl<T: Tensor, M: MimiModuleParameters<T>> MimiModuleParameters<T> for Vec<M> {
1817    fn visit_mimi_parameters<'a>(
1818        &'a self,
1819        prefix: &str,
1820        visitor: &mut dyn FnMut(ParameterMetadata, &'a T),
1821    ) {
1822        for (index, module) in self.iter().enumerate() {
1823            module.visit_mimi_parameters(&parameter_name(prefix, &index.to_string()), visitor);
1824        }
1825    }
1826
1827    fn visit_mimi_parameters_mut<'a>(
1828        &'a mut self,
1829        prefix: &str,
1830        visitor: &mut dyn FnMut(ParameterMetadata, &'a mut T),
1831    ) {
1832        for (index, module) in self.iter_mut().enumerate() {
1833            module.visit_mimi_parameters_mut(&parameter_name(prefix, &index.to_string()), visitor);
1834        }
1835    }
1836
1837    fn set_mimi_trainable(&mut self, trainable: bool) {
1838        for module in self {
1839            module.set_mimi_trainable(trainable);
1840        }
1841    }
1842}
1843
1844impl<T: Tensor, M: MimiModuleParameters<T>> MimiModuleParameters<T> for Option<M> {
1845    fn visit_mimi_parameters<'a>(
1846        &'a self,
1847        prefix: &str,
1848        visitor: &mut dyn FnMut(ParameterMetadata, &'a T),
1849    ) {
1850        if let Some(module) = self {
1851            module.visit_mimi_parameters(prefix, visitor);
1852        }
1853    }
1854
1855    fn visit_mimi_parameters_mut<'a>(
1856        &'a mut self,
1857        prefix: &str,
1858        visitor: &mut dyn FnMut(ParameterMetadata, &'a mut T),
1859    ) {
1860        if let Some(module) = self {
1861            module.visit_mimi_parameters_mut(prefix, visitor);
1862        }
1863    }
1864
1865    fn set_mimi_trainable(&mut self, trainable: bool) {
1866        if let Some(module) = self {
1867            module.set_mimi_trainable(trainable);
1868        }
1869    }
1870}
1871
1872macro_rules! module_parameters {
1873    ($module:ident { $($field:ident),+ $(,)? }) => {
1874        impl<T: Tensor> MimiModuleParameters<T> for $module<T> {
1875            fn visit_mimi_parameters<'a>(
1876                &'a self,
1877                prefix: &str,
1878                visitor: &mut dyn FnMut(ParameterMetadata, &'a T),
1879            ) {
1880                $(
1881                    self.$field.visit_mimi_parameters(
1882                        &parameter_name(prefix, stringify!($field)),
1883                        visitor,
1884                    );
1885                )+
1886            }
1887
1888            fn visit_mimi_parameters_mut<'a>(
1889                &'a mut self,
1890                prefix: &str,
1891                visitor: &mut dyn FnMut(ParameterMetadata, &'a mut T),
1892            ) {
1893                $(
1894                    self.$field.visit_mimi_parameters_mut(
1895                        &parameter_name(prefix, stringify!($field)),
1896                        visitor,
1897                    );
1898                )+
1899            }
1900
1901            fn set_mimi_trainable(&mut self, trainable: bool) {
1902                $(self.$field.set_mimi_trainable(trainable);)+
1903            }
1904        }
1905    };
1906}
1907
1908module_parameters!(Mimi {
1909    quantizer,
1910    encoder,
1911    encoder_transformer,
1912    downsample,
1913    upsample,
1914    decoder_transformer,
1915    decoder,
1916});
1917module_parameters!(SeaNetEncoder {
1918    init_conv1d,
1919    layers,
1920    final_conv1d,
1921});
1922module_parameters!(EncoderLayer {
1923    residuals,
1924    downsample,
1925});
1926module_parameters!(MimiTransformer { layers });
1927module_parameters!(MimiTransformerLayer {
1928    norm1,
1929    norm2,
1930    self_attn,
1931    mlp,
1932    layer_scale_1,
1933    layer_scale_2,
1934});
1935module_parameters!(LayerScale { scale });
1936module_parameters!(MimiMlp { linear1, linear2 });
1937module_parameters!(MimiSelfAttention { in_proj, out_proj });
1938module_parameters!(SeaNetDecoder {
1939    init_conv1d,
1940    layers,
1941    final_conv1d,
1942});
1943module_parameters!(DecoderLayer {
1944    upsample,
1945    residuals,
1946});
1947module_parameters!(SeaNetResnetBlock { block });
1948module_parameters!(StreamableConv1d { weight, bias });
1949module_parameters!(StreamableConvTranspose1d { weight, bias });
1950module_parameters!(SplitResidualVectorQuantizer {
1951    rvq_first,
1952    rvq_rest,
1953});
1954module_parameters!(ResidualVectorQuantizer {
1955    input_proj,
1956    output_proj,
1957    vq,
1958});
1959module_parameters!(ResidualVectorQuantization { layers });
1960module_parameters!(VectorQuantization { _codebook });
1961module_parameters!(EuclideanCodebook {
1962    _initialized,
1963    cluster_usage,
1964    embedding_sum,
1965});
1966module_parameters!(Conv1x1NoBias { weight });
1967
1968impl<T: Tensor> Parameterized<T> for Mimi<T> {
1969    fn visit_parameters<'a, V>(&'a self, visitor: &mut V)
1970    where
1971        V: ParameterVisitor<'a, T>,
1972    {
1973        self.visit_mimi_parameters("", &mut |metadata, value| {
1974            visitor.visit(metadata, value);
1975        });
1976    }
1977
1978    fn visit_parameters_mut<'a, V>(&'a mut self, visitor: &mut V)
1979    where
1980        V: ParameterVisitorMut<'a, T>,
1981    {
1982        self.visit_mimi_parameters_mut("", &mut |metadata, value| {
1983            visitor.visit_mut(metadata, value);
1984        });
1985    }
1986
1987    fn set_trainable(&mut self, trainable: bool) {
1988        self.set_mimi_trainable(trainable);
1989    }
1990}
1991
1992fn validate_latent<T: Tensor>(latent: &T) -> Result<(), Error> {
1993    if latent.shape().len() != 3 || latent.dim(1) != 512 {
1994        return Err(Error::InvalidShape(format!(
1995            "Mimi latent frames must have shape [batch, 512, frames], got {:?}",
1996            latent.shape()
1997        )));
1998    }
1999    Ok(())
2000}
2001
2002fn validate_pcm<T: Tensor>(pcm: &T) -> Result<(), Error> {
2003    if pcm.shape().len() != 3 || pcm.dim(1) != 1 {
2004        return Err(Error::InvalidShape(format!(
2005            "Mimi PCM must have shape [batch, 1, samples], got {:?}",
2006            pcm.shape()
2007        )));
2008    }
2009    Ok(())
2010}
2011
2012fn validate_codes<T: Tensor>(codes: &T, max_codebooks: i32) -> Result<(), Error> {
2013    if codes.shape().len() != 3 || codes.dim(1) <= 0 || codes.dim(1) > max_codebooks {
2014        return Err(Error::InvalidShape(format!(
2015            "Mimi codes must have shape [batch, 1..={max_codebooks}, frames], got {:?}",
2016            codes.shape()
2017        )));
2018    }
2019    Ok(())
2020}
2021
2022#[cfg(test)]
2023mod tests {
2024    use super::{
2025        checkpoint_tensor_plan, CheckpointTensorLayout, Config, Mimi, MimiModuleParameters,
2026    };
2027    use eredu_nn::{AttentionMask, Error as ComputeError, Index, PadMode, Tensor};
2028
2029    #[derive(Debug, Clone)]
2030    struct ShapeTensor(Vec<i32>);
2031
2032    impl ShapeTensor {
2033        fn unavailable() -> Result<Self, ComputeError> {
2034            unreachable!("shape-only test backend cannot execute tensor operations")
2035        }
2036    }
2037
2038    impl Tensor for ShapeTensor {
2039        type Context = ();
2040
2041        fn shape(&self) -> &[i32] {
2042            &self.0
2043        }
2044
2045        fn unloaded_f32(shape: &[i32], _: &Self::Context) -> Result<Self, ComputeError> {
2046            Ok(Self(shape.to_vec()))
2047        }
2048
2049        fn from_f32_slice(_: &[f32], _: &[i32], _: &Self::Context) -> Result<Self, ComputeError> {
2050            Self::unavailable()
2051        }
2052
2053        fn add(&self, _: &Self, _: &Self::Context) -> Result<Self, ComputeError> {
2054            Self::unavailable()
2055        }
2056
2057        fn subtract(&self, _: &Self, _: &Self::Context) -> Result<Self, ComputeError> {
2058            Self::unavailable()
2059        }
2060
2061        fn multiply(&self, _: &Self, _: &Self::Context) -> Result<Self, ComputeError> {
2062            Self::unavailable()
2063        }
2064
2065        fn multiply_scalar(&self, _: f32, _: &Self::Context) -> Result<Self, ComputeError> {
2066            Self::unavailable()
2067        }
2068
2069        fn divide(&self, _: &Self, _: &Self::Context) -> Result<Self, ComputeError> {
2070            Self::unavailable()
2071        }
2072
2073        fn square(&self, _: &Self::Context) -> Result<Self, ComputeError> {
2074            Self::unavailable()
2075        }
2076
2077        fn maximum_scalar(&self, _: f32, _: &Self::Context) -> Result<Self, ComputeError> {
2078            Self::unavailable()
2079        }
2080
2081        fn reshape(&self, _: &[i32], _: &Self::Context) -> Result<Self, ComputeError> {
2082            Self::unavailable()
2083        }
2084
2085        fn transpose_axes(&self, _: &[i32], _: &Self::Context) -> Result<Self, ComputeError> {
2086            Self::unavailable()
2087        }
2088
2089        fn swap_axes(&self, _: i32, _: i32, _: &Self::Context) -> Result<Self, ComputeError> {
2090            Self::unavailable()
2091        }
2092
2093        fn transpose(&self, _: &Self::Context) -> Result<Self, ComputeError> {
2094            Self::unavailable()
2095        }
2096
2097        fn expand_dims(&self, _: i32, _: &Self::Context) -> Result<Self, ComputeError> {
2098            Self::unavailable()
2099        }
2100
2101        fn squeeze_axes(&self, _: &[i32], _: &Self::Context) -> Result<Self, ComputeError> {
2102            Self::unavailable()
2103        }
2104
2105        fn index(&self, _: &[Index], _: &Self::Context) -> Result<Self, ComputeError> {
2106            Self::unavailable()
2107        }
2108
2109        fn take_axis(&self, _: &Self, _: i32, _: &Self::Context) -> Result<Self, ComputeError> {
2110            Self::unavailable()
2111        }
2112
2113        fn concatenate(_: &[Self], _: i32, _: &Self::Context) -> Result<Self, ComputeError> {
2114            Self::unavailable()
2115        }
2116
2117        fn stack(_: &[Self], _: i32, _: &Self::Context) -> Result<Self, ComputeError> {
2118            Self::unavailable()
2119        }
2120
2121        fn matmul(_: &Self, _: &Self, _: &Self::Context) -> Result<Self, ComputeError> {
2122            Self::unavailable()
2123        }
2124
2125        fn sum_axis(_: &Self, _: i32, _: bool, _: &Self::Context) -> Result<Self, ComputeError> {
2126            Self::unavailable()
2127        }
2128
2129        fn argmin_axis(_: &Self, _: i32, _: bool, _: &Self::Context) -> Result<Self, ComputeError> {
2130            Self::unavailable()
2131        }
2132
2133        fn pad(
2134            _: &Self,
2135            _: &[(i32, i32)],
2136            _: PadMode,
2137            _: &Self::Context,
2138        ) -> Result<Self, ComputeError> {
2139            Self::unavailable()
2140        }
2141
2142        fn conv1d(
2143            _: &Self,
2144            _: &Self,
2145            _: i32,
2146            _: i32,
2147            _: i32,
2148            _: i32,
2149            _: &Self::Context,
2150        ) -> Result<Self, ComputeError> {
2151            Self::unavailable()
2152        }
2153
2154        fn conv_transpose1d(
2155            _: &Self,
2156            _: &Self,
2157            _: i32,
2158            _: i32,
2159            _: i32,
2160            _: i32,
2161            _: i32,
2162            _: &Self::Context,
2163        ) -> Result<Self, ComputeError> {
2164            Self::unavailable()
2165        }
2166
2167        fn linear(
2168            _: &Self,
2169            _: &Self,
2170            _: Option<&Self>,
2171            _: &Self::Context,
2172        ) -> Result<Self, ComputeError> {
2173            Self::unavailable()
2174        }
2175
2176        fn layer_norm(
2177            _: &Self,
2178            _: Option<&Self>,
2179            _: Option<&Self>,
2180            _: f32,
2181            _: &Self::Context,
2182        ) -> Result<Self, ComputeError> {
2183            Self::unavailable()
2184        }
2185
2186        fn gelu(_: &Self, _: &Self::Context) -> Result<Self, ComputeError> {
2187            Self::unavailable()
2188        }
2189
2190        fn elu(_: &Self, _: f32, _: &Self::Context) -> Result<Self, ComputeError> {
2191            Self::unavailable()
2192        }
2193
2194        fn rope(
2195            _: &Self,
2196            _: i32,
2197            _: bool,
2198            _: f32,
2199            _: f32,
2200            _: i32,
2201            _: &Self::Context,
2202        ) -> Result<Self, ComputeError> {
2203            Self::unavailable()
2204        }
2205
2206        fn scaled_dot_product_attention(
2207            _: &Self,
2208            _: &Self,
2209            _: &Self,
2210            _: f32,
2211            _: AttentionMask<'_, Self>,
2212            _: &Self::Context,
2213        ) -> Result<Self, ComputeError> {
2214            Self::unavailable()
2215        }
2216    }
2217
2218    #[test]
2219    fn checkpoint_quantizer_keys_keep_the_model_root() {
2220        let key = "quantizer.rvq_first.vq.layers.0._codebook.embedding_sum";
2221        let plan = checkpoint_tensor_plan(key).unwrap();
2222        assert_eq!(plan.parameter, key);
2223        assert_eq!(plan.layout, CheckpointTensorLayout::Identity);
2224    }
2225
2226    #[test]
2227    fn checkpoint_plan_declares_canonical_convolution_layouts() {
2228        assert_eq!(
2229            checkpoint_tensor_plan("encoder.model.0.conv.conv.weight")
2230                .unwrap()
2231                .layout,
2232            CheckpointTensorLayout::Transpose3d([0, 2, 1])
2233        );
2234        assert_eq!(
2235            checkpoint_tensor_plan("decoder.model.2.convtr.convtr.weight")
2236                .unwrap()
2237                .layout,
2238            CheckpointTensorLayout::Transpose3d([1, 2, 0])
2239        );
2240        assert_eq!(
2241            checkpoint_tensor_plan("upsample.convtr.convtr.convtr.weight")
2242                .unwrap()
2243                .layout,
2244            CheckpointTensorLayout::Transpose3d([0, 2, 1])
2245        );
2246        assert!(checkpoint_tensor_plan("optimizer.state").is_none());
2247    }
2248
2249    #[test]
2250    fn parameter_names_are_unique_and_cover_checkpoint_mapping() {
2251        let model = Mimi::<ShapeTensor>::new(Config::v0_1(Some(8)), &()).unwrap();
2252        let mut names = Vec::new();
2253        model.visit_mimi_parameters("", &mut |metadata, _| {
2254            names.push(metadata.id.as_str().to_owned());
2255        });
2256
2257        let mut unique = names.clone();
2258        unique.sort();
2259        unique.dedup();
2260        assert_eq!(names.len(), unique.len(), "duplicate Mimi parameter name");
2261
2262        for checkpoint_name in [
2263            "quantizer.rvq_first.vq.layers.0._codebook.embedding_sum",
2264            "downsample.conv.conv.conv.weight",
2265            "encoder_transformer.transformer.layers.0.self_attn.in_proj_weight",
2266            "encoder.model.0.conv.conv.weight",
2267            "upsample.convtr.convtr.convtr.weight",
2268            "decoder_transformer.transformer.layers.0.linear1.weight",
2269            "decoder.model.0.conv.conv.weight",
2270        ] {
2271            let model_name = checkpoint_tensor_plan(checkpoint_name)
2272                .unwrap_or_else(|| panic!("checkpoint key was not mapped: {checkpoint_name}"))
2273                .parameter;
2274            assert!(
2275                unique.binary_search(&model_name).is_ok(),
2276                "mapped parameter is absent from Mimi: {checkpoint_name} -> {model_name}"
2277            );
2278        }
2279    }
2280
2281    #[test]
2282    fn v0_1_config_defaults_to_moshi_active_codebooks() {
2283        let cfg = Config::v0_1(None);
2284        assert_eq!(cfg.sample_rate, 24_000.0);
2285        assert_eq!(cfg.frame_rate, 12.5);
2286        assert_eq!(cfg.num_codebooks, 16);
2287        assert_eq!(cfg.total_codebooks, 32);
2288        assert_eq!(cfg.bins, 2_048);
2289    }
2290}