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_checkpoint::{
9    recipe::{DerivedWeightRecipe, RecipeCatalog, RecipeDtype, RecipeError, RecipeMetadata},
10    safetensors::{SafetensorsMetadataCatalog, SafetensorsShardError},
11    schema::{
12        CatalogPolicy, CheckpointPlanError, SafetensorsCheckpointPlan, SafetensorsTensorConstraint,
13        StoredDtypeConstraint,
14    },
15    store::{
16        PreparedCheckpointSource, PreparedTensorSource, ResolvedCheckpointSource,
17        SafetensorsWeightStore, SharedCheckpointSource, StoreError, TensorMetadata,
18        TensorSelection, TensorSourceProvenance,
19    },
20    validation::{resolve_safetensors_plan, CheckpointValidation, ResolvedCheckpointPlan},
21    SourceTensorEncoding, StoredDtype,
22};
23use eredu_nn::{
24    AttentionMask, Index, LayerNorm, Linear, LinearSpec, NeuralBackend, PadMode, Parameter,
25    ParameterId, ParameterMetadata, ParameterSpec, ParameterVisitor, ParameterVisitorMut,
26    Parameterized, Rope, Tensor,
27};
28use eredu_runtime::{
29    bind_materialized_unit, materialize_selected_bindings, select_bindings, ParameterBackend,
30    ParameterOrchestrationError, ResidencyDeclarationError, SelectedBindingPlan, WeightBinding,
31};
32use std::{
33    collections::{BTreeMap, BTreeSet},
34    path::Path,
35    sync::Arc,
36};
37
38use crate::{AudioTokenizer, AudioTokenizerConfig, Error};
39
40const EPSILON: f32 = 1e-5;
41
42fn parameter_name(prefix: &str, field: &str) -> String {
43    if prefix.is_empty() {
44        field.to_owned()
45    } else {
46        format!("{prefix}.{field}")
47    }
48}
49
50fn parameter_spec(id: &str) -> ParameterSpec {
51    ParameterSpec::trainable(id).expect("Mimi parameter identities are non-empty")
52}
53
54fn unloaded_parameter<T: Tensor>(
55    shape: &[i32],
56    context: &T::Context,
57) -> Result<Parameter<T>, eredu_nn::Error> {
58    Parameter::unloaded(parameter_spec("value"), shape, context)
59}
60
61fn unloaded_linear<T: Tensor>(
62    input: i32,
63    output: i32,
64    bias: bool,
65    context: &T::Context,
66) -> Result<Linear<T>, eredu_nn::Error> {
67    Linear::unloaded(
68        LinearSpec {
69            input,
70            output,
71            weight: parameter_spec("weight"),
72            bias: bias.then(|| parameter_spec("bias")),
73            format: eredu_nn::LinearFormatSpec::unscaled(eredu_checkpoint::LinearFormat::Dense)
74                .unwrap(),
75        },
76        context,
77    )
78}
79
80fn unloaded_layer_norm<T: Tensor>(
81    dimensions: i32,
82    epsilon: f32,
83    context: &T::Context,
84) -> Result<LayerNorm<T>, eredu_nn::Error> {
85    LayerNorm::unloaded(
86        dimensions,
87        epsilon,
88        Some(parameter_spec("weight")),
89        Some(parameter_spec("bias")),
90        context,
91    )
92}
93
94/// Mimi resampling strategy.
95#[derive(Debug, Clone, Copy, Eq, PartialEq)]
96pub enum ResampleMethod {
97    /// Learned convolutional resampling.
98    Conv,
99}
100
101/// Mimi codec configuration.
102#[derive(Debug, Clone)]
103pub struct Config {
104    /// Audio channels.
105    pub channels: i32,
106    /// PCM sample rate.
107    pub sample_rate: f64,
108    /// Codec frame rate.
109    pub frame_rate: f64,
110    /// Whether the original training path renormalized audio.
111    pub renormalize: bool,
112    /// Latent resampling method.
113    pub resample_method: ResampleMethod,
114    /// Active residual codebooks.
115    pub num_codebooks: i32,
116    /// Total codebooks available in the released checkpoint.
117    pub total_codebooks: i32,
118    /// Codebook cardinality.
119    pub bins: i32,
120    /// Codebook embedding dimension.
121    pub quantizer_dim: i32,
122    /// Model latent dimension.
123    pub latent_dim: i32,
124}
125
126impl Config {
127    /// Released Mimi v0.1 defaults, with a caller-selected active codebook count.
128    pub fn v0_1(num_codebooks: Option<i32>) -> Self {
129        Self {
130            channels: 1,
131            sample_rate: 24_000.0,
132            frame_rate: 12.5,
133            renormalize: true,
134            resample_method: ResampleMethod::Conv,
135            num_codebooks: num_codebooks.unwrap_or(16),
136            total_codebooks: 32,
137            bins: 2_048,
138            quantizer_dim: 256,
139            latent_dim: 512,
140        }
141    }
142
143    fn validate(&self) -> Result<(), Error> {
144        if self.channels <= 0
145            || !self.sample_rate.is_finite()
146            || self.sample_rate <= 0.0
147            || !self.frame_rate.is_finite()
148            || self.frame_rate <= 0.0
149            || self.num_codebooks <= 0
150            || self.num_codebooks > self.total_codebooks
151            || self.bins <= 0
152            || self.quantizer_dim <= 0
153            || self.latent_dim <= 0
154        {
155            return Err(Error::InvalidShape(format!(
156                "invalid Mimi config: channels={}, sample_rate={}, frame_rate={}, num_codebooks={}, total_codebooks={}, bins={}, quantizer_dim={}, latent_dim={}",
157                self.channels,
158                self.sample_rate,
159                self.frame_rate,
160                self.num_codebooks,
161                self.total_codebooks,
162                self.bins,
163                self.quantizer_dim,
164                self.latent_dim
165            )));
166        }
167        if self.channels != 1
168            || self.sample_rate != 24_000.0
169            || self.frame_rate != 12.5
170            || !self.renormalize
171            || self.resample_method != ResampleMethod::Conv
172            || self.total_codebooks != 32
173            || self.bins != 2_048
174            || self.quantizer_dim != 256
175            || self.latent_dim != 512
176        {
177            return Err(Error::InvalidShape(
178                "unsupported Mimi configuration; only the released v0.1 profile is admitted".into(),
179            ));
180        }
181        Ok(())
182    }
183}
184
185/// Mimi audio tokenizer.
186#[derive(Debug, Clone)]
187pub struct Mimi<T: Tensor> {
188    /// Split residual vector quantizer.
189    pub quantizer: SplitResidualVectorQuantizer<T>,
190    encoder: SeaNetEncoder<T>,
191    encoder_transformer: MimiTransformer<T>,
192    downsample: StreamableConv1d<T>,
193    upsample: StreamableConvTranspose1d<T>,
194    decoder_transformer: MimiTransformer<T>,
195    decoder: SeaNetDecoder<T>,
196    config: Config,
197}
198
199impl<T: Tensor> Mimi<T> {
200    /// Creates an unloaded Mimi tokenizer from config.
201    pub fn new(config: Config, context: &T::Context) -> Result<Self, Error> {
202        config.validate()?;
203        Ok(Self {
204            quantizer: SplitResidualVectorQuantizer::unloaded(&config, context)?,
205            encoder: SeaNetEncoder::unloaded(context)?,
206            encoder_transformer: MimiTransformer::unloaded(context)?,
207            downsample: StreamableConv1d::unloaded_with_pad_mode(
208                config.latent_dim,
209                config.latent_dim,
210                4,
211                2,
212                false,
213                PadMode::Edge,
214                context,
215            )?,
216            upsample: StreamableConvTranspose1d::unloaded(
217                config.latent_dim,
218                config.latent_dim,
219                4,
220                2,
221                config.latent_dim,
222                false,
223                context,
224            )?,
225            decoder_transformer: MimiTransformer::unloaded(context)?,
226            decoder: SeaNetDecoder::unloaded(context)?,
227            config,
228        })
229    }
230
231    /// Returns the Mimi configuration.
232    pub fn mimi_config(&self) -> &Config {
233        &self.config
234    }
235
236    /// Encodes latent frames shaped `[batch, 512, frames]` into Mimi tokens.
237    pub fn encode_latent(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
238        self.quantizer.encode(latent, context)
239    }
240
241    /// Encodes PCM shaped `[batch, 1, samples]` into Mimi tokens `[batch, codebooks, frames]`.
242    pub fn encode(&mut self, pcm: &T, context: &T::Context) -> Result<T, Error> {
243        let latent = self.encoder.forward(pcm, context)?;
244        let latent = self.encoder_transformer.forward(&latent, context)?;
245        let latent = self.downsample.forward(&latent, context)?;
246        self.quantizer.encode(&latent, context)
247    }
248
249    /// Resets state used by [`Mimi::encode_step`].
250    pub fn reset_encode_state(&mut self) {
251        self.encoder.reset_state();
252        self.encoder_transformer.reset_state();
253        self.downsample.reset_state();
254    }
255
256    /// Encodes one PCM frame into the next Mimi token frame.
257    ///
258    /// Accepts PCM shaped `[batch, 1, samples]`. Returns `None` until the
259    /// streaming encoder has enough samples to emit a complete codec frame.
260    pub fn encode_step(&mut self, pcm: &T, context: &T::Context) -> Result<Option<T>, Error> {
261        let latent = match self.encoder.step(pcm, context)? {
262            Some(latent) => latent,
263            None => return Ok(None),
264        };
265        let latent = self.encoder_transformer.step(&latent, context)?;
266        let latent = match self.downsample.step(&latent, context)? {
267            Some(latent) => latent,
268            None => return Ok(None),
269        };
270        Ok(Some(
271            self.quantizer
272                .encode(&latent, context)?
273                .squeeze_axes(&[2], context)?,
274        ))
275    }
276
277    /// Decodes Mimi tokens shaped `[batch, codebooks, frames]` into latent frames.
278    pub fn decode_latent(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
279        self.quantizer.decode(codes, context)
280    }
281
282    /// Decodes Mimi tokens shaped `[batch, codebooks, frames]` into PCM `[batch, 1, samples]`.
283    pub fn decode(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
284        let latent = self.quantizer.decode(codes, context)?;
285        let latent = self.upsample.forward(&latent, context)?;
286        let latent = self.decoder_transformer.forward(&latent, context)?;
287        self.decoder.forward(&latent, context)
288    }
289
290    /// Resets state used by [`Mimi::decode_step`].
291    pub fn reset_decode_state(&mut self) {
292        self.upsample.reset_state();
293        self.decoder_transformer.reset_state();
294        self.decoder.reset_state();
295    }
296
297    /// Decodes one Mimi token frame into the next PCM chunk.
298    ///
299    /// Accepts codes shaped `[batch, codebooks]` or `[batch, codebooks, 1]`.
300    pub fn decode_step(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
301        let codes = match codes.shape() {
302            [_, _] => codes.expand_dims(2, context)?,
303            [_, _, 1] => codes.clone(),
304            _ => {
305                return Err(Error::InvalidShape(format!(
306                    "Mimi decode_step expects [batch, codebooks] or [batch, codebooks, 1], got {:?}",
307                    codes.shape()
308                )));
309            }
310        };
311        let latent = self.quantizer.decode(&codes, context)?;
312        let latent = self.upsample.step(&latent, context)?;
313        let latent = self.decoder_transformer.step(&latent, context)?;
314        self.decoder.step(&latent, context)
315    }
316}
317
318/// Exact released-checkpoint requirement for one Mimi parameter.
319#[derive(Debug, Clone, Eq, PartialEq)]
320pub struct MimiParameterRequirement {
321    logical_name: String,
322    checkpoint_key: String,
323    physical_shape: Vec<usize>,
324    logical_shape: Vec<usize>,
325    source_dtype: StoredDtype,
326    output_dtype: RecipeDtype,
327    source_encoding: SourceTensorEncoding,
328    source_bytes: u64,
329    output_bytes: u64,
330    recipe: DerivedWeightRecipe,
331    active: bool,
332}
333
334impl MimiParameterRequirement {
335    /// Stable parameter identity in [`Mimi`]'s authoritative traversal.
336    pub fn logical_name(&self) -> &str {
337        &self.logical_name
338    }
339
340    /// Exact physical tensor name in the released SafeTensors artifact.
341    pub fn checkpoint_key(&self) -> &str {
342        &self.checkpoint_key
343    }
344
345    /// Exact physical checkpoint geometry.
346    pub fn physical_shape(&self) -> &[usize] {
347        &self.physical_shape
348    }
349
350    /// Exact logical destination geometry after the neutral recipe.
351    pub fn logical_shape(&self) -> &[usize] {
352        &self.logical_shape
353    }
354
355    /// Exact scalar representation in the released SafeTensors payload.
356    pub const fn source_dtype(&self) -> &StoredDtype {
357        &self.source_dtype
358    }
359
360    /// Exact scalar representation produced by the neutral recipe.
361    pub const fn output_dtype(&self) -> &RecipeDtype {
362        &self.output_dtype
363    }
364
365    /// Exact physical container encoding required from the admitted source.
366    pub const fn source_encoding(&self) -> &SourceTensorEncoding {
367        &self.source_encoding
368    }
369
370    /// Exact selected source byte count.
371    pub const fn source_bytes(&self) -> u64 {
372        self.source_bytes
373    }
374
375    /// Exact logical output byte count.
376    pub const fn output_bytes(&self) -> u64 {
377        self.output_bytes
378    }
379
380    /// Neutral source-to-logical layout recipe.
381    pub const fn recipe(&self) -> &DerivedWeightRecipe {
382        &self.recipe
383    }
384
385    /// Whether this parameter is materialized for the selected active codebooks.
386    pub const fn is_active(&self) -> bool {
387        self.active
388    }
389}
390
391/// Header-admitted Mimi artifact with an immutable exact parameter plan.
392pub struct PreparedMimiArtifact {
393    config: Config,
394    source: SharedCheckpointSource,
395    requirements: Vec<MimiParameterRequirement>,
396    bindings: Vec<WeightBinding>,
397}
398
399impl std::fmt::Debug for PreparedMimiArtifact {
400    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
401        formatter
402            .debug_struct("PreparedMimiArtifact")
403            .field("config", &self.config)
404            .field("requirements", &self.requirements)
405            .field("bindings", &self.bindings)
406            .finish_non_exhaustive()
407    }
408}
409
410impl PreparedMimiArtifact {
411    /// Released configuration fixed by neutral preparation.
412    pub const fn config(&self) -> &Config {
413        &self.config
414    }
415
416    /// Exact full-checkpoint requirements, including inactive codebooks.
417    pub fn requirements(&self) -> &[MimiParameterRequirement] {
418        &self.requirements
419    }
420
421    /// Active generic runtime bindings selected by preparation.
422    pub fn bindings(&self) -> &[WeightBinding] {
423        &self.bindings
424    }
425
426    /// Performs backend mechanism admission without reading payloads or
427    /// constructing backend-native Mimi tensors.
428    pub fn select<B>(
429        self,
430    ) -> Result<SelectedMimiArtifact<B>, MimiConstructionError<B::ParameterError>>
431    where
432        B: ParameterBackend<Parameter = <B as NeuralBackend>::Tensor>,
433    {
434        let Self {
435            config,
436            source,
437            requirements,
438            bindings,
439        } = self;
440        let selected = select_bindings::<B>(source, bindings)?;
441        Ok(SelectedMimiArtifact {
442            config,
443            requirements,
444            selected,
445        })
446    }
447}
448
449/// Backend-admitted Mimi plan ready for native materialization.
450pub struct SelectedMimiArtifact<B>
451where
452    B: ParameterBackend<Parameter = <B as NeuralBackend>::Tensor>,
453{
454    config: Config,
455    requirements: Vec<MimiParameterRequirement>,
456    selected: SelectedBindingPlan<B>,
457}
458
459impl<B> std::fmt::Debug for SelectedMimiArtifact<B>
460where
461    B: ParameterBackend<Parameter = <B as NeuralBackend>::Tensor>,
462{
463    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
464        formatter
465            .debug_struct("SelectedMimiArtifact")
466            .field("config", &self.config)
467            .field("requirements", &self.requirements)
468            .field("selected", &self.selected)
469            .finish_non_exhaustive()
470    }
471}
472
473impl<B> SelectedMimiArtifact<B>
474where
475    B: ParameterBackend<Parameter = <B as NeuralBackend>::Tensor>,
476{
477    /// Constructs, materializes, validates, and atomically binds the ordinary
478    /// backend-neutral Mimi type.
479    pub fn construct(
480        self,
481        tensor_context: &<<B as NeuralBackend>::Tensor as Tensor>::Context,
482        materialization_context: &B::MaterializationContext,
483    ) -> Result<Mimi<<B as NeuralBackend>::Tensor>, MimiConstructionError<B::ParameterError>> {
484        let SelectedMimiArtifact {
485            config,
486            requirements: _,
487            selected,
488        } = self;
489        let mut mimi = Mimi::new(config, tensor_context)?;
490        let materialized = materialize_selected_bindings::<B>(selected, materialization_context)?;
491        bind_materialized_unit::<B, _>(&mut mimi, materialized)?;
492        Ok(mimi)
493    }
494}
495
496/// Failure while preparing a released Mimi artifact from neutral metadata.
497#[derive(Debug, thiserror::Error)]
498pub enum MimiArtifactError {
499    /// Mimi configuration or pure topology was invalid.
500    #[error(transparent)]
501    Codec(#[from] Error),
502    /// Canonical SafeTensors discovery or header admission failed.
503    #[error(transparent)]
504    Safetensors(#[from] SafetensorsShardError),
505    /// A declarative checkpoint plan was internally invalid.
506    #[error(transparent)]
507    Plan(#[from] CheckpointPlanError),
508    /// The exact released catalog contract was not satisfied.
509    #[error("Mimi SafeTensors catalog is not exact: {0:?}")]
510    Catalog(CheckpointValidation),
511    /// Neutral checkpoint storage could not be opened or inspected.
512    #[error(transparent)]
513    Store(#[from] StoreError),
514    /// A neutral layout recipe was invalid.
515    #[error(transparent)]
516    Recipe(#[from] RecipeError),
517    /// A generic runtime binding declaration was invalid.
518    #[error(transparent)]
519    Binding(#[from] ResidencyDeclarationError),
520    /// Codec-owned parameter topology contradicted its released schema.
521    #[error("invalid Mimi parameter topology: {0}")]
522    Topology(String),
523}
524
525/// Failure while selecting, materializing, or constructing Mimi on a backend.
526#[derive(Debug, thiserror::Error)]
527pub enum MimiConstructionError<E>
528where
529    E: std::error::Error + Send + Sync + 'static,
530{
531    /// Backend-neutral Mimi construction failed.
532    #[error(transparent)]
533    Codec(#[from] Error),
534    /// Generic parameter selection, materialization, validation, or binding failed.
535    #[error(transparent)]
536    Parameters(#[from] ParameterOrchestrationError<E>),
537}
538
539/// Inspects and admits an exact released Mimi SafeTensors artifact without
540/// reading tensor payloads.
541pub fn prepare_checkpoint(
542    path: impl AsRef<Path>,
543    config: Config,
544) -> Result<PreparedMimiArtifact, MimiArtifactError> {
545    let (plan, requirements, bindings) = prepare_catalog(&config)?;
546    let catalog = SafetensorsMetadataCatalog::discover(path)?;
547    let resolution =
548        resolve_safetensors_plan(&catalog, &plan).map_err(MimiArtifactError::Catalog)?;
549    let store = Arc::new(SafetensorsWeightStore::open_admitted(
550        catalog.into_admitted_shards(),
551        1,
552    )?);
553    prepared_from_resolution(config, store, resolution, requirements, bindings)
554}
555
556/// Returns the complete released-checkpoint schema derived from Mimi's
557/// authoritative parameter topology without inspecting an artifact.
558pub fn released_checkpoint_requirements(
559    config: &Config,
560) -> Result<Vec<MimiParameterRequirement>, MimiArtifactError> {
561    prepare_catalog(config).map(|(_, requirements, _)| requirements)
562}
563
564/// Admits an already opened backend-neutral SafeTensors-compatible source.
565/// No tensor payload is acquired during preparation.
566pub fn prepare_source(
567    source: SharedCheckpointSource,
568    config: Config,
569) -> Result<PreparedMimiArtifact, MimiArtifactError> {
570    let (plan, requirements, bindings) = prepare_catalog(&config)?;
571    let resolution =
572        resolve_safetensors_plan(source.as_ref(), &plan).map_err(MimiArtifactError::Catalog)?;
573    prepared_from_resolution(config, source, resolution, requirements, bindings)
574}
575
576/// Selects generic backend mechanisms, then constructs Mimi through the exact
577/// neutral artifact plan.
578pub fn construct<B>(
579    prepared: PreparedMimiArtifact,
580    tensor_context: &<<B as NeuralBackend>::Tensor as Tensor>::Context,
581    materialization_context: &B::MaterializationContext,
582) -> Result<Mimi<<B as NeuralBackend>::Tensor>, MimiConstructionError<B::ParameterError>>
583where
584    B: ParameterBackend<Parameter = <B as NeuralBackend>::Tensor>,
585{
586    prepared
587        .select::<B>()?
588        .construct(tensor_context, materialization_context)
589}
590
591fn checkpoint_parameter_for_key(key: &str) -> Option<String> {
592    transform_decoder_key(key)
593}
594
595fn prepare_catalog(
596    config: &Config,
597) -> Result<
598    (
599        SafetensorsCheckpointPlan,
600        Vec<MimiParameterRequirement>,
601        Vec<WeightBinding>,
602    ),
603    MimiArtifactError,
604> {
605    config.validate()?;
606    let active_topology = parameter_topology(config.clone())?;
607    let mut full_config = config.clone();
608    full_config.num_codebooks = full_config.total_codebooks;
609    let full_topology = parameter_topology(full_config)?;
610    if active_topology
611        .keys()
612        .any(|name| !full_topology.contains_key(name))
613    {
614        return Err(MimiArtifactError::Topology(
615            "active codebook topology is not contained by the released topology".into(),
616        ));
617    }
618
619    let mut physical_names = BTreeSet::new();
620    let mut requirements = Vec::with_capacity(full_topology.len());
621    let mut bindings = Vec::with_capacity(active_topology.len());
622    for (logical_name, logical_shape) in full_topology {
623        let checkpoint_key = checkpoint_key_for_parameter(&logical_name).ok_or_else(|| {
624            MimiArtifactError::Topology(format!(
625                "parameter {logical_name:?} has no released checkpoint identity"
626            ))
627        })?;
628        if !physical_names.insert(checkpoint_key.clone()) {
629            return Err(MimiArtifactError::Topology(format!(
630                "released checkpoint identity {checkpoint_key:?} is duplicated"
631            )));
632        }
633        if checkpoint_parameter_for_key(&checkpoint_key).as_deref() != Some(&logical_name) {
634            return Err(MimiArtifactError::Topology(format!(
635                "checkpoint identity {checkpoint_key:?} does not round-trip to {logical_name:?}"
636            )));
637        }
638        let axes = checkpoint_layout_axes(&logical_name);
639        let physical_shape = match axes {
640            Some(axes) => inverse_permuted_shape(&logical_shape, axes)?,
641            None => logical_shape.clone(),
642        };
643        let source_bytes = f32_bytes(&physical_shape, &checkpoint_key)?;
644        let output_bytes = f32_bytes(&logical_shape, &logical_name)?;
645        let source = DerivedWeightRecipe::source(&checkpoint_key, TensorSelection::Full);
646        let recipe = match axes {
647            Some(axes) => DerivedWeightRecipe::Transpose {
648                input: Box::new(source),
649                axes: axes.to_vec(),
650            },
651            None => source,
652        };
653        let active = active_topology.contains_key(&logical_name);
654        if active {
655            bindings.push(WeightBinding::from_recipe(
656                &logical_name,
657                recipe.clone(),
658                output_bytes,
659            )?);
660        }
661        requirements.push(MimiParameterRequirement {
662            logical_name,
663            checkpoint_key,
664            physical_shape,
665            logical_shape,
666            source_dtype: StoredDtype::F32,
667            output_dtype: RecipeDtype::F32,
668            source_encoding: SourceTensorEncoding::Safetensors(StoredDtype::F32),
669            source_bytes,
670            output_bytes,
671            recipe,
672            active,
673        });
674    }
675    requirements.sort_by(|left, right| left.checkpoint_key.cmp(&right.checkpoint_key));
676    bindings.sort_by(|left, right| left.name().cmp(right.name()));
677    let constraints = requirements
678        .iter()
679        .map(|requirement| {
680            SafetensorsTensorConstraint::required(
681                &requirement.checkpoint_key,
682                requirement.physical_shape.clone(),
683                StoredDtypeConstraint::Exact(requirement.source_dtype.clone()),
684            )
685        })
686        .collect();
687    let plan = SafetensorsCheckpointPlan::new(
688        "mimi-v0.1",
689        constraints,
690        Vec::new(),
691        CatalogPolicy::strict(),
692    )?;
693    Ok((plan, requirements, bindings))
694}
695
696fn prepared_from_resolution(
697    config: Config,
698    source: SharedCheckpointSource,
699    resolution: ResolvedCheckpointPlan,
700    requirements: Vec<MimiParameterRequirement>,
701    bindings: Vec<WeightBinding>,
702) -> Result<PreparedMimiArtifact, MimiArtifactError> {
703    let mut exact_catalog = BTreeMap::new();
704    for requirement in &requirements {
705        let key = requirement.checkpoint_key();
706        let actual_metadata = source.source_metadata(key)?;
707        let expected_metadata = TensorMetadata {
708            name: key.to_owned(),
709            logical_shape: requirement.physical_shape.clone(),
710            physical_shape: requirement.physical_shape.clone(),
711            stored_dtype: requirement.source_dtype.clone(),
712            encoded_byte_len: requirement.source_bytes,
713            backing_shard: actual_metadata.backing_shard.clone(),
714        };
715        let actual_provenance = source.source_provenance(key)?;
716        let expected_provenance = TensorSourceProvenance {
717            catalog_key: key.to_owned(),
718            physical_tensor: key.to_owned(),
719            output: key.to_owned(),
720            backing_shard: expected_metadata.backing_shard.clone(),
721            source_encoding: requirement.source_encoding.clone(),
722        };
723        if actual_metadata != expected_metadata || actual_provenance != expected_provenance {
724            return Err(MimiArtifactError::Topology(format!(
725                "checkpoint tensor {key:?} does not have the exact released SafeTensors provenance"
726            )));
727        }
728        exact_catalog.insert(
729            key.to_owned(),
730            PreparedTensorSource {
731                metadata: expected_metadata,
732                provenance: expected_provenance,
733            },
734        );
735    }
736    let source: SharedCheckpointSource =
737        Arc::new(PreparedCheckpointSource::new(source, exact_catalog)?);
738    let source: SharedCheckpointSource =
739        Arc::new(ResolvedCheckpointSource::new(source, resolution));
740    validate_requirement_recipes(source.as_ref(), &requirements)?;
741    Ok(PreparedMimiArtifact {
742        config,
743        source,
744        requirements,
745        bindings,
746    })
747}
748
749fn validate_requirement_recipes(
750    catalog: &(impl RecipeCatalog + ?Sized),
751    requirements: &[MimiParameterRequirement],
752) -> Result<(), MimiArtifactError> {
753    for requirement in requirements {
754        let actual = requirement.recipe.infer(catalog)?;
755        let expected = RecipeMetadata {
756            shape: requirement.logical_shape.clone(),
757            dtype: requirement.output_dtype.clone(),
758            byte_len: requirement.output_bytes,
759        };
760        if actual != expected {
761            return Err(MimiArtifactError::Topology(format!(
762                "recipe for {:?} produced {actual:?}, expected {expected:?}",
763                requirement.logical_name
764            )));
765        }
766        let metadata = catalog.tensor_metadata(&requirement.checkpoint_key)?;
767        if metadata.encoded_byte_len != requirement.source_bytes {
768            return Err(MimiArtifactError::Topology(format!(
769                "checkpoint tensor {:?} declares {} source bytes, expected {}",
770                requirement.checkpoint_key, metadata.encoded_byte_len, requirement.source_bytes
771            )));
772        }
773    }
774    Ok(())
775}
776
777#[derive(Debug, Clone)]
778struct PlanningTensor(Vec<i32>);
779
780impl PlanningTensor {
781    fn unavailable() -> Result<Self, eredu_nn::Error> {
782        Err(eredu_nn::Error::backend(
783            "Mimi planning tensors cannot execute neural operations",
784        ))
785    }
786}
787
788impl Tensor for PlanningTensor {
789    type Context = ();
790
791    fn shape(&self) -> &[i32] {
792        &self.0
793    }
794
795    fn unloaded_f32(shape: &[i32], _: &Self::Context) -> Result<Self, eredu_nn::Error> {
796        Ok(Self(shape.to_vec()))
797    }
798
799    fn from_f32_slice(_: &[f32], _: &[i32], _: &Self::Context) -> Result<Self, eredu_nn::Error> {
800        Self::unavailable()
801    }
802
803    fn add(&self, _: &Self, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
804        Self::unavailable()
805    }
806
807    fn subtract(&self, _: &Self, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
808        Self::unavailable()
809    }
810
811    fn multiply(&self, _: &Self, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
812        Self::unavailable()
813    }
814
815    fn multiply_scalar(&self, _: f32, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
816        Self::unavailable()
817    }
818
819    fn divide(&self, _: &Self, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
820        Self::unavailable()
821    }
822
823    fn square(&self, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
824        Self::unavailable()
825    }
826
827    fn maximum_scalar(&self, _: f32, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
828        Self::unavailable()
829    }
830
831    fn reshape(&self, _: &[i32], _: &Self::Context) -> Result<Self, eredu_nn::Error> {
832        Self::unavailable()
833    }
834
835    fn transpose_axes(&self, _: &[i32], _: &Self::Context) -> Result<Self, eredu_nn::Error> {
836        Self::unavailable()
837    }
838
839    fn swap_axes(&self, _: i32, _: i32, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
840        Self::unavailable()
841    }
842
843    fn transpose(&self, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
844        Self::unavailable()
845    }
846
847    fn expand_dims(&self, _: i32, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
848        Self::unavailable()
849    }
850
851    fn squeeze_axes(&self, _: &[i32], _: &Self::Context) -> Result<Self, eredu_nn::Error> {
852        Self::unavailable()
853    }
854
855    fn index(&self, _: &[Index], _: &Self::Context) -> Result<Self, eredu_nn::Error> {
856        Self::unavailable()
857    }
858
859    fn take_axis(&self, _: &Self, _: i32, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
860        Self::unavailable()
861    }
862
863    fn concatenate(_: &[Self], _: i32, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
864        Self::unavailable()
865    }
866
867    fn stack(_: &[Self], _: i32, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
868        Self::unavailable()
869    }
870
871    fn matmul(_: &Self, _: &Self, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
872        Self::unavailable()
873    }
874
875    fn sum_axis(_: &Self, _: i32, _: bool, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
876        Self::unavailable()
877    }
878
879    fn argmin_axis(_: &Self, _: i32, _: bool, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
880        Self::unavailable()
881    }
882
883    fn pad(
884        _: &Self,
885        _: &[(i32, i32)],
886        _: PadMode,
887        _: &Self::Context,
888    ) -> Result<Self, eredu_nn::Error> {
889        Self::unavailable()
890    }
891
892    fn conv1d(
893        _: &Self,
894        _: &Self,
895        _: i32,
896        _: i32,
897        _: i32,
898        _: i32,
899        _: &Self::Context,
900    ) -> Result<Self, eredu_nn::Error> {
901        Self::unavailable()
902    }
903
904    fn conv_transpose1d(
905        _: &Self,
906        _: &Self,
907        _: i32,
908        _: i32,
909        _: i32,
910        _: i32,
911        _: i32,
912        _: &Self::Context,
913    ) -> Result<Self, eredu_nn::Error> {
914        Self::unavailable()
915    }
916
917    fn linear(
918        _: &Self,
919        _: &Self,
920        _: Option<&Self>,
921        _: &Self::Context,
922    ) -> Result<Self, eredu_nn::Error> {
923        Self::unavailable()
924    }
925
926    fn layer_norm(
927        _: &Self,
928        _: Option<&Self>,
929        _: Option<&Self>,
930        _: f32,
931        _: &Self::Context,
932    ) -> Result<Self, eredu_nn::Error> {
933        Self::unavailable()
934    }
935
936    fn gelu(_: &Self, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
937        Self::unavailable()
938    }
939
940    fn elu(_: &Self, _: f32, _: &Self::Context) -> Result<Self, eredu_nn::Error> {
941        Self::unavailable()
942    }
943
944    fn rope(
945        _: &Self,
946        _: i32,
947        _: bool,
948        _: f32,
949        _: f32,
950        _: i32,
951        _: &Self::Context,
952    ) -> Result<Self, eredu_nn::Error> {
953        Self::unavailable()
954    }
955
956    fn scaled_dot_product_attention(
957        _: &Self,
958        _: &Self,
959        _: &Self,
960        _: f32,
961        _: AttentionMask<'_, Self>,
962        _: &Self::Context,
963    ) -> Result<Self, eredu_nn::Error> {
964        Self::unavailable()
965    }
966}
967
968fn parameter_topology(config: Config) -> Result<BTreeMap<String, Vec<usize>>, MimiArtifactError> {
969    let mimi = Mimi::<PlanningTensor>::new(config, &())?;
970    let mut topology = BTreeMap::new();
971    let mut duplicate = None;
972    mimi.visit_mimi_parameters("", &mut |metadata, parameter| {
973        let shape = parameter
974            .shape()
975            .iter()
976            .map(|dimension| {
977                usize::try_from(*dimension).map_err(|_| {
978                    MimiArtifactError::Topology(format!(
979                        "parameter {:?} has invalid shape {:?}",
980                        metadata.id,
981                        parameter.shape()
982                    ))
983                })
984            })
985            .collect::<Result<Vec<_>, _>>();
986        match shape {
987            Ok(shape) => {
988                if topology.insert(metadata.id.to_string(), shape).is_some() {
989                    duplicate = Some(metadata.id.to_string());
990                }
991            }
992            Err(error) => duplicate = Some(error.to_string()),
993        }
994    });
995    if let Some(duplicate) = duplicate {
996        return Err(MimiArtifactError::Topology(format!(
997            "duplicate or invalid parameter identity {duplicate:?}"
998        )));
999    }
1000    Ok(topology)
1001}
1002
1003fn checkpoint_layout_axes(parameter: &str) -> Option<[usize; 3]> {
1004    (parameter.ends_with(".weight") && is_conv_weight_key(parameter)).then(|| {
1005        if parameter.contains(".upsample.") {
1006            [1, 2, 0]
1007        } else {
1008            [0, 2, 1]
1009        }
1010    })
1011}
1012
1013fn inverse_permuted_shape(
1014    logical_shape: &[usize],
1015    axes: [usize; 3],
1016) -> Result<Vec<usize>, MimiArtifactError> {
1017    if logical_shape.len() != axes.len() {
1018        return Err(MimiArtifactError::Topology(format!(
1019            "rank-{} parameter cannot use transpose axes {axes:?}",
1020            logical_shape.len()
1021        )));
1022    }
1023    let mut physical = vec![0; axes.len()];
1024    for (logical_axis, physical_axis) in axes.into_iter().enumerate() {
1025        physical[physical_axis] = logical_shape[logical_axis];
1026    }
1027    Ok(physical)
1028}
1029
1030fn f32_bytes(shape: &[usize], name: &str) -> Result<u64, MimiArtifactError> {
1031    shape.iter().try_fold(4u64, |bytes, dimension| {
1032        bytes.checked_mul(*dimension as u64).ok_or_else(|| {
1033            MimiArtifactError::Topology(format!("byte count overflows for {name:?}"))
1034        })
1035    })
1036}
1037
1038fn checkpoint_key_for_parameter(parameter: &str) -> Option<String> {
1039    if parameter.starts_with("quantizer.") {
1040        return Some(parameter.to_owned());
1041    }
1042    if parameter == "downsample.weight" {
1043        return Some("downsample.conv.conv.conv.weight".into());
1044    }
1045    if let Some(key) = parameter.strip_prefix("encoder_transformer.") {
1046        let key = key
1047            .replace(".self_attn.in_proj.weight", ".self_attn.in_proj_weight")
1048            .replace(".mlp.linear1.", ".linear1.")
1049            .replace(".mlp.linear2.", ".linear2.");
1050        return Some(format!("encoder_transformer.transformer.{key}"));
1051    }
1052    if let Some(key) = reverse_seanet_encoder_key(parameter) {
1053        return Some(format!("encoder.model.{key}"));
1054    }
1055    if parameter == "upsample.weight" {
1056        return Some("upsample.convtr.convtr.convtr.weight".into());
1057    }
1058    if let Some(key) = parameter.strip_prefix("decoder_transformer.") {
1059        let key = key
1060            .replace(".self_attn.in_proj.weight", ".self_attn.in_proj_weight")
1061            .replace(".mlp.linear1.", ".linear1.")
1062            .replace(".mlp.linear2.", ".linear2.");
1063        return Some(format!("decoder_transformer.transformer.{key}"));
1064    }
1065    reverse_seanet_decoder_key(parameter).map(|key| format!("decoder.model.{key}"))
1066}
1067
1068const SEANET_ENCODER_KEY_MAPPINGS: &[(&str, &str)] = &[
1069    ("0.conv.conv.", "encoder.init_conv1d."),
1070    (
1071        "1.block.1.conv.conv.",
1072        "encoder.layers.0.residuals.0.block.0.",
1073    ),
1074    (
1075        "1.block.3.conv.conv.",
1076        "encoder.layers.0.residuals.0.block.1.",
1077    ),
1078    ("3.conv.conv.", "encoder.layers.0.downsample."),
1079    (
1080        "4.block.1.conv.conv.",
1081        "encoder.layers.1.residuals.0.block.0.",
1082    ),
1083    (
1084        "4.block.3.conv.conv.",
1085        "encoder.layers.1.residuals.0.block.1.",
1086    ),
1087    ("6.conv.conv.", "encoder.layers.1.downsample."),
1088    (
1089        "7.block.1.conv.conv.",
1090        "encoder.layers.2.residuals.0.block.0.",
1091    ),
1092    (
1093        "7.block.3.conv.conv.",
1094        "encoder.layers.2.residuals.0.block.1.",
1095    ),
1096    ("9.conv.conv.", "encoder.layers.2.downsample."),
1097    (
1098        "10.block.1.conv.conv.",
1099        "encoder.layers.3.residuals.0.block.0.",
1100    ),
1101    (
1102        "10.block.3.conv.conv.",
1103        "encoder.layers.3.residuals.0.block.1.",
1104    ),
1105    ("12.conv.conv.", "encoder.layers.3.downsample."),
1106    ("14.conv.conv.", "encoder.final_conv1d."),
1107];
1108
1109const SEANET_DECODER_KEY_MAPPINGS: &[(&str, &str)] = &[
1110    ("0.conv.conv.", "decoder.init_conv1d."),
1111    ("2.convtr.convtr.", "decoder.layers.0.upsample."),
1112    (
1113        "3.block.1.conv.conv.",
1114        "decoder.layers.0.residuals.0.block.0.",
1115    ),
1116    (
1117        "3.block.3.conv.conv.",
1118        "decoder.layers.0.residuals.0.block.1.",
1119    ),
1120    ("5.convtr.convtr.", "decoder.layers.1.upsample."),
1121    (
1122        "6.block.1.conv.conv.",
1123        "decoder.layers.1.residuals.0.block.0.",
1124    ),
1125    (
1126        "6.block.3.conv.conv.",
1127        "decoder.layers.1.residuals.0.block.1.",
1128    ),
1129    ("8.convtr.convtr.", "decoder.layers.2.upsample."),
1130    (
1131        "9.block.1.conv.conv.",
1132        "decoder.layers.2.residuals.0.block.0.",
1133    ),
1134    (
1135        "9.block.3.conv.conv.",
1136        "decoder.layers.2.residuals.0.block.1.",
1137    ),
1138    ("11.convtr.convtr.", "decoder.layers.3.upsample."),
1139    (
1140        "12.block.1.conv.conv.",
1141        "decoder.layers.3.residuals.0.block.0.",
1142    ),
1143    (
1144        "12.block.3.conv.conv.",
1145        "decoder.layers.3.residuals.0.block.1.",
1146    ),
1147    ("14.conv.conv.", "decoder.final_conv1d."),
1148];
1149
1150fn reverse_seanet_key(parameter: &str, mappings: &[(&str, &str)]) -> Option<String> {
1151    let &(source, target) = mappings
1152        .iter()
1153        .find(|(_, target)| parameter.starts_with(target))?;
1154    Some(format!("{source}{}", &parameter[target.len()..]))
1155}
1156
1157fn reverse_seanet_encoder_key(parameter: &str) -> Option<String> {
1158    reverse_seanet_key(parameter, SEANET_ENCODER_KEY_MAPPINGS)
1159}
1160
1161fn reverse_seanet_decoder_key(parameter: &str) -> Option<String> {
1162    reverse_seanet_key(parameter, SEANET_DECODER_KEY_MAPPINGS)
1163}
1164
1165fn transform_decoder_key(key: &str) -> Option<String> {
1166    if key.starts_with("quantizer.") {
1167        return Some(key.to_string());
1168    }
1169    if key == "downsample.conv.conv.conv.weight" {
1170        return Some("downsample.weight".to_string());
1171    }
1172    if let Some(key) = key.strip_prefix("encoder_transformer.transformer.") {
1173        let key = key
1174            .replace(".self_attn.in_proj_weight", ".self_attn.in_proj.weight")
1175            .replace(".linear1.", ".mlp.linear1.")
1176            .replace(".linear2.", ".mlp.linear2.");
1177        return Some(format!("encoder_transformer.{key}"));
1178    }
1179    if let Some(key) = key.strip_prefix("encoder.model.") {
1180        return transform_seanet_encoder_key(key);
1181    }
1182    if key == "upsample.convtr.convtr.convtr.weight" {
1183        return Some("upsample.weight".to_string());
1184    }
1185    if let Some(key) = key.strip_prefix("decoder_transformer.transformer.") {
1186        let key = key
1187            .replace(".self_attn.in_proj_weight", ".self_attn.in_proj.weight")
1188            .replace(".linear1.", ".mlp.linear1.")
1189            .replace(".linear2.", ".mlp.linear2.");
1190        return Some(format!("decoder_transformer.{key}"));
1191    }
1192    if let Some(key) = key.strip_prefix("decoder.model.") {
1193        return transform_seanet_decoder_key(key);
1194    }
1195    None
1196}
1197
1198fn transform_seanet_encoder_key(key: &str) -> Option<String> {
1199    transform_seanet_key(key, SEANET_ENCODER_KEY_MAPPINGS)
1200}
1201
1202fn transform_seanet_decoder_key(key: &str) -> Option<String> {
1203    transform_seanet_key(key, SEANET_DECODER_KEY_MAPPINGS)
1204}
1205
1206fn transform_seanet_key(key: &str, mappings: &[(&str, &str)]) -> Option<String> {
1207    let &(source, target) = mappings
1208        .iter()
1209        .find(|(source, _)| key.starts_with(source))?;
1210    Some(format!("{target}{}", &key[source.len()..]))
1211}
1212
1213fn is_conv_weight_key(key: &str) -> bool {
1214    key.starts_with("upsample.")
1215        || key.starts_with("downsample.")
1216        || key.contains(".upsample.")
1217        || key.contains(".downsample.")
1218        || key.contains(".init_conv1d.")
1219        || key.contains(".final_conv1d.")
1220        || key.contains(".block.")
1221}
1222
1223impl<T: Tensor> AudioTokenizer for Mimi<T> {
1224    type Tensor = T;
1225
1226    fn config(&self) -> AudioTokenizerConfig {
1227        AudioTokenizerConfig {
1228            sample_rate: self.config.sample_rate,
1229            frame_rate: self.config.frame_rate,
1230            channels: self.config.channels,
1231            codebooks: self.config.num_codebooks,
1232            cardinality: self.config.bins,
1233        }
1234    }
1235
1236    fn encode(&mut self, pcm: &T, context: &T::Context) -> Result<T, Error> {
1237        self.encode(pcm, context)
1238    }
1239
1240    fn decode(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
1241        self.decode(codes, context)
1242    }
1243}
1244
1245#[derive(Debug, Clone)]
1246struct SeaNetEncoder<T: Tensor> {
1247    init_conv1d: StreamableConv1d<T>,
1248    layers: Vec<EncoderLayer<T>>,
1249    final_conv1d: StreamableConv1d<T>,
1250}
1251
1252impl<T: Tensor> SeaNetEncoder<T> {
1253    fn unloaded(context: &T::Context) -> Result<Self, Error> {
1254        let ratios = [4, 5, 6, 8];
1255        let mut channels = 64;
1256        let mut layers = Vec::with_capacity(ratios.len());
1257        for ratio in ratios {
1258            layers.push(EncoderLayer::unloaded(
1259                channels,
1260                channels * 2,
1261                ratio,
1262                context,
1263            )?);
1264            channels *= 2;
1265        }
1266        Ok(Self {
1267            init_conv1d: StreamableConv1d::unloaded(1, 64, 7, 1, context)?,
1268            layers,
1269            final_conv1d: StreamableConv1d::unloaded(1024, 512, 3, 1, context)?,
1270        })
1271    }
1272
1273    fn forward(&mut self, pcm: &T, context: &T::Context) -> Result<T, Error> {
1274        validate_pcm(pcm)?;
1275        let mut x = self.init_conv1d.forward(pcm, context)?;
1276        for layer in &mut self.layers {
1277            x = layer.forward(&x, context)?;
1278        }
1279        self.final_conv1d
1280            .forward(&T::elu(&x, 1.0, context)?, context)
1281    }
1282
1283    fn reset_state(&mut self) {
1284        self.init_conv1d.reset_state();
1285        for layer in &mut self.layers {
1286            layer.reset_state();
1287        }
1288        self.final_conv1d.reset_state();
1289    }
1290
1291    fn step(&mut self, pcm: &T, context: &T::Context) -> Result<Option<T>, Error> {
1292        validate_pcm(pcm)?;
1293        let mut x = match self.init_conv1d.step(pcm, context)? {
1294            Some(x) => x,
1295            None => return Ok(None),
1296        };
1297        for layer in &mut self.layers {
1298            x = match layer.step(&x, context)? {
1299                Some(x) => x,
1300                None => return Ok(None),
1301            };
1302        }
1303        self.final_conv1d.step(&T::elu(&x, 1.0, context)?, context)
1304    }
1305}
1306
1307#[derive(Debug, Clone)]
1308struct EncoderLayer<T: Tensor> {
1309    residuals: Vec<SeaNetResnetBlock<T>>,
1310    downsample: StreamableConv1d<T>,
1311}
1312
1313impl<T: Tensor> EncoderLayer<T> {
1314    fn unloaded(
1315        in_channels: i32,
1316        out_channels: i32,
1317        ratio: i32,
1318        context: &T::Context,
1319    ) -> Result<Self, Error> {
1320        Ok(Self {
1321            residuals: vec![SeaNetResnetBlock::unloaded(in_channels, context)?],
1322            downsample: StreamableConv1d::unloaded(
1323                in_channels,
1324                out_channels,
1325                ratio * 2,
1326                ratio,
1327                context,
1328            )?,
1329        })
1330    }
1331
1332    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1333        let mut x = x.clone();
1334        for residual in &mut self.residuals {
1335            x = residual.forward(&x, context)?;
1336        }
1337        self.downsample.forward(&T::elu(&x, 1.0, context)?, context)
1338    }
1339
1340    fn reset_state(&mut self) {
1341        for residual in &mut self.residuals {
1342            residual.reset_state();
1343        }
1344        self.downsample.reset_state();
1345    }
1346
1347    fn step(&mut self, x: &T, context: &T::Context) -> Result<Option<T>, Error> {
1348        let mut x = x.clone();
1349        for residual in &mut self.residuals {
1350            x = residual.step(&x, context)?;
1351        }
1352        self.downsample.step(&T::elu(&x, 1.0, context)?, context)
1353    }
1354}
1355
1356#[derive(Debug, Clone)]
1357struct MimiTransformer<T: Tensor> {
1358    layers: Vec<MimiTransformerLayer<T>>,
1359}
1360
1361impl<T: Tensor> MimiTransformer<T> {
1362    fn unloaded(context: &T::Context) -> Result<Self, Error> {
1363        Ok(Self {
1364            layers: (0..8)
1365                .map(|_| MimiTransformerLayer::unloaded(context))
1366                .collect::<Result<Vec<_>, _>>()?,
1367        })
1368    }
1369
1370    fn forward(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
1371        let mut x = latent.swap_axes(1, 2, context)?;
1372        for layer in &mut self.layers {
1373            x = layer.forward(&x, context)?;
1374        }
1375        Ok(x.swap_axes(1, 2, context)?)
1376    }
1377
1378    fn reset_state(&mut self) {
1379        for layer in &mut self.layers {
1380            layer.reset_state();
1381        }
1382    }
1383
1384    fn step(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
1385        let mut x = latent.swap_axes(1, 2, context)?;
1386        for layer in &mut self.layers {
1387            x = layer.step(&x, context)?;
1388        }
1389        Ok(x.swap_axes(1, 2, context)?)
1390    }
1391}
1392
1393#[derive(Debug, Clone)]
1394struct MimiTransformerLayer<T: Tensor> {
1395    norm1: LayerNorm<T>,
1396    norm2: LayerNorm<T>,
1397    self_attn: MimiSelfAttention<T>,
1398    mlp: MimiMlp<T>,
1399    layer_scale_1: LayerScale<T>,
1400    layer_scale_2: LayerScale<T>,
1401}
1402
1403impl<T: Tensor> MimiTransformerLayer<T> {
1404    fn unloaded(context: &T::Context) -> Result<Self, Error> {
1405        Ok(Self {
1406            norm1: unloaded_layer_norm(512, 1e-5, context)?,
1407            norm2: unloaded_layer_norm(512, 1e-5, context)?,
1408            self_attn: MimiSelfAttention::unloaded(context)?,
1409            mlp: MimiMlp::unloaded(context)?,
1410            layer_scale_1: LayerScale::unloaded(512, context)?,
1411            layer_scale_2: LayerScale::unloaded(512, context)?,
1412        })
1413    }
1414
1415    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1416        let normed = self.norm1.forward(x, context)?;
1417        let attended = self
1418            .self_attn
1419            .forward(&normed, context)?
1420            .multiply(self.layer_scale_1.scale.as_ref(), context)?;
1421        let x = x.add(&attended, context)?;
1422        let normed = self.norm2.forward(&x, context)?;
1423        let mlp = self
1424            .mlp
1425            .forward(&normed, context)?
1426            .multiply(self.layer_scale_2.scale.as_ref(), context)?;
1427        Ok(x.add(&mlp, context)?)
1428    }
1429
1430    fn reset_state(&mut self) {
1431        self.self_attn.reset_state();
1432    }
1433
1434    fn step(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1435        let normed = self.norm1.forward(x, context)?;
1436        let attended = self
1437            .self_attn
1438            .step(&normed, context)?
1439            .multiply(self.layer_scale_1.scale.as_ref(), context)?;
1440        let x = x.add(&attended, context)?;
1441        let normed = self.norm2.forward(&x, context)?;
1442        let mlp = self
1443            .mlp
1444            .forward(&normed, context)?
1445            .multiply(self.layer_scale_2.scale.as_ref(), context)?;
1446        Ok(x.add(&mlp, context)?)
1447    }
1448}
1449
1450#[derive(Debug, Clone)]
1451struct LayerScale<T: Tensor> {
1452    scale: Parameter<T>,
1453}
1454
1455impl<T: Tensor> LayerScale<T> {
1456    fn unloaded(dim: i32, context: &T::Context) -> Result<Self, Error> {
1457        Ok(Self {
1458            scale: unloaded_parameter(&[dim], context)?,
1459        })
1460    }
1461}
1462
1463#[derive(Debug, Clone)]
1464struct MimiMlp<T: Tensor> {
1465    linear1: Linear<T>,
1466    linear2: Linear<T>,
1467}
1468
1469impl<T: Tensor> MimiMlp<T> {
1470    fn unloaded(context: &T::Context) -> Result<Self, Error> {
1471        Ok(Self {
1472            linear1: unloaded_linear(512, 2048, false, context)?,
1473            linear2: unloaded_linear(2048, 512, false, context)?,
1474        })
1475    }
1476
1477    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1478        let x = self.linear1.forward(x, context)?;
1479        let x = T::gelu(&x, context)?;
1480        Ok(self.linear2.forward(&x, context)?)
1481    }
1482}
1483
1484#[derive(Debug, Clone)]
1485struct MimiSelfAttention<T: Tensor> {
1486    in_proj: Linear<T>,
1487    out_proj: Linear<T>,
1488    rope: Rope,
1489    num_heads: i32,
1490    head_dim: i32,
1491    scale: f32,
1492    context: i32,
1493    key_cache: Option<T>,
1494    value_cache: Option<T>,
1495}
1496
1497impl<T: Tensor> MimiSelfAttention<T> {
1498    fn unloaded(context: &T::Context) -> Result<Self, Error> {
1499        let head_dim = 64;
1500        Ok(Self {
1501            in_proj: unloaded_linear(512, 1536, false, context)?,
1502            out_proj: unloaded_linear(512, 512, false, context)?,
1503            rope: Rope::new(head_dim, true, 10_000.0, 1.0),
1504            num_heads: 8,
1505            head_dim,
1506            scale: (head_dim as f32).sqrt().recip(),
1507            context: 250,
1508            key_cache: None,
1509            value_cache: None,
1510        })
1511    }
1512
1513    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1514        let shape = x.shape();
1515        if shape.len() != 3 || shape[2] != 512 {
1516            return Err(Error::InvalidShape(format!(
1517                "Mimi decoder transformer expects [batch, frames, 512], got {:?}",
1518                x.shape()
1519            )));
1520        }
1521        let (batch, seq, dim) = (shape[0], shape[1], shape[2]);
1522        let qkv = self
1523            .in_proj
1524            .forward(x, context)?
1525            .reshape(&[batch, seq, 3, self.num_heads, self.head_dim], context)?;
1526        let mut q = qkv
1527            .index(
1528                &[
1529                    Index::Full,
1530                    Index::Full,
1531                    Index::At(0),
1532                    Index::Full,
1533                    Index::Full,
1534                ],
1535                context,
1536            )?
1537            .transpose_axes(&[0, 2, 1, 3], context)?;
1538        let mut k = qkv
1539            .index(
1540                &[
1541                    Index::Full,
1542                    Index::Full,
1543                    Index::At(1),
1544                    Index::Full,
1545                    Index::Full,
1546                ],
1547                context,
1548            )?
1549            .transpose_axes(&[0, 2, 1, 3], context)?;
1550        let v = qkv
1551            .index(
1552                &[
1553                    Index::Full,
1554                    Index::Full,
1555                    Index::At(2),
1556                    Index::Full,
1557                    Index::Full,
1558                ],
1559                context,
1560            )?
1561            .transpose_axes(&[0, 2, 1, 3], context)?;
1562        q = self.rope.forward(&q, 0, context)?;
1563        k = self.rope.forward(&k, 0, context)?;
1564        let attended = T::scaled_dot_product_attention(
1565            &q,
1566            &k,
1567            &v,
1568            self.scale,
1569            AttentionMask::Causal,
1570            context,
1571        )?
1572        .transpose_axes(&[0, 2, 1, 3], context)?
1573        .reshape(&[batch, seq, dim], context)?;
1574        Ok(self.out_proj.forward(&attended, context)?)
1575    }
1576
1577    fn reset_state(&mut self) {
1578        self.key_cache = None;
1579        self.value_cache = None;
1580    }
1581
1582    fn step(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1583        let shape = x.shape();
1584        if shape.len() != 3 || shape[2] != 512 {
1585            return Err(Error::InvalidShape(format!(
1586                "Mimi decoder transformer step expects [batch, frames, 512], got {:?}",
1587                x.shape()
1588            )));
1589        }
1590        let (batch, seq, dim) = (shape[0], shape[1], shape[2]);
1591        let prev_len = self
1592            .key_cache
1593            .as_ref()
1594            .map(|cache| cache.dim(2))
1595            .unwrap_or(0);
1596        let qkv = self
1597            .in_proj
1598            .forward(x, context)?
1599            .reshape(&[batch, seq, 3, self.num_heads, self.head_dim], context)?;
1600        let mut q = qkv
1601            .index(
1602                &[
1603                    Index::Full,
1604                    Index::Full,
1605                    Index::At(0),
1606                    Index::Full,
1607                    Index::Full,
1608                ],
1609                context,
1610            )?
1611            .transpose_axes(&[0, 2, 1, 3], context)?;
1612        let mut k = qkv
1613            .index(
1614                &[
1615                    Index::Full,
1616                    Index::Full,
1617                    Index::At(1),
1618                    Index::Full,
1619                    Index::Full,
1620                ],
1621                context,
1622            )?
1623            .transpose_axes(&[0, 2, 1, 3], context)?;
1624        let v = qkv
1625            .index(
1626                &[
1627                    Index::Full,
1628                    Index::Full,
1629                    Index::At(2),
1630                    Index::Full,
1631                    Index::Full,
1632                ],
1633                context,
1634            )?
1635            .transpose_axes(&[0, 2, 1, 3], context)?;
1636        q = self.rope.forward(&q, prev_len, context)?;
1637        k = self.rope.forward(&k, prev_len, context)?;
1638
1639        let mut keys = match self.key_cache.take() {
1640            Some(prev) => T::concatenate(&[prev, k], 2, context)?,
1641            None => k,
1642        };
1643        let mut values = match self.value_cache.take() {
1644            Some(prev) => T::concatenate(&[prev, v], 2, context)?,
1645            None => v,
1646        };
1647        let key_len = keys.dim(2);
1648        if key_len > self.context + seq {
1649            let start = key_len - (self.context + seq);
1650            keys = keys.index(
1651                &[
1652                    Index::Full,
1653                    Index::Full,
1654                    Index::Range(start, key_len),
1655                    Index::Full,
1656                ],
1657                context,
1658            )?;
1659            values = values.index(
1660                &[
1661                    Index::Full,
1662                    Index::Full,
1663                    Index::Range(start, key_len),
1664                    Index::Full,
1665                ],
1666                context,
1667            )?;
1668        }
1669        let retained_prev_len = keys.dim(2) - seq;
1670        let mask =
1671            streaming_attention_mask::<T>(batch, seq, retained_prev_len, self.context, context)?;
1672        let attended = T::scaled_dot_product_attention(
1673            &q,
1674            &keys,
1675            &values,
1676            self.scale,
1677            AttentionMask::Tensor(&mask),
1678            context,
1679        )?
1680        .transpose_axes(&[0, 2, 1, 3], context)?
1681        .reshape(&[batch, seq, dim], context)?;
1682        self.key_cache = Some(keys);
1683        self.value_cache = Some(values);
1684        Ok(self.out_proj.forward(&attended, context)?)
1685    }
1686}
1687
1688fn streaming_attention_mask<T: Tensor>(
1689    batch: i32,
1690    query_len: i32,
1691    prev_len: i32,
1692    attention_context: i32,
1693    execution: &T::Context,
1694) -> Result<T, Error> {
1695    let key_len = prev_len + query_len;
1696    let mut mask = Vec::with_capacity((batch * query_len * key_len) as usize);
1697    for _ in 0..batch {
1698        for q in 0..query_len {
1699            let q_pos = prev_len + q;
1700            for k in 0..key_len {
1701                if k <= q_pos && q_pos <= k + attention_context {
1702                    mask.push(0.0f32);
1703                } else {
1704                    mask.push(f32::NEG_INFINITY);
1705                }
1706            }
1707        }
1708    }
1709    Ok(T::from_f32_slice(
1710        &mask,
1711        &[batch, 1, query_len, key_len],
1712        execution,
1713    )?)
1714}
1715
1716#[derive(Debug, Clone)]
1717struct SeaNetDecoder<T: Tensor> {
1718    init_conv1d: StreamableConv1d<T>,
1719    layers: Vec<DecoderLayer<T>>,
1720    final_conv1d: StreamableConv1d<T>,
1721}
1722
1723impl<T: Tensor> SeaNetDecoder<T> {
1724    fn unloaded(context: &T::Context) -> Result<Self, Error> {
1725        let ratios = [8, 6, 5, 4];
1726        let mut channels = 1024;
1727        let mut layers = Vec::with_capacity(ratios.len());
1728        for ratio in ratios {
1729            let out_channels = channels / 2;
1730            layers.push(DecoderLayer::unloaded(
1731                channels,
1732                out_channels,
1733                ratio,
1734                context,
1735            )?);
1736            channels = out_channels;
1737        }
1738        Ok(Self {
1739            init_conv1d: StreamableConv1d::unloaded(512, 1024, 7, 1, context)?,
1740            layers,
1741            final_conv1d: StreamableConv1d::unloaded(64, 1, 3, 1, context)?,
1742        })
1743    }
1744
1745    fn forward(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
1746        let mut x = self.init_conv1d.forward(latent, context)?;
1747        for layer in &mut self.layers {
1748            x = layer.forward(&T::elu(&x, 1.0, context)?, context)?;
1749        }
1750        self.final_conv1d
1751            .forward(&T::elu(&x, 1.0, context)?, context)
1752    }
1753
1754    fn reset_state(&mut self) {
1755        self.init_conv1d.reset_state();
1756        for layer in &mut self.layers {
1757            layer.reset_state();
1758        }
1759        self.final_conv1d.reset_state();
1760    }
1761
1762    fn step(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
1763        let mut x = self.init_conv1d.step(latent, context)?.ok_or_else(|| {
1764            Error::InvalidShape("Mimi decoder init conv produced no streaming output".into())
1765        })?;
1766        for layer in &mut self.layers {
1767            x = layer.step(&T::elu(&x, 1.0, context)?, context)?;
1768        }
1769        self.final_conv1d
1770            .step(&T::elu(&x, 1.0, context)?, context)?
1771            .ok_or_else(|| Error::InvalidShape("Mimi decoder final conv produced no output".into()))
1772    }
1773}
1774
1775#[derive(Debug, Clone)]
1776struct DecoderLayer<T: Tensor> {
1777    upsample: StreamableConvTranspose1d<T>,
1778    residuals: Vec<SeaNetResnetBlock<T>>,
1779}
1780
1781impl<T: Tensor> DecoderLayer<T> {
1782    fn unloaded(
1783        in_channels: i32,
1784        out_channels: i32,
1785        ratio: i32,
1786        context: &T::Context,
1787    ) -> Result<Self, Error> {
1788        Ok(Self {
1789            upsample: StreamableConvTranspose1d::unloaded(
1790                in_channels,
1791                out_channels,
1792                ratio * 2,
1793                ratio,
1794                1,
1795                true,
1796                context,
1797            )?,
1798            residuals: vec![SeaNetResnetBlock::unloaded(out_channels, context)?],
1799        })
1800    }
1801
1802    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1803        let mut x = self.upsample.forward(x, context)?;
1804        for residual in &mut self.residuals {
1805            x = residual.forward(&x, context)?;
1806        }
1807        Ok(x)
1808    }
1809
1810    fn reset_state(&mut self) {
1811        self.upsample.reset_state();
1812        for residual in &mut self.residuals {
1813            residual.reset_state();
1814        }
1815    }
1816
1817    fn step(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1818        let mut x = self.upsample.step(x, context)?;
1819        for residual in &mut self.residuals {
1820            x = residual.step(&x, context)?;
1821        }
1822        Ok(x)
1823    }
1824}
1825
1826#[derive(Debug, Clone)]
1827struct SeaNetResnetBlock<T: Tensor> {
1828    block: Vec<StreamableConv1d<T>>,
1829}
1830
1831impl<T: Tensor> SeaNetResnetBlock<T> {
1832    fn unloaded(channels: i32, context: &T::Context) -> Result<Self, Error> {
1833        Ok(Self {
1834            block: vec![
1835                StreamableConv1d::unloaded(channels, channels / 2, 3, 1, context)?,
1836                StreamableConv1d::unloaded(channels / 2, channels, 1, 1, context)?,
1837            ],
1838        })
1839    }
1840
1841    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1842        let mut y = x.clone();
1843        for conv in &mut self.block {
1844            y = conv.forward(&T::elu(&y, 1.0, context)?, context)?;
1845        }
1846        Ok(y.add(x, context)?)
1847    }
1848
1849    fn reset_state(&mut self) {
1850        for conv in &mut self.block {
1851            conv.reset_state();
1852        }
1853    }
1854
1855    fn step(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1856        let mut y = x.clone();
1857        for conv in &mut self.block {
1858            y = conv
1859                .step(&T::elu(&y, 1.0, context)?, context)?
1860                .ok_or_else(|| {
1861                    Error::InvalidShape("Mimi residual conv produced no output".into())
1862                })?;
1863        }
1864        Ok(y.add(x, context)?)
1865    }
1866}
1867
1868#[derive(Debug, Clone)]
1869struct StreamableConv1d<T: Tensor> {
1870    weight: Parameter<T>,
1871    bias: Option<Parameter<T>>,
1872    stride: i32,
1873    dilation: i32,
1874    groups: i32,
1875    pad_mode: PadMode,
1876    state_prev_xs: Option<T>,
1877    left_pad_applied: bool,
1878}
1879
1880impl<T: Tensor> StreamableConv1d<T> {
1881    fn unloaded(
1882        in_channels: i32,
1883        out_channels: i32,
1884        kernel_size: i32,
1885        stride: i32,
1886        context: &T::Context,
1887    ) -> Result<Self, Error> {
1888        Self::unloaded_with_pad_mode(
1889            in_channels,
1890            out_channels,
1891            kernel_size,
1892            stride,
1893            true,
1894            PadMode::Constant,
1895            context,
1896        )
1897    }
1898
1899    fn unloaded_with_pad_mode(
1900        in_channels: i32,
1901        out_channels: i32,
1902        kernel_size: i32,
1903        stride: i32,
1904        bias: bool,
1905        pad_mode: PadMode,
1906        context: &T::Context,
1907    ) -> Result<Self, Error> {
1908        Ok(Self {
1909            weight: unloaded_parameter(&[out_channels, kernel_size, in_channels], context)?,
1910            bias: bias
1911                .then(|| unloaded_parameter(&[out_channels], context))
1912                .transpose()?,
1913            stride,
1914            dilation: 1,
1915            groups: 1,
1916            pad_mode,
1917            state_prev_xs: None,
1918            left_pad_applied: false,
1919        })
1920    }
1921
1922    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1923        let kernel_size = self.weight.as_ref().dim(1);
1924        let effective_kernel = (kernel_size - 1) * self.dilation + 1;
1925        let padding_total = effective_kernel - self.stride;
1926        let extra_padding =
1927            extra_padding_for_conv1d(x.dim(2), effective_kernel, self.stride, padding_total);
1928        let x = pad_bct(x, padding_total, extra_padding, self.pad_mode, context)?;
1929        let x = x.swap_axes(1, 2, context)?;
1930        let mut y = T::conv1d(
1931            &x,
1932            self.weight.as_ref(),
1933            self.stride,
1934            0,
1935            self.dilation,
1936            self.groups,
1937            context,
1938        )?;
1939        if let Some(bias) = &self.bias {
1940            y = y.add(bias.as_ref(), context)?;
1941        }
1942        Ok(y.swap_axes(1, 2, context)?)
1943    }
1944
1945    fn reset_state(&mut self) {
1946        self.state_prev_xs = None;
1947        self.left_pad_applied = false;
1948    }
1949
1950    fn step(&mut self, x: &T, context: &T::Context) -> Result<Option<T>, Error> {
1951        let kernel_size = self.weight.as_ref().dim(1);
1952        let effective_kernel = (kernel_size - 1) * self.dilation + 1;
1953        let padding_total = effective_kernel - self.stride;
1954        let x = if self.left_pad_applied {
1955            x.clone()
1956        } else {
1957            self.left_pad_applied = true;
1958            pad_bct(x, padding_total, 0, self.pad_mode, context)?
1959        };
1960        let x = match self.state_prev_xs.take() {
1961            Some(prev) => T::concatenate(&[prev, x], 2, context)?,
1962            None => x,
1963        };
1964        let seq_len = x.dim(2);
1965        let num_frames = (seq_len + self.stride).saturating_sub(effective_kernel) / self.stride;
1966        if num_frames <= 0 {
1967            self.state_prev_xs = Some(x);
1968            return Ok(None);
1969        }
1970        let offset = num_frames * self.stride;
1971        self.state_prev_xs = Some(x.index(
1972            &[Index::Full, Index::Full, Index::Range(offset, seq_len)],
1973            context,
1974        )?);
1975        let in_len = (num_frames - 1) * self.stride + effective_kernel;
1976        let x = x.index(
1977            &[Index::Full, Index::Full, Index::Range(0, in_len)],
1978            context,
1979        )?;
1980        self.forward_unpadded(&x, context).map(Some)
1981    }
1982
1983    fn forward_unpadded(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
1984        let x = x.swap_axes(1, 2, context)?;
1985        let mut y = T::conv1d(
1986            &x,
1987            self.weight.as_ref(),
1988            self.stride,
1989            0,
1990            self.dilation,
1991            self.groups,
1992            context,
1993        )?;
1994        if let Some(bias) = &self.bias {
1995            y = y.add(bias.as_ref(), context)?;
1996        }
1997        Ok(y.swap_axes(1, 2, context)?)
1998    }
1999}
2000
2001#[derive(Debug, Clone)]
2002struct StreamableConvTranspose1d<T: Tensor> {
2003    weight: Parameter<T>,
2004    bias: Option<Parameter<T>>,
2005    kernel_size: i32,
2006    stride: i32,
2007    groups: i32,
2008    state_prev_ys: Option<T>,
2009}
2010
2011impl<T: Tensor> StreamableConvTranspose1d<T> {
2012    fn unloaded(
2013        in_channels: i32,
2014        out_channels: i32,
2015        kernel_size: i32,
2016        stride: i32,
2017        groups: i32,
2018        bias: bool,
2019        context: &T::Context,
2020    ) -> Result<Self, Error> {
2021        Ok(Self {
2022            weight: unloaded_parameter(
2023                &[out_channels, kernel_size, in_channels / groups],
2024                context,
2025            )?,
2026            bias: bias
2027                .then(|| unloaded_parameter(&[out_channels], context))
2028                .transpose()?,
2029            kernel_size,
2030            stride,
2031            groups,
2032            state_prev_ys: None,
2033        })
2034    }
2035
2036    fn forward(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
2037        let y = self.forward_untrimmed(x, context)?;
2038        let padding_total = self.kernel_size.saturating_sub(self.stride);
2039        unpad_bct(&y, 0, padding_total, context)
2040    }
2041
2042    fn reset_state(&mut self) {
2043        self.state_prev_ys = None;
2044    }
2045
2046    fn step(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
2047        let y = self.forward_untrimmed(x, context)?;
2048        let out_len = y.dim(2);
2049        let y = match self.state_prev_ys.take() {
2050            None => y,
2051            Some(prev) => {
2052                let prev_len = prev.dim(2);
2053                let prev = match &self.bias {
2054                    None => prev,
2055                    Some(bias) => prev.subtract(
2056                        &bias
2057                            .as_ref()
2058                            .reshape(&[1, bias.as_ref().dim(0), 1], context)?,
2059                        context,
2060                    )?,
2061                };
2062                let y1 = y
2063                    .index(
2064                        &[Index::Full, Index::Full, Index::Range(0, prev_len)],
2065                        context,
2066                    )?
2067                    .add(&prev, context)?;
2068                let y2 = y.index(
2069                    &[Index::Full, Index::Full, Index::Range(prev_len, out_len)],
2070                    context,
2071                )?;
2072                T::concatenate(&[y1, y2], 2, context)?
2073            }
2074        };
2075        let invalid_steps = self.kernel_size - self.stride;
2076        let split = out_len - invalid_steps;
2077        let out = y.index(&[Index::Full, Index::Full, Index::Range(0, split)], context)?;
2078        self.state_prev_ys = Some(y.index(
2079            &[Index::Full, Index::Full, Index::Range(split, out_len)],
2080            context,
2081        )?);
2082        Ok(out)
2083    }
2084
2085    fn forward_untrimmed(&mut self, x: &T, context: &T::Context) -> Result<T, Error> {
2086        let x = x.swap_axes(1, 2, context)?;
2087        let mut y = T::conv_transpose1d(
2088            &x,
2089            self.weight.as_ref(),
2090            self.stride,
2091            0,
2092            1,
2093            0,
2094            self.groups,
2095            context,
2096        )?;
2097        if let Some(bias) = &self.bias {
2098            y = y.add(bias.as_ref(), context)?;
2099        }
2100        Ok(y.swap_axes(1, 2, context)?)
2101    }
2102}
2103
2104fn extra_padding_for_conv1d(len: i32, kernel_size: i32, stride: i32, padding_total: i32) -> i32 {
2105    let n_frames = (len + padding_total - kernel_size) as f64 / stride as f64 + 1.0;
2106    let ideal_len = ((n_frames.ceil() as i32 - 1) * stride + kernel_size) - padding_total;
2107    ideal_len.saturating_sub(len)
2108}
2109
2110fn pad_bct<T: Tensor>(
2111    x: &T,
2112    left: i32,
2113    right: i32,
2114    mode: PadMode,
2115    context: &T::Context,
2116) -> Result<T, Error> {
2117    Ok(T::pad(x, &[(0, 0), (0, 0), (left, right)], mode, context)?)
2118}
2119
2120fn unpad_bct<T: Tensor>(x: &T, left: i32, right: i32, context: &T::Context) -> Result<T, Error> {
2121    let len = x.dim(2);
2122    if len < left + right {
2123        return Err(Error::InvalidShape(format!(
2124            "cannot unpad Mimi tensor of length {len} by {left}+{right}"
2125        )));
2126    }
2127    Ok(x.index(
2128        &[Index::Full, Index::Full, Index::Range(left, len - right)],
2129        context,
2130    )?)
2131}
2132
2133/// Split residual vector quantizer used by Mimi.
2134#[derive(Debug, Clone)]
2135pub struct SplitResidualVectorQuantizer<T: Tensor> {
2136    /// First semantic codebook branch.
2137    pub rvq_first: ResidualVectorQuantizer<T>,
2138    /// Remaining acoustic codebook branch.
2139    pub rvq_rest: ResidualVectorQuantizer<T>,
2140    n_q: i32,
2141}
2142
2143impl<T: Tensor> SplitResidualVectorQuantizer<T> {
2144    fn unloaded(config: &Config, context: &T::Context) -> Result<Self, Error> {
2145        Ok(Self {
2146            rvq_first: ResidualVectorQuantizer::unloaded(
2147                config.latent_dim,
2148                config.quantizer_dim,
2149                1,
2150                config.bins,
2151                context,
2152            )?,
2153            rvq_rest: ResidualVectorQuantizer::unloaded(
2154                config.latent_dim,
2155                config.quantizer_dim,
2156                config.num_codebooks - 1,
2157                config.bins,
2158                context,
2159            )?,
2160            n_q: config.num_codebooks,
2161        })
2162    }
2163
2164    /// Encodes latent frames shaped `[batch, 512, frames]`.
2165    pub fn encode(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
2166        validate_latent(latent)?;
2167        let first = self.rvq_first.encode(latent, context)?;
2168        if self.n_q == 1 {
2169            Ok(first)
2170        } else {
2171            let rest = self.rvq_rest.encode(latent, context)?;
2172            Ok(T::concatenate(&[first, rest], 1, context)?)
2173        }
2174    }
2175
2176    /// Decodes tokens shaped `[batch, codebooks, frames]`.
2177    pub fn decode(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
2178        validate_codes(codes, self.n_q)?;
2179        let first_codes = codes.index(&[Index::Full, Index::Range(0, 1), Index::Full], context)?;
2180        let mut quantized = self.rvq_first.decode(&first_codes, context)?;
2181        if codes.dim(1) > 1 {
2182            let rest_codes = codes.index(
2183                &[Index::Full, Index::Range(1, codes.dim(1)), Index::Full],
2184                context,
2185            )?;
2186            quantized = quantized.add(&self.rvq_rest.decode(&rest_codes, context)?, context)?;
2187        }
2188        Ok(quantized)
2189    }
2190}
2191
2192/// Residual vector quantizer branch.
2193#[derive(Debug, Clone)]
2194pub struct ResidualVectorQuantizer<T: Tensor> {
2195    /// Input projection from Mimi latent dimension into codebook dimension.
2196    pub input_proj: Conv1x1NoBias<T>,
2197    /// Output projection from codebook dimension back into Mimi latent dimension.
2198    pub output_proj: Conv1x1NoBias<T>,
2199    /// Residual codebook layers.
2200    pub vq: ResidualVectorQuantization<T>,
2201}
2202
2203impl<T: Tensor> ResidualVectorQuantizer<T> {
2204    fn unloaded(
2205        latent_dim: i32,
2206        codebook_dim: i32,
2207        layers: i32,
2208        bins: i32,
2209        context: &T::Context,
2210    ) -> Result<Self, Error> {
2211        Ok(Self {
2212            input_proj: Conv1x1NoBias::unloaded(latent_dim, codebook_dim, context)?,
2213            output_proj: Conv1x1NoBias::unloaded(codebook_dim, latent_dim, context)?,
2214            vq: ResidualVectorQuantization::unloaded(layers, codebook_dim, bins, context)?,
2215        })
2216    }
2217
2218    fn encode(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
2219        self.vq
2220            .encode(&self.input_proj.forward(latent, context)?, context)
2221    }
2222
2223    fn decode(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
2224        self.output_proj
2225            .forward(&self.vq.decode(codes, context)?, context)
2226    }
2227}
2228
2229/// Residual vector quantization layers.
2230#[derive(Debug, Clone)]
2231pub struct ResidualVectorQuantization<T: Tensor> {
2232    /// Ordered residual quantization layers.
2233    pub layers: Vec<VectorQuantization<T>>,
2234}
2235
2236impl<T: Tensor> ResidualVectorQuantization<T> {
2237    fn unloaded(layers: i32, dim: i32, bins: i32, context: &T::Context) -> Result<Self, Error> {
2238        Ok(Self {
2239            layers: (0..layers)
2240                .map(|_| VectorQuantization::unloaded(dim, bins, context))
2241                .collect::<Result<Vec<_>, _>>()?,
2242        })
2243    }
2244
2245    fn encode(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
2246        if self.layers.is_empty() {
2247            return Err(Error::InvalidShape("Mimi RVQ has no layers".into()));
2248        }
2249        let mut residual = latent.clone();
2250        let mut codes = Vec::with_capacity(self.layers.len());
2251        for layer in &mut self.layers {
2252            let indices = layer.encode(&residual, context)?;
2253            let quantized = layer.decode_one(&indices, context)?;
2254            residual = residual.subtract(&quantized, context)?;
2255            codes.push(indices);
2256        }
2257        Ok(T::stack(&codes, 1, context)?)
2258    }
2259
2260    fn decode(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
2261        if codes.dim(1) != self.layers.len() as i32 {
2262            return Err(Error::InvalidShape(format!(
2263                "Mimi RVQ expected {} codebooks, got {:?}",
2264                self.layers.len(),
2265                codes.shape()
2266            )));
2267        }
2268        let mut out: Option<T> = None;
2269        for (index, layer) in self.layers.iter_mut().enumerate() {
2270            let code = codes.index(
2271                &[Index::Full, Index::At(index as i32), Index::Full],
2272                context,
2273            )?;
2274            let quantized = layer.decode_one(&code, context)?;
2275            out = Some(match out {
2276                None => quantized,
2277                Some(prev) => prev.add(&quantized, context)?,
2278            });
2279        }
2280        out.ok_or_else(|| Error::InvalidShape("Mimi RVQ has no layers".into()))
2281    }
2282}
2283
2284/// Single vector-quantization layer.
2285#[derive(Debug, Clone)]
2286pub struct VectorQuantization<T: Tensor> {
2287    /// Euclidean codebook.
2288    pub _codebook: EuclideanCodebook<T>,
2289}
2290
2291impl<T: Tensor> VectorQuantization<T> {
2292    fn unloaded(dim: i32, bins: i32, context: &T::Context) -> Result<Self, Error> {
2293        Ok(Self {
2294            _codebook: EuclideanCodebook::unloaded(dim, bins, context)?,
2295        })
2296    }
2297
2298    fn encode(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
2299        let latent = latent.swap_axes(1, 2, context)?;
2300        self._codebook.encode(&latent, context)
2301    }
2302
2303    fn decode_one(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
2304        self._codebook
2305            .decode(codes, context)?
2306            .swap_axes(1, 2, context)
2307            .map_err(Into::into)
2308    }
2309}
2310
2311/// Euclidean codebook backed by EMA cluster statistics.
2312#[derive(Debug, Clone)]
2313pub struct EuclideanCodebook<T: Tensor> {
2314    /// Checkpoint initialization flag.
2315    pub _initialized: Parameter<T>,
2316    /// EMA cluster usage.
2317    pub cluster_usage: Parameter<T>,
2318    /// EMA embedding sum.
2319    pub embedding_sum: Parameter<T>,
2320}
2321
2322impl<T: Tensor> EuclideanCodebook<T> {
2323    fn unloaded(dim: i32, bins: i32, context: &T::Context) -> Result<Self, Error> {
2324        Ok(Self {
2325            _initialized: unloaded_parameter(&[1], context)?,
2326            cluster_usage: unloaded_parameter(&[bins], context)?,
2327            embedding_sum: unloaded_parameter(&[bins, dim], context)?,
2328        })
2329    }
2330
2331    fn embedding(&self, context: &T::Context) -> Result<T, Error> {
2332        let usage = self
2333            .cluster_usage
2334            .as_ref()
2335            .maximum_scalar(EPSILON, context)?
2336            .expand_dims(1, context)?;
2337        Ok(self.embedding_sum.as_ref().divide(&usage, context)?)
2338    }
2339
2340    fn encode(&self, latent_btd: &T, context: &T::Context) -> Result<T, Error> {
2341        if latent_btd.shape().len() != 3 {
2342            return Err(Error::InvalidShape(format!(
2343                "Mimi codebook encode expects [batch, frames, dim], got {:?}",
2344                latent_btd.shape()
2345            )));
2346        }
2347        let batch = latent_btd.dim(0);
2348        let frames = latent_btd.dim(1);
2349        let dim = latent_btd.dim(2);
2350        let flat = latent_btd.reshape(&[batch * frames, dim], context)?;
2351        let embedding = self.embedding(context)?;
2352        let x2 = T::sum_axis(&flat.square(context)?, -1, true, context)?;
2353        let e2 = T::sum_axis(&embedding.square(context)?, -1, false, context)?
2354            .expand_dims(0, context)?;
2355        let dot = T::matmul(&flat, &embedding.transpose(context)?, context)?;
2356        let dists = x2
2357            .add(&e2, context)?
2358            .subtract(&dot.multiply_scalar(2.0, context)?, context)?;
2359        Ok(T::argmin_axis(&dists, -1, false, context)?.reshape(&[batch, frames], context)?)
2360    }
2361
2362    fn decode(&self, codes: &T, context: &T::Context) -> Result<T, Error> {
2363        if codes.shape().len() != 2 {
2364            return Err(Error::InvalidShape(format!(
2365                "Mimi codebook decode expects [batch, frames], got {:?}",
2366                codes.shape()
2367            )));
2368        }
2369        let batch = codes.dim(0);
2370        let frames = codes.dim(1);
2371        let embedding = self.embedding(context)?;
2372        let flat = codes.reshape(&[batch * frames], context)?;
2373        Ok(embedding
2374            .take_axis(&flat, 0, context)?
2375            .reshape(&[batch, frames, embedding.dim(1)], context)?)
2376    }
2377}
2378
2379/// Bias-free 1x1 convolution over `[batch, channels, frames]` tensors.
2380#[derive(Debug, Clone)]
2381pub struct Conv1x1NoBias<T: Tensor> {
2382    /// Weight shaped `[out_channels, in_channels, 1]`.
2383    pub weight: Parameter<T>,
2384}
2385
2386impl<T: Tensor> Conv1x1NoBias<T> {
2387    fn unloaded(in_channels: i32, out_channels: i32, context: &T::Context) -> Result<Self, Error> {
2388        Ok(Self {
2389            weight: unloaded_parameter(&[out_channels, in_channels, 1], context)?,
2390        })
2391    }
2392
2393    fn forward(&self, latent: &T, context: &T::Context) -> Result<T, Error> {
2394        if latent.shape().len() != 3 {
2395            return Err(Error::InvalidShape(format!(
2396                "Mimi 1x1 projection expects [batch, channels, frames], got {:?}",
2397                latent.shape()
2398            )));
2399        }
2400        let x = latent.swap_axes(1, 2, context)?;
2401        let weight = self.weight.as_ref().squeeze_axes(&[-1], context)?;
2402        Ok(T::matmul(&x, &weight.transpose(context)?, context)?.swap_axes(1, 2, context)?)
2403    }
2404}
2405
2406trait MimiModuleParameters<T: Tensor> {
2407    fn visit_mimi_parameters<'a>(
2408        &'a self,
2409        prefix: &str,
2410        visitor: &mut dyn FnMut(ParameterMetadata, &'a T),
2411    );
2412    fn visit_mimi_parameters_mut<'a>(
2413        &'a mut self,
2414        prefix: &str,
2415        visitor: &mut dyn FnMut(ParameterMetadata, &'a mut T),
2416    );
2417    fn set_mimi_trainable(&mut self, trainable: bool);
2418}
2419
2420struct PrefixVisitor<'a, F: ?Sized> {
2421    prefix: &'a str,
2422    visitor: &'a mut F,
2423    exact: bool,
2424}
2425
2426impl<'a, 'value, T, F: ?Sized> ParameterVisitor<'value, T> for PrefixVisitor<'a, F>
2427where
2428    T: 'value,
2429    F: FnMut(ParameterMetadata, &'value T),
2430{
2431    fn visit(&mut self, mut metadata: ParameterMetadata, value: &'value T) {
2432        let id = if self.exact {
2433            self.prefix.to_owned()
2434        } else {
2435            parameter_name(self.prefix, metadata.id.as_str())
2436        };
2437        metadata.id = ParameterId::new(id).expect("Mimi parameter identities are non-empty");
2438        (self.visitor)(metadata, value);
2439    }
2440}
2441
2442struct PrefixVisitorMut<'a, F: ?Sized> {
2443    prefix: &'a str,
2444    visitor: &'a mut F,
2445    exact: bool,
2446}
2447
2448impl<'a, 'value, T, F: ?Sized> ParameterVisitorMut<'value, T> for PrefixVisitorMut<'a, F>
2449where
2450    T: 'value,
2451    F: FnMut(ParameterMetadata, &'value mut T),
2452{
2453    fn visit_mut(&mut self, mut metadata: ParameterMetadata, value: &'value mut T) {
2454        let id = if self.exact {
2455            self.prefix.to_owned()
2456        } else {
2457            parameter_name(self.prefix, metadata.id.as_str())
2458        };
2459        metadata.id = ParameterId::new(id).expect("Mimi parameter identities are non-empty");
2460        (self.visitor)(metadata, value);
2461    }
2462}
2463
2464impl<T: Tensor> MimiModuleParameters<T> for Parameter<T> {
2465    fn visit_mimi_parameters<'a>(
2466        &'a self,
2467        prefix: &str,
2468        visitor: &mut dyn FnMut(ParameterMetadata, &'a T),
2469    ) {
2470        self.visit_parameters(&mut PrefixVisitor {
2471            prefix,
2472            visitor,
2473            exact: true,
2474        });
2475    }
2476
2477    fn visit_mimi_parameters_mut<'a>(
2478        &'a mut self,
2479        prefix: &str,
2480        visitor: &mut dyn FnMut(ParameterMetadata, &'a mut T),
2481    ) {
2482        self.visit_parameters_mut(&mut PrefixVisitorMut {
2483            prefix,
2484            visitor,
2485            exact: true,
2486        });
2487    }
2488
2489    fn set_mimi_trainable(&mut self, trainable: bool) {
2490        self.set_trainable(trainable);
2491    }
2492}
2493
2494macro_rules! structured_leaf_parameters {
2495    ($type:ty) => {
2496        impl<T: Tensor> MimiModuleParameters<T> for $type {
2497            fn visit_mimi_parameters<'a>(
2498                &'a self,
2499                prefix: &str,
2500                visitor: &mut dyn FnMut(ParameterMetadata, &'a T),
2501            ) {
2502                self.visit_parameters(&mut PrefixVisitor {
2503                    prefix,
2504                    visitor,
2505                    exact: false,
2506                });
2507            }
2508
2509            fn visit_mimi_parameters_mut<'a>(
2510                &'a mut self,
2511                prefix: &str,
2512                visitor: &mut dyn FnMut(ParameterMetadata, &'a mut T),
2513            ) {
2514                self.visit_parameters_mut(&mut PrefixVisitorMut {
2515                    prefix,
2516                    visitor,
2517                    exact: false,
2518                });
2519            }
2520
2521            fn set_mimi_trainable(&mut self, trainable: bool) {
2522                self.set_trainable(trainable);
2523            }
2524        }
2525    };
2526}
2527
2528structured_leaf_parameters!(Linear<T>);
2529structured_leaf_parameters!(LayerNorm<T>);
2530
2531impl<T: Tensor, M: MimiModuleParameters<T>> MimiModuleParameters<T> for Vec<M> {
2532    fn visit_mimi_parameters<'a>(
2533        &'a self,
2534        prefix: &str,
2535        visitor: &mut dyn FnMut(ParameterMetadata, &'a T),
2536    ) {
2537        for (index, module) in self.iter().enumerate() {
2538            module.visit_mimi_parameters(&parameter_name(prefix, &index.to_string()), visitor);
2539        }
2540    }
2541
2542    fn visit_mimi_parameters_mut<'a>(
2543        &'a mut self,
2544        prefix: &str,
2545        visitor: &mut dyn FnMut(ParameterMetadata, &'a mut T),
2546    ) {
2547        for (index, module) in self.iter_mut().enumerate() {
2548            module.visit_mimi_parameters_mut(&parameter_name(prefix, &index.to_string()), visitor);
2549        }
2550    }
2551
2552    fn set_mimi_trainable(&mut self, trainable: bool) {
2553        for module in self {
2554            module.set_mimi_trainable(trainable);
2555        }
2556    }
2557}
2558
2559impl<T: Tensor, M: MimiModuleParameters<T>> MimiModuleParameters<T> for Option<M> {
2560    fn visit_mimi_parameters<'a>(
2561        &'a self,
2562        prefix: &str,
2563        visitor: &mut dyn FnMut(ParameterMetadata, &'a T),
2564    ) {
2565        if let Some(module) = self {
2566            module.visit_mimi_parameters(prefix, visitor);
2567        }
2568    }
2569
2570    fn visit_mimi_parameters_mut<'a>(
2571        &'a mut self,
2572        prefix: &str,
2573        visitor: &mut dyn FnMut(ParameterMetadata, &'a mut T),
2574    ) {
2575        if let Some(module) = self {
2576            module.visit_mimi_parameters_mut(prefix, visitor);
2577        }
2578    }
2579
2580    fn set_mimi_trainable(&mut self, trainable: bool) {
2581        if let Some(module) = self {
2582            module.set_mimi_trainable(trainable);
2583        }
2584    }
2585}
2586
2587macro_rules! module_parameters {
2588    ($module:ident { $($field:ident),+ $(,)? }) => {
2589        impl<T: Tensor> MimiModuleParameters<T> for $module<T> {
2590            fn visit_mimi_parameters<'a>(
2591                &'a self,
2592                prefix: &str,
2593                visitor: &mut dyn FnMut(ParameterMetadata, &'a T),
2594            ) {
2595                $(
2596                    self.$field.visit_mimi_parameters(
2597                        &parameter_name(prefix, stringify!($field)),
2598                        visitor,
2599                    );
2600                )+
2601            }
2602
2603            fn visit_mimi_parameters_mut<'a>(
2604                &'a mut self,
2605                prefix: &str,
2606                visitor: &mut dyn FnMut(ParameterMetadata, &'a mut T),
2607            ) {
2608                $(
2609                    self.$field.visit_mimi_parameters_mut(
2610                        &parameter_name(prefix, stringify!($field)),
2611                        visitor,
2612                    );
2613                )+
2614            }
2615
2616            fn set_mimi_trainable(&mut self, trainable: bool) {
2617                $(self.$field.set_mimi_trainable(trainable);)+
2618            }
2619        }
2620    };
2621}
2622
2623module_parameters!(Mimi {
2624    quantizer,
2625    encoder,
2626    encoder_transformer,
2627    downsample,
2628    upsample,
2629    decoder_transformer,
2630    decoder,
2631});
2632module_parameters!(SeaNetEncoder {
2633    init_conv1d,
2634    layers,
2635    final_conv1d,
2636});
2637module_parameters!(EncoderLayer {
2638    residuals,
2639    downsample,
2640});
2641module_parameters!(MimiTransformer { layers });
2642module_parameters!(MimiTransformerLayer {
2643    norm1,
2644    norm2,
2645    self_attn,
2646    mlp,
2647    layer_scale_1,
2648    layer_scale_2,
2649});
2650module_parameters!(LayerScale { scale });
2651module_parameters!(MimiMlp { linear1, linear2 });
2652module_parameters!(MimiSelfAttention { in_proj, out_proj });
2653module_parameters!(SeaNetDecoder {
2654    init_conv1d,
2655    layers,
2656    final_conv1d,
2657});
2658module_parameters!(DecoderLayer {
2659    upsample,
2660    residuals,
2661});
2662module_parameters!(SeaNetResnetBlock { block });
2663module_parameters!(StreamableConv1d { weight, bias });
2664module_parameters!(StreamableConvTranspose1d { weight, bias });
2665module_parameters!(SplitResidualVectorQuantizer {
2666    rvq_first,
2667    rvq_rest,
2668});
2669module_parameters!(ResidualVectorQuantizer {
2670    input_proj,
2671    output_proj,
2672    vq,
2673});
2674module_parameters!(ResidualVectorQuantization { layers });
2675module_parameters!(VectorQuantization { _codebook });
2676module_parameters!(EuclideanCodebook {
2677    _initialized,
2678    cluster_usage,
2679    embedding_sum,
2680});
2681module_parameters!(Conv1x1NoBias { weight });
2682
2683impl<T: Tensor> Parameterized<T> for Mimi<T> {
2684    fn visit_parameters<'a, V>(&'a self, visitor: &mut V)
2685    where
2686        V: ParameterVisitor<'a, T>,
2687    {
2688        self.visit_mimi_parameters("", &mut |metadata, value| {
2689            visitor.visit(metadata, value);
2690        });
2691    }
2692
2693    fn visit_parameters_mut<'a, V>(&'a mut self, visitor: &mut V)
2694    where
2695        V: ParameterVisitorMut<'a, T>,
2696    {
2697        self.visit_mimi_parameters_mut("", &mut |metadata, value| {
2698            visitor.visit_mut(metadata, value);
2699        });
2700    }
2701
2702    fn set_trainable(&mut self, trainable: bool) {
2703        self.set_mimi_trainable(trainable);
2704    }
2705}
2706
2707fn validate_latent<T: Tensor>(latent: &T) -> Result<(), Error> {
2708    if latent.shape().len() != 3 || latent.dim(1) != 512 {
2709        return Err(Error::InvalidShape(format!(
2710            "Mimi latent frames must have shape [batch, 512, frames], got {:?}",
2711            latent.shape()
2712        )));
2713    }
2714    Ok(())
2715}
2716
2717fn validate_pcm<T: Tensor>(pcm: &T) -> Result<(), Error> {
2718    if pcm.shape().len() != 3 || pcm.dim(1) != 1 {
2719        return Err(Error::InvalidShape(format!(
2720            "Mimi PCM must have shape [batch, 1, samples], got {:?}",
2721            pcm.shape()
2722        )));
2723    }
2724    Ok(())
2725}
2726
2727fn validate_codes<T: Tensor>(codes: &T, max_codebooks: i32) -> Result<(), Error> {
2728    if codes.shape().len() != 3 || codes.dim(1) <= 0 || codes.dim(1) > max_codebooks {
2729        return Err(Error::InvalidShape(format!(
2730            "Mimi codes must have shape [batch, 1..={max_codebooks}, frames], got {:?}",
2731            codes.shape()
2732        )));
2733    }
2734    Ok(())
2735}
2736
2737#[cfg(test)]
2738mod tests {
2739    use std::{
2740        collections::BTreeMap,
2741        sync::{
2742            atomic::{AtomicUsize, Ordering},
2743            Arc,
2744        },
2745    };
2746
2747    use super::{
2748        checkpoint_key_for_parameter, checkpoint_layout_axes, checkpoint_parameter_for_key,
2749        parameter_topology, prepare_catalog, prepare_source, released_checkpoint_requirements,
2750        Config, Mimi, MimiArtifactError, MimiParameterRequirement, RecipeDtype,
2751    };
2752    use eredu_checkpoint::store::{
2753        CheckpointLease, CheckpointSource, TensorMetadata, TensorReadRequest,
2754        TensorSourceProvenance, WeightStoreBackend, WeightStoreDiagnostics,
2755    };
2756    use eredu_checkpoint::{SourceTensorEncoding, StoredDtype};
2757
2758    struct MetadataSource {
2759        tensors: BTreeMap<String, TensorMetadata>,
2760        payload_reads: AtomicUsize,
2761        encoding: SourceTensorEncoding,
2762    }
2763
2764    impl MetadataSource {
2765        fn exact(requirements: &[MimiParameterRequirement]) -> Self {
2766            Self {
2767                tensors: requirements
2768                    .iter()
2769                    .map(|requirement| {
2770                        (
2771                            requirement.checkpoint_key().to_owned(),
2772                            TensorMetadata {
2773                                name: requirement.checkpoint_key().to_owned(),
2774                                logical_shape: requirement.physical_shape().to_vec(),
2775                                physical_shape: requirement.physical_shape().to_vec(),
2776                                stored_dtype: requirement.source_dtype().clone(),
2777                                encoded_byte_len: requirement.source_bytes(),
2778                                backing_shard: None,
2779                            },
2780                        )
2781                    })
2782                    .collect(),
2783                payload_reads: AtomicUsize::new(0),
2784                encoding: SourceTensorEncoding::Safetensors(StoredDtype::F32),
2785            }
2786        }
2787    }
2788
2789    impl CheckpointSource for MetadataSource {
2790        fn source_keys(&self) -> Vec<String> {
2791            self.tensors.keys().cloned().collect()
2792        }
2793
2794        fn source_metadata(
2795            &self,
2796            key: &str,
2797        ) -> Result<TensorMetadata, eredu_checkpoint::store::StoreError> {
2798            self.tensors.get(key).cloned().ok_or_else(|| {
2799                eredu_checkpoint::store::StoreError::UnknownTensor { key: key.into() }
2800            })
2801        }
2802
2803        fn acquire_lease(
2804            &self,
2805            _: TensorReadRequest,
2806        ) -> Result<CheckpointLease, eredu_checkpoint::store::StoreError> {
2807            self.payload_reads.fetch_add(1, Ordering::Relaxed);
2808            Err(eredu_checkpoint::store::StoreError::Internal(
2809                "metadata-only test source cannot read payloads".into(),
2810            ))
2811        }
2812
2813        fn source_diagnostics(
2814            &self,
2815        ) -> Result<WeightStoreDiagnostics, eredu_checkpoint::store::StoreError> {
2816            Ok(WeightStoreDiagnostics {
2817                backend: WeightStoreBackend::Memory,
2818                cache_hits: 0,
2819                cache_misses: 0,
2820                evictions: 0,
2821                currently_cached_shards: 0,
2822                touched_shard_paths: Vec::new(),
2823                payload_shard_paths: Vec::new(),
2824                physical_reads: self.payload_reads.load(Ordering::Relaxed) as u64,
2825                physical_read_bytes: 0,
2826                coalesced_group_hits: 0,
2827            })
2828        }
2829
2830        fn source_provenance(
2831            &self,
2832            key: &str,
2833        ) -> Result<TensorSourceProvenance, eredu_checkpoint::store::StoreError> {
2834            let metadata = self.source_metadata(key)?;
2835            Ok(TensorSourceProvenance {
2836                catalog_key: key.to_owned(),
2837                physical_tensor: key.to_owned(),
2838                output: key.to_owned(),
2839                backing_shard: metadata.backing_shard,
2840                source_encoding: self.encoding.clone(),
2841            })
2842        }
2843    }
2844
2845    fn exact_source() -> (Arc<MetadataSource>, Vec<MimiParameterRequirement>) {
2846        let (_, requirements, _) = prepare_catalog(&Config::v0_1(Some(8))).unwrap();
2847        (Arc::new(MetadataSource::exact(&requirements)), requirements)
2848    }
2849
2850    #[test]
2851    fn checkpoint_quantizer_keys_keep_the_model_root() {
2852        let key = "quantizer.rvq_first.vq.layers.0._codebook.embedding_sum";
2853        assert_eq!(checkpoint_parameter_for_key(key).as_deref(), Some(key));
2854        assert_eq!(checkpoint_key_for_parameter(key).as_deref(), Some(key));
2855        assert_eq!(checkpoint_layout_axes(key), None);
2856    }
2857
2858    #[test]
2859    fn checkpoint_plan_declares_canonical_convolution_layouts() {
2860        assert_eq!(
2861            checkpoint_layout_axes("encoder.init_conv1d.weight"),
2862            Some([0, 2, 1])
2863        );
2864        assert_eq!(
2865            checkpoint_layout_axes("decoder.layers.0.upsample.weight"),
2866            Some([1, 2, 0])
2867        );
2868        assert_eq!(checkpoint_layout_axes("upsample.weight"), Some([0, 2, 1]));
2869        assert!(checkpoint_parameter_for_key("optimizer.state").is_none());
2870    }
2871
2872    #[test]
2873    fn parameter_names_are_unique_and_cover_checkpoint_mapping() {
2874        let active = parameter_topology(Config::v0_1(Some(8))).unwrap();
2875        let full = parameter_topology(Config::v0_1(Some(32))).unwrap();
2876        assert_eq!(active.len(), 246);
2877        assert_eq!(full.len(), 318);
2878        for model_name in full.keys() {
2879            let checkpoint_name = checkpoint_key_for_parameter(model_name)
2880                .unwrap_or_else(|| panic!("model parameter was not mapped: {model_name}"));
2881            assert_eq!(
2882                checkpoint_parameter_for_key(&checkpoint_name).as_deref(),
2883                Some(model_name.as_str()),
2884                "checkpoint mapping did not round-trip"
2885            );
2886        }
2887    }
2888
2889    #[test]
2890    fn exact_catalog_preparation_validates_total_topology_without_payload_reads() {
2891        for active in 1..=32 {
2892            let (_, requirements) = exact_source();
2893            let source = Arc::new(MetadataSource::exact(&requirements));
2894            let prepared = prepare_source(source.clone(), Config::v0_1(Some(active))).unwrap();
2895            assert_eq!(prepared.requirements().len(), 318);
2896            assert_eq!(prepared.bindings().len(), 3 * active as usize + 222);
2897            assert_eq!(
2898                prepared
2899                    .requirements()
2900                    .iter()
2901                    .filter(|requirement| requirement.is_active())
2902                    .count(),
2903                3 * active as usize + 222
2904            );
2905            assert!(prepared.requirements().iter().all(|requirement| {
2906                requirement.source_dtype() == &StoredDtype::F32
2907                    && requirement.output_dtype() == &RecipeDtype::F32
2908                    && requirement.source_encoding()
2909                        == &SourceTensorEncoding::Safetensors(StoredDtype::F32)
2910            }));
2911            assert_eq!(
2912                prepared
2913                    .requirements()
2914                    .iter()
2915                    .filter(|requirement| matches!(
2916                        requirement.recipe(),
2917                        eredu_checkpoint::recipe::DerivedWeightRecipe::Transpose { .. }
2918                    ))
2919                    .count(),
2920                30
2921            );
2922            assert_eq!(
2923                prepared
2924                    .requirements()
2925                    .iter()
2926                    .filter(|requirement| requirement.recipe().source_keys().len() == 1)
2927                    .count(),
2928                318
2929            );
2930            assert_eq!(source.payload_reads.load(Ordering::Relaxed), 0);
2931        }
2932    }
2933
2934    #[test]
2935    fn corrupt_catalogs_fail_before_payload_reads() {
2936        let (source, requirements) = exact_source();
2937        let missing_key = requirements[0].checkpoint_key().to_owned();
2938        let mut missing = MetadataSource::exact(&requirements);
2939        missing.tensors.remove(&missing_key);
2940        let missing = Arc::new(missing);
2941        assert!(matches!(
2942            prepare_source(missing.clone(), Config::v0_1(Some(8))),
2943            Err(MimiArtifactError::Catalog(_))
2944        ));
2945        assert_eq!(missing.payload_reads.load(Ordering::Relaxed), 0);
2946
2947        let mut extra = MetadataSource::exact(&requirements);
2948        extra.tensors.insert(
2949            "optimizer.state".into(),
2950            TensorMetadata {
2951                name: "optimizer.state".into(),
2952                logical_shape: vec![1],
2953                physical_shape: vec![1],
2954                stored_dtype: StoredDtype::F32,
2955                encoded_byte_len: 4,
2956                backing_shard: None,
2957            },
2958        );
2959        let extra = Arc::new(extra);
2960        assert!(matches!(
2961            prepare_source(extra.clone(), Config::v0_1(Some(8))),
2962            Err(MimiArtifactError::Catalog(_))
2963        ));
2964        assert_eq!(extra.payload_reads.load(Ordering::Relaxed), 0);
2965
2966        let corrupt = requirements
2967            .iter()
2968            .find(|requirement| requirement.physical_shape().len() == 3)
2969            .unwrap();
2970        let mut wrong_shape = MetadataSource::exact(&requirements);
2971        wrong_shape
2972            .tensors
2973            .get_mut(corrupt.checkpoint_key())
2974            .unwrap()
2975            .logical_shape[0] += 1;
2976        let wrong_shape = Arc::new(wrong_shape);
2977        assert!(matches!(
2978            prepare_source(wrong_shape.clone(), Config::v0_1(Some(8))),
2979            Err(MimiArtifactError::Catalog(_))
2980        ));
2981        assert_eq!(wrong_shape.payload_reads.load(Ordering::Relaxed), 0);
2982
2983        let mut wrong_dtype = MetadataSource::exact(&requirements);
2984        wrong_dtype
2985            .tensors
2986            .get_mut(corrupt.checkpoint_key())
2987            .unwrap()
2988            .stored_dtype = StoredDtype::F16;
2989        let wrong_dtype = Arc::new(wrong_dtype);
2990        assert!(matches!(
2991            prepare_source(wrong_dtype.clone(), Config::v0_1(Some(8))),
2992            Err(MimiArtifactError::Catalog(_))
2993        ));
2994        assert_eq!(wrong_dtype.payload_reads.load(Ordering::Relaxed), 0);
2995
2996        let mut wrong_bytes = MetadataSource::exact(&requirements);
2997        wrong_bytes
2998            .tensors
2999            .get_mut(corrupt.checkpoint_key())
3000            .unwrap()
3001            .encoded_byte_len -= 4;
3002        let wrong_bytes = Arc::new(wrong_bytes);
3003        assert!(matches!(
3004            prepare_source(wrong_bytes.clone(), Config::v0_1(Some(8))),
3005            Err(MimiArtifactError::Topology(_))
3006        ));
3007        assert_eq!(wrong_bytes.payload_reads.load(Ordering::Relaxed), 0);
3008
3009        let mut wrong_encoding = MetadataSource::exact(&requirements);
3010        wrong_encoding.encoding = SourceTensorEncoding::RecipeOutput(StoredDtype::F32);
3011        let wrong_encoding = Arc::new(wrong_encoding);
3012        assert!(matches!(
3013            prepare_source(wrong_encoding.clone(), Config::v0_1(Some(8))),
3014            Err(MimiArtifactError::Topology(_))
3015        ));
3016        assert_eq!(wrong_encoding.payload_reads.load(Ordering::Relaxed), 0);
3017        assert_eq!(source.payload_reads.load(Ordering::Relaxed), 0);
3018    }
3019
3020    #[test]
3021    fn v0_1_config_defaults_to_moshi_active_codebooks() {
3022        let cfg = Config::v0_1(None);
3023        assert_eq!(cfg.sample_rate, 24_000.0);
3024        assert_eq!(cfg.frame_rate, 12.5);
3025        assert_eq!(cfg.num_codebooks, 16);
3026        assert_eq!(cfg.total_codebooks, 32);
3027        assert_eq!(cfg.bins, 2_048);
3028    }
3029
3030    #[test]
3031    fn unsupported_or_non_finite_profiles_fail_before_tensor_construction() {
3032        let mut invalid = Config::v0_1(Some(8));
3033        invalid.sample_rate = f64::NAN;
3034        assert!(Mimi::<super::PlanningTensor>::new(invalid, &()).is_err());
3035
3036        let mut unsupported = Config::v0_1(Some(8));
3037        unsupported.total_codebooks = 16;
3038        assert!(Mimi::<super::PlanningTensor>::new(unsupported, &()).is_err());
3039
3040        for active in [-1, 0, 33] {
3041            assert!(released_checkpoint_requirements(&Config::v0_1(Some(active))).is_err());
3042        }
3043    }
3044}