1use 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#[derive(Debug, Clone, Copy, Eq, PartialEq)]
96pub enum ResampleMethod {
97 Conv,
99}
100
101#[derive(Debug, Clone)]
103pub struct Config {
104 pub channels: i32,
106 pub sample_rate: f64,
108 pub frame_rate: f64,
110 pub renormalize: bool,
112 pub resample_method: ResampleMethod,
114 pub num_codebooks: i32,
116 pub total_codebooks: i32,
118 pub bins: i32,
120 pub quantizer_dim: i32,
122 pub latent_dim: i32,
124}
125
126impl Config {
127 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#[derive(Debug, Clone)]
187pub struct Mimi<T: Tensor> {
188 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 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 pub fn mimi_config(&self) -> &Config {
233 &self.config
234 }
235
236 pub fn encode_latent(&mut self, latent: &T, context: &T::Context) -> Result<T, Error> {
238 self.quantizer.encode(latent, context)
239 }
240
241 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 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 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 pub fn decode_latent(&mut self, codes: &T, context: &T::Context) -> Result<T, Error> {
279 self.quantizer.decode(codes, context)
280 }
281
282 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 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 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#[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 pub fn logical_name(&self) -> &str {
337 &self.logical_name
338 }
339
340 pub fn checkpoint_key(&self) -> &str {
342 &self.checkpoint_key
343 }
344
345 pub fn physical_shape(&self) -> &[usize] {
347 &self.physical_shape
348 }
349
350 pub fn logical_shape(&self) -> &[usize] {
352 &self.logical_shape
353 }
354
355 pub const fn source_dtype(&self) -> &StoredDtype {
357 &self.source_dtype
358 }
359
360 pub const fn output_dtype(&self) -> &RecipeDtype {
362 &self.output_dtype
363 }
364
365 pub const fn source_encoding(&self) -> &SourceTensorEncoding {
367 &self.source_encoding
368 }
369
370 pub const fn source_bytes(&self) -> u64 {
372 self.source_bytes
373 }
374
375 pub const fn output_bytes(&self) -> u64 {
377 self.output_bytes
378 }
379
380 pub const fn recipe(&self) -> &DerivedWeightRecipe {
382 &self.recipe
383 }
384
385 pub const fn is_active(&self) -> bool {
387 self.active
388 }
389}
390
391pub 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 pub const fn config(&self) -> &Config {
413 &self.config
414 }
415
416 pub fn requirements(&self) -> &[MimiParameterRequirement] {
418 &self.requirements
419 }
420
421 pub fn bindings(&self) -> &[WeightBinding] {
423 &self.bindings
424 }
425
426 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
449pub 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 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#[derive(Debug, thiserror::Error)]
498pub enum MimiArtifactError {
499 #[error(transparent)]
501 Codec(#[from] Error),
502 #[error(transparent)]
504 Safetensors(#[from] SafetensorsShardError),
505 #[error(transparent)]
507 Plan(#[from] CheckpointPlanError),
508 #[error("Mimi SafeTensors catalog is not exact: {0:?}")]
510 Catalog(CheckpointValidation),
511 #[error(transparent)]
513 Store(#[from] StoreError),
514 #[error(transparent)]
516 Recipe(#[from] RecipeError),
517 #[error(transparent)]
519 Binding(#[from] ResidencyDeclarationError),
520 #[error("invalid Mimi parameter topology: {0}")]
522 Topology(String),
523}
524
525#[derive(Debug, thiserror::Error)]
527pub enum MimiConstructionError<E>
528where
529 E: std::error::Error + Send + Sync + 'static,
530{
531 #[error(transparent)]
533 Codec(#[from] Error),
534 #[error(transparent)]
536 Parameters(#[from] ParameterOrchestrationError<E>),
537}
538
539pub 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
556pub fn released_checkpoint_requirements(
559 config: &Config,
560) -> Result<Vec<MimiParameterRequirement>, MimiArtifactError> {
561 prepare_catalog(config).map(|(_, requirements, _)| requirements)
562}
563
564pub 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
576pub 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}{}", ¶meter[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#[derive(Debug, Clone)]
2135pub struct SplitResidualVectorQuantizer<T: Tensor> {
2136 pub rvq_first: ResidualVectorQuantizer<T>,
2138 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 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 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#[derive(Debug, Clone)]
2194pub struct ResidualVectorQuantizer<T: Tensor> {
2195 pub input_proj: Conv1x1NoBias<T>,
2197 pub output_proj: Conv1x1NoBias<T>,
2199 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#[derive(Debug, Clone)]
2231pub struct ResidualVectorQuantization<T: Tensor> {
2232 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#[derive(Debug, Clone)]
2286pub struct VectorQuantization<T: Tensor> {
2287 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#[derive(Debug, Clone)]
2313pub struct EuclideanCodebook<T: Tensor> {
2314 pub _initialized: Parameter<T>,
2316 pub cluster_usage: Parameter<T>,
2318 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#[derive(Debug, Clone)]
2381pub struct Conv1x1NoBias<T: Tensor> {
2382 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(¶meter_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(¶meter_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 ¶meter_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 ¶meter_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}