Skip to main content

cortiq_engine/
audiovae.rs

1//! MiniMax-H3's audio VAE decoder: BigVGAN at 32 kHz, stereo.
2//!
3//! 32 latent channels at 40 frames a second become a waveform at 800
4//! samples a frame, through seven transposed-convolution stages and
5//! their AMP residual blocks. The two stereo channels are two
6//! independent mono passes, which is how the reference batches them.
7//!
8//! The activations are the interesting part and the easy thing to get
9//! subtly wrong: every nonlinearity is wrapped in a 2× kaiser-sinc
10//! upsample, the pointwise SnakeBeta, and a 2× lowpass back down. The
11//! filter is designed here rather than shipped, from the same
12//! `kaiser_sinc_filter1d(cutoff, half_width, 12)` the reference calls,
13//! so it cannot drift out of step with a checkpoint that does not
14//! contain it.
15
16use crate::pool::Pool;
17use cortiq_core::CmfModel;
18use std::sync::Arc;
19
20/// Both resampling filters in the alias-free activation are designed at
21/// this kernel length.
22const FILTER_LEN: usize = 12;
23
24struct Conv1d {
25    w: Vec<f32>, // [out, in, k]
26    b: Option<Vec<f32>>,
27    out_ch: usize,
28    in_ch: usize,
29    k: usize,
30    pad: usize,
31    dilation: usize,
32}
33
34impl Conv1d {
35    fn load(model: &Arc<CmfModel>, name: &str, pad: usize, dilation: usize) -> Result<Self, String> {
36        let e = model
37            .tensor(&format!("{name}.weight"))
38            .ok_or_else(|| format!("missing {name}.weight"))?;
39        let w = crate::dit::cmf_f32(model, &format!("{name}.weight"))?;
40        let b = crate::dit::cmf_f32(model, &format!("{name}.bias")).ok();
41        Ok(Self {
42            out_ch: e.shape[0],
43            in_ch: e.shape[1],
44            k: e.shape[2],
45            w,
46            b,
47            pad,
48            dilation,
49        })
50    }
51
52    /// `x` is `[in_ch, n]`; the result is `[out_ch, n]` for the
53    /// paddings used here (all `same`).
54    fn apply(&self, x: &[f32], n: usize, pool: Option<&Pool>) -> Vec<f32> {
55        let out_n = (n + 2 * self.pad).saturating_sub(self.dilation * (self.k - 1));
56        let mut out = vec![0f32; self.out_ch * out_n];
57        let ptr = SendPtr(out.as_mut_ptr());
58        let work = |lo: usize, hi: usize| {
59            for o in lo..hi {
60                // SAFETY: workers own disjoint output channels.
61                let dst = unsafe { ptr.row(o * out_n, out_n) };
62                let bias = self.b.as_ref().map_or(0.0, |b| b[o]);
63                dst.fill(bias);
64                for i in 0..self.in_ch {
65                    let ker = &self.w[(o * self.in_ch + i) * self.k..(o * self.in_ch + i + 1) * self.k];
66                    let src = &x[i * n..(i + 1) * n];
67                    for (t, d) in dst.iter_mut().enumerate() {
68                        let mut acc = 0f32;
69                        for (j, &kv) in ker.iter().enumerate() {
70                            let p = (t + j * self.dilation) as isize - self.pad as isize;
71                            if p >= 0 && (p as usize) < n {
72                                acc += kv * src[p as usize];
73                            }
74                        }
75                        *d += acc;
76                    }
77                }
78            }
79        };
80        match pool {
81            Some(p) => p.run_rows(self.out_ch, &work),
82            None => work(0, self.out_ch),
83        }
84        out
85    }
86}
87
88struct ConvT1d {
89    w: Vec<f32>, // [in, out, k]
90    b: Vec<f32>,
91    in_ch: usize,
92    out_ch: usize,
93    k: usize,
94    stride: usize,
95    pad: usize,
96}
97
98impl ConvT1d {
99    fn load(model: &Arc<CmfModel>, name: &str, stride: usize) -> Result<Self, String> {
100        let e = model
101            .tensor(&format!("{name}.weight"))
102            .ok_or_else(|| format!("missing {name}.weight"))?;
103        let (in_ch, out_ch, k) = (e.shape[0], e.shape[1], e.shape[2]);
104        Ok(Self {
105            w: crate::dit::cmf_f32(model, &format!("{name}.weight"))?,
106            b: crate::dit::cmf_f32(model, &format!("{name}.bias"))?,
107            in_ch,
108            out_ch,
109            k,
110            stride,
111            pad: (k - stride) / 2,
112        })
113    }
114
115    fn apply(&self, x: &[f32], n: usize, pool: Option<&Pool>) -> Vec<f32> {
116        let full = (n - 1) * self.stride + self.k;
117        let out_n = full - 2 * self.pad;
118        let mut out = vec![0f32; self.out_ch * out_n];
119        let ptr = SendPtr(out.as_mut_ptr());
120        let work = |lo: usize, hi: usize| {
121            for o in lo..hi {
122                // SAFETY: workers own disjoint output channels.
123                let dst = unsafe { ptr.row(o * out_n, out_n) };
124                dst.fill(self.b[o]);
125                for i in 0..self.in_ch {
126                    let ker = &self.w[(i * self.out_ch + o) * self.k..(i * self.out_ch + o + 1) * self.k];
127                    let src = &x[i * n..(i + 1) * n];
128                    for (t, &sv) in src.iter().enumerate() {
129                        if sv == 0.0 {
130                            continue;
131                        }
132                        let base = t * self.stride;
133                        for (j, &kv) in ker.iter().enumerate() {
134                            let p = base + j;
135                            if p >= self.pad && p - self.pad < out_n {
136                                dst[p - self.pad] += sv * kv;
137                            }
138                        }
139                    }
140                }
141            }
142        };
143        match pool {
144            Some(p) => p.run_rows(self.out_ch, &work),
145            None => work(0, self.out_ch),
146        }
147        out
148    }
149}
150
151/// `x + sin²(α·x)/β`, with α and β stored in log scale.
152struct SnakeBeta {
153    alpha: Vec<f32>,
154    beta: Vec<f32>,
155}
156
157impl SnakeBeta {
158    fn load(model: &Arc<CmfModel>, name: &str) -> Result<Self, String> {
159        Ok(Self {
160            alpha: crate::dit::cmf_f32(model, &format!("{name}.alpha"))?
161                .iter()
162                .map(|v| v.exp())
163                .collect(),
164            beta: crate::dit::cmf_f32(model, &format!("{name}.beta"))?
165                .iter()
166                .map(|v| v.exp())
167                .collect(),
168        })
169    }
170
171    fn apply(&self, x: &mut [f32], n: usize) {
172        for (c, row) in x.chunks_exact_mut(n).enumerate() {
173            let (a, b) = (self.alpha[c], 1.0 / (self.beta[c] + 1e-9));
174            for v in row.iter_mut() {
175                let s = (a * *v).sin();
176                *v += s * s * b;
177            }
178        }
179    }
180}
181
182fn bessel_i0(x: f64) -> f64 {
183    // Series; the argument here is ~4.7, where a dozen terms is exact
184    // to double precision.
185    let mut sum = 1.0;
186    let mut term = 1.0;
187    for k in 1..40 {
188        term *= (x / (2.0 * k as f64)).powi(2);
189        sum += term;
190        if term < 1e-18 * sum {
191            break;
192        }
193    }
194    sum
195}
196
197fn sinc(x: f64) -> f64 {
198    if x == 0.0 {
199        1.0
200    } else {
201        (std::f64::consts::PI * x).sin() / (std::f64::consts::PI * x)
202    }
203}
204
205/// The reference's `kaiser_sinc_filter1d`, normalized to unit sum.
206fn kaiser_sinc(cutoff: f64, half_width: f64, k: usize) -> Vec<f32> {
207    let half = k / 2;
208    let delta_f = 4.0 * half_width;
209    let a = 2.285 * (half as f64 - 1.0) * std::f64::consts::PI * delta_f + 7.95;
210    let beta = if a > 50.0 {
211        0.1102 * (a - 8.7)
212    } else if a >= 21.0 {
213        0.5842 * (a - 21.0).powf(0.4) + 0.078_86 * (a - 21.0)
214    } else {
215        0.0
216    };
217    let denom = bessel_i0(beta);
218    let n = k as f64 - 1.0;
219    let mut f: Vec<f64> = (0..k)
220        .map(|i| {
221            let r = (2.0 * i as f64 / n) - 1.0;
222            let win = bessel_i0(beta * (1.0 - r * r).max(0.0).sqrt()) / denom;
223            // even length: sample points sit on half-integers
224            let t = -(half as f64) + i as f64 + 0.5;
225            2.0 * cutoff * win * sinc(2.0 * cutoff * t)
226        })
227        .collect();
228    let s: f64 = f.iter().sum();
229    for v in f.iter_mut() {
230        *v /= s;
231    }
232    f.into_iter().map(|v| v as f32).collect()
233}
234
235/// Replicate-pad, then a per-channel FIR.
236fn fir_pad(x: &[f32], ch: usize, n: usize, f: &[f32], pad_l: usize, pad_r: usize, stride: usize) -> (Vec<f32>, usize) {
237    let padded = n + pad_l + pad_r;
238    let out_n = (padded - f.len()) / stride + 1;
239    let mut out = vec![0f32; ch * out_n];
240    let mut buf = vec![0f32; padded];
241    for c in 0..ch {
242        let src = &x[c * n..(c + 1) * n];
243        for (i, b) in buf.iter_mut().enumerate() {
244            let p = i as isize - pad_l as isize;
245            *b = src[p.clamp(0, n as isize - 1) as usize];
246        }
247        for t in 0..out_n {
248            let mut acc = 0f32;
249            for (j, &kv) in f.iter().enumerate() {
250                acc += kv * buf[t * stride + j];
251            }
252            out[c * out_n + t] = acc;
253        }
254    }
255    (out, out_n)
256}
257
258/// Upsample ×2, apply, downsample ×2 — the anti-aliased activation.
259struct Activation1d {
260    act: SnakeBeta,
261    up: Vec<f32>,
262    down: Vec<f32>,
263}
264
265impl Activation1d {
266    /// `name` is the Activation1d module, not its `.act`. The release
267    /// ships both resampling filters as buffers — 254 of them — so read
268    /// them rather than re-designing them, and keep `kaiser_sinc` as
269    /// the fallback for a checkpoint that drops them.
270    fn load(model: &Arc<CmfModel>, name: &str) -> Result<Self, String> {
271        let designed = || kaiser_sinc(0.25, 0.3, FILTER_LEN);
272        Ok(Self {
273            act: SnakeBeta::load(model, &format!("{name}.act"))?,
274            up: crate::dit::cmf_f32(model, &format!("{name}.upsample.filter"))
275                .unwrap_or_else(|_| designed()),
276            down: crate::dit::cmf_f32(model, &format!("{name}.downsample.lowpass.filter"))
277                .unwrap_or_else(|_| designed()),
278        })
279    }
280
281    fn apply(&self, x: &[f32], ch: usize, n: usize) -> (Vec<f32>, usize) {
282        // conv_transpose1d(pad(x, 5, 5), filter, stride 2) · 2, then the
283        // 15-sample margins the reference trims off each end.
284        let pad = FILTER_LEN / 2 - 1;
285        let pad_l = pad * 2 + (FILTER_LEN - 2) / 2;
286        let pad_r = pad * 2 + (FILTER_LEN - 2 + 1) / 2;
287        let pn = n + 2 * pad;
288        let full = (pn - 1) * 2 + FILTER_LEN;
289        let mut up = vec![0f32; ch * full];
290        for c in 0..ch {
291            let src = &x[c * n..(c + 1) * n];
292            let dst = &mut up[c * full..(c + 1) * full];
293            for i in 0..pn {
294                let p = i as isize - pad as isize;
295                let v = src[p.clamp(0, n as isize - 1) as usize] * 2.0;
296                if v == 0.0 {
297                    continue;
298                }
299                for (j, &kv) in self.up.iter().enumerate() {
300                    dst[i * 2 + j] += v * kv;
301                }
302            }
303        }
304        let keep = full - pad_l - pad_r;
305        let mut mid = vec![0f32; ch * keep];
306        for c in 0..ch {
307            mid[c * keep..(c + 1) * keep]
308                .copy_from_slice(&up[c * full + pad_l..c * full + pad_l + keep]);
309        }
310        self.act.apply(&mut mid, keep);
311        // LowPassFilter1d at stride 2: even kernel pads 5 left, 6 right.
312        fir_pad(&mid, ch, keep, &self.down, FILTER_LEN / 2 - 1, FILTER_LEN / 2, 2)
313    }
314}
315
316struct AmpBlock {
317    convs1: Vec<Conv1d>,
318    convs2: Vec<Conv1d>,
319    acts: Vec<Activation1d>,
320}
321
322pub struct AudioVae {
323    dec_in: Conv1d,
324    conv_pre: Conv1d,
325    ups: Vec<ConvT1d>,
326    resblocks: Vec<AmpBlock>,
327    act_post: Activation1d,
328    conv_post: Conv1d,
329    latents_mean: Vec<f32>,
330    latents_std: Vec<f32>,
331    pool: Option<Arc<Pool>>,
332    n_kernels: usize,
333    pub sample_rate: usize,
334}
335
336fn get_padding(k: usize, d: usize) -> usize {
337    (k * d - d) / 2
338}
339
340impl AudioVae {
341    pub fn from_cmf(model: &Arc<CmfModel>) -> Result<Self, String> {
342        let cfg: serde_json::Value = serde_json::from_slice(
343            model.tensor_bytes("avae.config_json").map_err(|e| e.to_string())?,
344        )
345        .map_err(|e| format!("avae.config_json: {e}"))?;
346        let rates: Vec<usize> = cfg["upsample_rates"]
347            .as_array()
348            .ok_or("upsample_rates")?
349            .iter()
350            .map(|v| v.as_u64().unwrap_or(1) as usize)
351            .collect();
352        let rk: Vec<usize> = cfg["resblock_kernel_sizes"]
353            .as_array()
354            .ok_or("resblock_kernel_sizes")?
355            .iter()
356            .map(|v| v.as_u64().unwrap_or(3) as usize)
357            .collect();
358        let rd: Vec<Vec<usize>> = cfg["resblock_dilation_sizes"]
359            .as_array()
360            .ok_or("resblock_dilation_sizes")?
361            .iter()
362            .map(|a| {
363                a.as_array()
364                    .unwrap()
365                    .iter()
366                    .map(|v| v.as_u64().unwrap_or(1) as usize)
367                    .collect()
368            })
369            .collect();
370
371        let mut ups = Vec::new();
372        for (i, &u) in rates.iter().enumerate() {
373            ups.push(ConvT1d::load(model, &format!("avae.decoder.ups.{i}.0"), u)?);
374        }
375        let mut resblocks = Vec::new();
376        for i in 0..rates.len() {
377            for (j, (&k, d)) in rk.iter().zip(&rd).enumerate() {
378                let p = format!("avae.decoder.resblocks.{}", i * rk.len() + j);
379                let convs1 = (0..d.len())
380                    .map(|q| Conv1d::load(model, &format!("{p}.convs1.{q}"), get_padding(k, d[q]), d[q]))
381                    .collect::<Result<Vec<_>, _>>()?;
382                let convs2 = (0..d.len())
383                    .map(|q| Conv1d::load(model, &format!("{p}.convs2.{q}"), get_padding(k, 1), 1))
384                    .collect::<Result<Vec<_>, _>>()?;
385                let acts = (0..convs1.len() + convs2.len())
386                    .map(|q| Activation1d::load(model, &format!("{p}.activations.{q}")))
387                    .collect::<Result<Vec<_>, _>>()?;
388                resblocks.push(AmpBlock { convs1, convs2, acts });
389            }
390        }
391        Ok(Self {
392            dec_in: Conv1d::load(model, "avae.dec_in_proj", 0, 1)?,
393            conv_pre: Conv1d::load(model, "avae.decoder.conv_pre", 3, 1)?,
394            ups,
395            resblocks,
396            act_post: Activation1d::load(model, "avae.decoder.activation_post")?,
397            conv_post: Conv1d::load(model, "avae.decoder.conv_post", 3, 1)?,
398            latents_mean: crate::dit::cmf_f32(model, "avae.latents_mean")?,
399            latents_std: crate::dit::cmf_f32(model, "avae.latents_std")?,
400            pool: Pool::from_env(),
401            n_kernels: rk.len(),
402            sample_rate: cfg["sample_rate"].as_u64().unwrap_or(32000) as usize,
403        })
404    }
405
406    /// Normalized latents `[C, 2, T]` → stereo `[2, L]` in [-1, 1].
407    pub fn decode(&self, z: &[f32], c: usize, t: usize) -> (Vec<f32>, usize) {
408        let pool = self.pool.as_deref();
409        let mut chans: Vec<Vec<f32>> = Vec::with_capacity(2);
410        for ch in 0..2 {
411            let mut lat = vec![0f32; c * t];
412            for ci in 0..c {
413                let (m, s) = (self.latents_mean[ci], self.latents_std[ci]);
414                for ti in 0..t {
415                    lat[ci * t + ti] = z[(ci * 2 + ch) * t + ti] * s + m;
416                }
417            }
418            // `CMF_AVAE_PROF=1`: per-stage rms, to diff against the
419            // reference stage by stage rather than at the waveform.
420            let prof = std::env::var_os("CMF_AVAE_PROF").is_some();
421            let rms = |x: &[f32]| (x.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>()
422                / x.len() as f64)
423                .sqrt();
424            let mut x = self.dec_in.apply(&lat, t, pool);
425            let mut n = t;
426            if prof {
427                eprintln!("ch{ch} dec_in rms {:.6e} n {n}", rms(&x));
428            }
429            x = self.conv_pre.apply(&x, n, pool);
430            if prof {
431                eprintln!("ch{ch} conv_pre rms {:.6e} n {n}", rms(&x));
432            }
433            for i in 0..self.ups.len() {
434                let up = &self.ups[i];
435                x = up.apply(&x, n, pool);
436                n = (n - 1) * up.stride + up.k - 2 * up.pad;
437                let ch_n = up.out_ch;
438                let mut acc = vec![0f32; ch_n * n];
439                for j in 0..self.n_kernels {
440                    let r = self.resblocks[i * self.n_kernels + j].apply(&x, ch_n, n, pool);
441                    for (a, b) in acc.iter_mut().zip(&r) {
442                        *a += b;
443                    }
444                }
445                let inv = 1.0 / self.n_kernels as f32;
446                for v in acc.iter_mut() {
447                    *v *= inv;
448                }
449                x = acc;
450                if prof {
451                    eprintln!("ch{ch} up{i} rms {:.6e} ch {ch_n} n {n}", rms(&x));
452                }
453            }
454            let last_ch = self.ups[self.ups.len() - 1].out_ch;
455            let (mut y, yn) = self.act_post.apply(&x, last_ch, n);
456            y = self.conv_post.apply(&y, yn, pool);
457            for v in y.iter_mut() {
458                *v = v.clamp(-1.0, 1.0);
459            }
460            chans.push(y);
461            n = yn;
462            let _ = n;
463        }
464        let len = chans[0].len().min(chans[1].len());
465        let mut out = vec![0f32; 2 * len];
466        for (ch, c) in chans.iter().enumerate() {
467            out[ch * len..(ch + 1) * len].copy_from_slice(&c[..len]);
468        }
469        (out, len)
470    }
471}
472
473impl AmpBlock {
474    fn apply(&self, x: &[f32], ch: usize, n: usize, pool: Option<&Pool>) -> Vec<f32> {
475        let mut cur = x.to_vec();
476        for i in 0..self.convs1.len() {
477            let (a1, a2) = (&self.acts[i * 2], &self.acts[i * 2 + 1]);
478            let (xt, tn) = a1.apply(&cur, ch, n);
479            let xt = self.convs1[i].apply(&xt, tn, pool);
480            let (xt, tn2) = a2.apply(&xt, ch, tn);
481            let xt = self.convs2[i].apply(&xt, tn2, pool);
482            for (a, b) in cur.iter_mut().zip(&xt) {
483                *a += b;
484            }
485        }
486        cur
487    }
488}
489
490/// Test hook: the designed 12-tap resampling filter.
491#[doc(hidden)]
492pub fn kaiser_sinc_for_test() -> Vec<f32> {
493    kaiser_sinc(0.25, 0.3, FILTER_LEN)
494}
495
496struct SendPtr(*mut f32);
497unsafe impl Send for SendPtr {}
498unsafe impl Sync for SendPtr {}
499impl SendPtr {
500    /// SAFETY: caller guarantees disjoint `[off, off+len)` per worker.
501    #[allow(clippy::mut_from_ref)]
502    unsafe fn row(&self, off: usize, len: usize) -> &mut [f32] {
503        unsafe { std::slice::from_raw_parts_mut(self.0.add(off), len) }
504    }
505}
506
507#[cfg(test)]
508mod tests {
509    use super::*;
510
511    #[test]
512    fn the_resampling_filter_is_the_references() {
513        let f = kaiser_sinc(0.25, 0.3, FILTER_LEN);
514        assert_eq!(f.len(), FILTER_LEN);
515        // Unit sum: without it a constant input leaks amplitude, which
516        // is the whole reason the reference normalizes.
517        assert!((f.iter().sum::<f32>() - 1.0).abs() < 1e-6);
518        // Symmetric about the centre, and its peak is at the centre.
519        for i in 0..FILTER_LEN / 2 {
520            assert!((f[i] - f[FILTER_LEN - 1 - i]).abs() < 1e-6, "asymmetric at {i}");
521        }
522        let peak = f.iter().cloned().fold(f32::MIN, f32::max);
523        assert!((f[5] - peak).abs() < 1e-6);
524
525    }
526
527    #[test]
528    fn bessel_i0_matches_known_values() {
529        // The third is the β the 12-tap filter's Kaiser window is
530        // designed at, so it is the value that actually gets used.
531        for (x, want) in [(0.0, 1.0), (1.0, 1.266_065_878), (4.664, 20.204_6)] {
532            let got = bessel_i0(x);
533            assert!((got - want).abs() < 1e-3 * want.max(1.0), "I0({x}) = {got}");
534        }
535    }
536}