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