Skip to main content

cortiq_engine/
ltxaudio.rs

1//! The LTX-2.5 audio path: the spectrogram VAE decoder and the BigVGAN
2//! vocoder that turns its output into a waveform.
3//!
4//! The transformer denoises sound in the same 48 blocks as the picture, so
5//! the soundtrack arrives as a `[8, T, 16]` latent. Turning that into audio
6//! is two models:
7//!
8//! 1. **The audio VAE decoder** — 2-D convolutions over (time, mel bin) with
9//!    PixelNorm and *height-causal* padding, so a frame never sees the
10//!    future. Two mid blocks, then three levels of three residual blocks
11//!    with a nearest ×2 between them; the first row after each upsample is
12//!    dropped, because the causal padding on the following convolution
13//!    already accounts for it. Out comes a 64-bin log-mel spectrogram.
14//! 2. **BigVGAN v2 with bandwidth extension** — `conv_pre`, six transposed
15//!    convolutions each followed by three anti-aliased multi-receptive-field
16//!    blocks whose outputs are averaged, and `conv_post`. Every activation is
17//!    a SnakeBeta sandwiched between a ×2 sinc upsample and a ×2 sinc
18//!    downsample, which is what keeps the harmonics it generates from
19//!    aliasing. That gives 16 kHz stereo; a second generator predicts a
20//!    residual from its mel spectrogram and adds it to a sinc-resampled copy
21//!    at 48 kHz.
22//!
23//! Everything is f32 throughout: the reference notes that bf16 accumulation
24//! through 108 sequential convolutions costs 40-90 % on spectral metrics.
25
26use crate::pool::Pool;
27use cortiq_core::CmfModel;
28use std::sync::Arc;
29
30fn tensor_f32(model: &Arc<CmfModel>, name: &str) -> Result<(Vec<f32>, Vec<usize>), String> {
31    let e = model.tensor(name).ok_or_else(|| format!("missing tensor {name}"))?;
32    let mut out = vec![0.0f32; e.n_elems()];
33    cortiq_core::quant::dequant_tensor(e, model.entry_bytes(e), &mut out)?;
34    Ok((out, e.shape.clone()))
35}
36
37fn silu(v: f32) -> f32 {
38    v / (1.0 + (-v).exp())
39}
40
41/// A `[C, H, W]` plane: channels, time, mel bin.
42#[derive(Clone)]
43pub struct Grid {
44    pub c: usize,
45    pub h: usize,
46    pub w: usize,
47    pub data: Vec<f32>,
48}
49
50impl Grid {
51    fn zeros(c: usize, h: usize, w: usize) -> Grid {
52        Grid { c, h, w, data: vec![0.0; c * h * w] }
53    }
54    fn n(&self) -> usize {
55        self.h * self.w
56    }
57}
58
59/// 2-D convolution, padded symmetrically on the mel axis and **causally on
60/// the time axis** — the whole kernel extent on the left, nothing on the
61/// right.
62struct Conv2d {
63    w: Vec<f32>,
64    b: Vec<f32>,
65    c_out: usize,
66    c_in: usize,
67    kh: usize,
68    kw: usize,
69}
70
71impl Conv2d {
72    fn load(model: &Arc<CmfModel>, name: &str) -> Result<Conv2d, String> {
73        let (w, s) = tensor_f32(model, &format!("{name}.weight"))?;
74        let (b, _) = tensor_f32(model, &format!("{name}.bias"))?;
75        Ok(Conv2d { w, b, c_out: s[0], c_in: s[1], kh: s[2], kw: s[3] })
76    }
77
78    fn forward(&self, x: &Grid, pool: Option<&Pool>) -> Grid {
79        let (h, w) = (x.h, x.w);
80        let npos = h * w;
81        let k = self.c_in * self.kh * self.kw;
82        let mut out = Grid::zeros(self.c_out, h, w);
83        let pad_h = self.kh - 1; // causal: all of it on the left
84        let pad_w = (self.kw - 1) / 2;
85        const CHUNK: usize = 8192;
86        let mut patches = vec![0f32; CHUNK.min(npos) * k];
87        let mut ys = vec![0f32; CHUNK.min(npos) * self.c_out];
88        let mut p0 = 0usize;
89        while p0 < npos {
90            let n = CHUNK.min(npos - p0);
91            patches[..n * k].fill(0.0);
92            for i in 0..n {
93                let p = p0 + i;
94                let (pwi, phi) = (p % w, p / w);
95                for ci in 0..self.c_in {
96                    for a in 0..self.kh {
97                        let sh = phi as isize + a as isize - pad_h as isize;
98                        if sh < 0 || sh >= h as isize {
99                            continue;
100                        }
101                        for bb in 0..self.kw {
102                            let sw = pwi as isize + bb as isize - pad_w as isize;
103                            if sw < 0 || sw >= w as isize {
104                                continue;
105                            }
106                            patches[i * k + (ci * self.kh + a) * self.kw + bb] =
107                                x.data[(ci * h + sh as usize) * w + sw as usize];
108                        }
109                    }
110                }
111            }
112            crate::fcd_ops::gemm_nt(
113                &patches[..n * k],
114                &self.w,
115                &mut ys[..n * self.c_out],
116                n,
117                k,
118                self.c_out,
119                pool,
120            );
121            for i in 0..n {
122                for co in 0..self.c_out {
123                    out.data[co * npos + p0 + i] = ys[i * self.c_out + co] + self.b[co];
124                }
125            }
126            p0 += n;
127        }
128        out
129    }
130}
131
132/// RMS across channels at each (time, mel) location — no learned weight.
133fn pixel_norm(x: &mut Grid) {
134    let n = x.n();
135    for p in 0..n {
136        let mut ss = 0f64;
137        for c in 0..x.c {
138            let v = x.data[c * n + p] as f64;
139            ss += v * v;
140        }
141        let inv = 1.0 / (ss / x.c as f64 + 1e-6).sqrt();
142        for c in 0..x.c {
143            x.data[c * n + p] = (x.data[c * n + p] as f64 * inv) as f32;
144        }
145    }
146}
147
148struct ResnetBlock {
149    conv1: Conv2d,
150    conv2: Conv2d,
151    shortcut: Option<Conv2d>,
152}
153
154impl ResnetBlock {
155    fn load(model: &Arc<CmfModel>, p: &str) -> Result<ResnetBlock, String> {
156        Ok(ResnetBlock {
157            conv1: Conv2d::load(model, &format!("{p}.conv1.conv"))?,
158            conv2: Conv2d::load(model, &format!("{p}.conv2.conv"))?,
159            shortcut: match model.tensor(&format!("{p}.nin_shortcut.conv.weight")) {
160                Some(_) => Some(Conv2d::load(model, &format!("{p}.nin_shortcut.conv"))?),
161                None => None,
162            },
163        })
164    }
165
166    fn forward(&self, x: &Grid, pool: Option<&Pool>) -> Grid {
167        let mut h = x.clone();
168        pixel_norm(&mut h);
169        h.data.iter_mut().for_each(|v| *v = silu(*v));
170        let mut h = self.conv1.forward(&h, pool);
171        pixel_norm(&mut h);
172        h.data.iter_mut().for_each(|v| *v = silu(*v));
173        let mut h = self.conv2.forward(&h, pool);
174        let res = match &self.shortcut {
175            Some(c) => c.forward(x, pool),
176            None => x.clone(),
177        };
178        for (v, &r) in h.data.iter_mut().zip(&res.data) {
179            *v += r;
180        }
181        h
182    }
183}
184
185/// Nearest ×2 on both axes, a convolution, then the first time row dropped
186/// — the causal padding on the convolution has already reproduced it.
187fn upsample2(x: &Grid, conv: &Conv2d, pool: Option<&Pool>) -> Grid {
188    let (h2, w2) = (x.h * 2, x.w * 2);
189    let mut up = Grid::zeros(x.c, h2, w2);
190    for c in 0..x.c {
191        for y in 0..h2 {
192            for z in 0..w2 {
193                up.data[(c * h2 + y) * w2 + z] = x.data[(c * x.h + y / 2) * x.w + z / 2];
194            }
195        }
196    }
197    let conved = conv.forward(&up, pool);
198    let mut out = Grid::zeros(conved.c, h2 - 1, w2);
199    for c in 0..conved.c {
200        for y in 1..h2 {
201            for z in 0..w2 {
202                out.data[(c * (h2 - 1) + y - 1) * w2 + z] = conved.data[(c * h2 + y) * w2 + z];
203            }
204        }
205    }
206    out
207}
208
209pub struct AudioVaeDecoder {
210    conv_in: Conv2d,
211    mid: Vec<ResnetBlock>,
212    levels: Vec<(Vec<ResnetBlock>, Option<Conv2d>)>,
213    conv_out: Conv2d,
214    mean: Vec<f32>,
215    std: Vec<f32>,
216}
217
218impl AudioVaeDecoder {
219    pub fn from_cmf(model: &Arc<CmfModel>) -> Result<AudioVaeDecoder, String> {
220        let mut levels = Vec::new();
221        let mut lv = 0usize;
222        while model
223            .tensor(&format!("avae.decoder.up.{lv}.block.0.conv1.conv.weight"))
224            .is_some()
225        {
226            let mut blocks = Vec::new();
227            let mut bi = 0usize;
228            while model
229                .tensor(&format!("avae.decoder.up.{lv}.block.{bi}.conv1.conv.weight"))
230                .is_some()
231            {
232                blocks.push(ResnetBlock::load(model, &format!("avae.decoder.up.{lv}.block.{bi}"))?);
233                bi += 1;
234            }
235            let up = match model.tensor(&format!("avae.decoder.up.{lv}.upsample.conv.conv.weight")) {
236                Some(_) => Some(Conv2d::load(model, &format!("avae.decoder.up.{lv}.upsample.conv.conv"))?),
237                None => None,
238            };
239            levels.push((blocks, up));
240            lv += 1;
241        }
242        Ok(AudioVaeDecoder {
243            conv_in: Conv2d::load(model, "avae.decoder.conv_in.conv")?,
244            mid: vec![
245                ResnetBlock::load(model, "avae.decoder.mid.block_1")?,
246                ResnetBlock::load(model, "avae.decoder.mid.block_2")?,
247            ],
248            levels,
249            conv_out: Conv2d::load(model, "avae.decoder.conv_out.conv")?,
250            mean: tensor_f32(model, "avae.per_channel_statistics.mean-of-means")?.0,
251            std: tensor_f32(model, "avae.per_channel_statistics.std-of-means")?.0,
252        })
253    }
254
255    /// `[8, T, 16]` latent → `[2, 4T-3, 64]` log-mel spectrogram.
256    pub fn decode(&self, latent: &Grid, pool: Option<&Pool>) -> Grid {
257        // The statistics are per *patchified* channel (channel-major over mel
258        // bins), which is how the transformer sees them.
259        let mut x = latent.clone();
260        let n = x.n();
261        for c in 0..x.c {
262            for wi in 0..x.w {
263                let idx = c * x.w + wi;
264                let (m, s) = (self.mean[idx], self.std[idx]);
265                for hi in 0..x.h {
266                    let o = (c * x.h + hi) * x.w + wi;
267                    x.data[o] = x.data[o] * s + m;
268                }
269            }
270        }
271        let _ = n;
272        let mut h = self.conv_in.forward(&x, pool);
273        for b in &self.mid {
274            h = b.forward(&h, pool);
275        }
276        for (blocks, up) in self.levels.iter().rev() {
277            for b in blocks {
278                h = b.forward(&h, pool);
279            }
280            if let Some(c) = up {
281                h = upsample2(&h, c, pool);
282            }
283        }
284        pixel_norm(&mut h);
285        h.data.iter_mut().for_each(|v| *v = silu(*v));
286        let out = self.conv_out.forward(&h, pool);
287        // the causal decoder produces 4T-3 frames; crop to that
288        let target = (latent.h * 4).saturating_sub(3).max(1);
289        if out.h == target {
290            return out;
291        }
292        let mut cropped = Grid::zeros(out.c, target.min(out.h), out.w);
293        for c in 0..out.c {
294            for y in 0..cropped.h {
295                for z in 0..out.w {
296                    cropped.data[(c * cropped.h + y) * out.w + z] = out.data[(c * out.h + y) * out.w + z];
297                }
298            }
299        }
300        cropped
301    }
302}
303
304// ------------------------------------------------------------- 1-D layers
305
306/// A multi-channel signal `[C, T]`.
307#[derive(Clone)]
308pub struct Sig {
309    pub c: usize,
310    pub t: usize,
311    pub data: Vec<f32>,
312}
313
314impl Sig {
315    fn zeros(c: usize, t: usize) -> Sig {
316        Sig { c, t, data: vec![0.0; c * t] }
317    }
318}
319
320struct Conv1d {
321    w: Vec<f32>,
322    b: Option<Vec<f32>>,
323    c_out: usize,
324    c_in: usize,
325    k: usize,
326    dilation: usize,
327    pad: usize,
328}
329
330impl Conv1d {
331    fn load(model: &Arc<CmfModel>, name: &str, dilation: usize) -> Result<Conv1d, String> {
332        let (w, s) = tensor_f32(model, &format!("{name}.weight"))?;
333        let b = tensor_f32(model, &format!("{name}.bias")).ok().map(|x| x.0);
334        let k = s[2];
335        Ok(Conv1d { w, b, c_out: s[0], c_in: s[1], k, dilation, pad: (k - 1) * dilation / 2 })
336    }
337
338    fn forward(&self, x: &Sig, pool: Option<&Pool>) -> Sig {
339        let t = x.t;
340        let kk = self.c_in * self.k;
341        let mut patches = vec![0f32; t * kk];
342        for p in 0..t {
343            for ci in 0..self.c_in {
344                for a in 0..self.k {
345                    let s = p as isize + (a * self.dilation) as isize - self.pad as isize;
346                    if s >= 0 && s < t as isize {
347                        patches[p * kk + ci * self.k + a] = x.data[ci * t + s as usize];
348                    }
349                }
350            }
351        }
352        let mut ys = vec![0f32; t * self.c_out];
353        crate::fcd_ops::gemm_nt(&patches, &self.w, &mut ys, t, kk, self.c_out, pool);
354        let mut out = Sig::zeros(self.c_out, t);
355        for p in 0..t {
356            for co in 0..self.c_out {
357                out.data[co * t + p] = ys[p * self.c_out + co] + self.b.as_ref().map_or(0.0, |b| b[co]);
358            }
359        }
360        out
361    }
362}
363
364/// `ConvTranspose1d(in, out, k, stride, padding)`, weights `[in, out, k]`.
365struct ConvT1d {
366    w: Vec<f32>,
367    b: Option<Vec<f32>>,
368    c_in: usize,
369    c_out: usize,
370    k: usize,
371    stride: usize,
372    pad: usize,
373}
374
375impl ConvT1d {
376    fn load(model: &Arc<CmfModel>, name: &str, stride: usize) -> Result<ConvT1d, String> {
377        let (w, s) = tensor_f32(model, &format!("{name}.weight"))?;
378        let b = tensor_f32(model, &format!("{name}.bias")).ok().map(|x| x.0);
379        let k = s[2];
380        Ok(ConvT1d { w, b, c_in: s[0], c_out: s[1], k, stride, pad: (k - stride) / 2 })
381    }
382
383    fn forward(&self, x: &Sig) -> Sig {
384        let t_out = (x.t - 1) * self.stride + self.k - 2 * self.pad;
385        let mut out = Sig::zeros(self.c_out, t_out);
386        for ci in 0..self.c_in {
387            for p in 0..x.t {
388                let v = &x.data[ci * x.t + p];
389                if *v == 0.0 {
390                    continue;
391                }
392                let base = p * self.stride;
393                for a in 0..self.k {
394                    let o = base + a;
395                    if o < self.pad || o - self.pad >= t_out {
396                        continue;
397                    }
398                    let oo = o - self.pad;
399                    for co in 0..self.c_out {
400                        out.data[co * t_out + oo] += v * self.w[(ci * self.c_out + co) * self.k + a];
401                    }
402                }
403            }
404        }
405        if let Some(b) = &self.b {
406            for co in 0..self.c_out {
407                for v in out.data[co * t_out..(co + 1) * t_out].iter_mut() {
408                    *v += b[co];
409                }
410            }
411        }
412        out
413    }
414}
415
416/// The anti-aliasing pair around every activation: a ×2 sinc upsample, the
417/// nonlinearity, a ×2 sinc downsample. Both filters ship in the checkpoint.
418struct Aliasing {
419    up: Vec<f32>,
420    down: Vec<f32>,
421    ratio: usize,
422}
423
424impl Aliasing {
425    fn load(model: &Arc<CmfModel>, p: &str) -> Result<Aliasing, String> {
426        Ok(Aliasing {
427            up: tensor_f32(model, &format!("{p}.upsample.filter"))?.0,
428            down: tensor_f32(model, &format!("{p}.downsample.lowpass.filter"))?.0,
429            ratio: 2,
430        })
431    }
432
433    fn upsample(&self, x: &Sig) -> Sig {
434        let k = self.up.len();
435        let stride = self.ratio;
436        let pad = k / stride - 1;
437        let pad_left = pad * stride + (k - stride) / 2;
438        let pad_right = pad * stride + (k - stride).div_ceil(2);
439        // replicate-pad, transposed convolution, then the same trim the
440        // reference takes
441        let tp = x.t + 2 * pad;
442        let full = (tp - 1) * stride + k;
443        let mut out = Sig::zeros(x.c, full);
444        for c in 0..x.c {
445            for p in 0..tp {
446                let src = (p as isize - pad as isize).clamp(0, x.t as isize - 1) as usize;
447                let v = x.data[c * x.t + src] * self.ratio as f32;
448                if v == 0.0 {
449                    continue;
450                }
451                for a in 0..k {
452                    out.data[c * full + p * stride + a] += v * self.up[a];
453                }
454            }
455        }
456        let (lo, hi) = (pad_left, full - pad_right);
457        let t2 = hi - lo;
458        let mut trimmed = Sig::zeros(x.c, t2);
459        for c in 0..x.c {
460            trimmed.data[c * t2..(c + 1) * t2].copy_from_slice(&out.data[c * full + lo..c * full + hi]);
461        }
462        trimmed
463    }
464
465    fn downsample(&self, x: &Sig) -> Sig {
466        let k = self.down.len();
467        let pad_left = k / 2 - if k % 2 == 0 { 1 } else { 0 };
468        let pad_right = k / 2;
469        let tp = x.t + pad_left + pad_right;
470        let t2 = (tp - k) / self.ratio + 1;
471        let mut out = Sig::zeros(x.c, t2);
472        for c in 0..x.c {
473            for p in 0..t2 {
474                let mut acc = 0f32;
475                for a in 0..k {
476                    let s = (p * self.ratio + a) as isize - pad_left as isize;
477                    let s = s.clamp(0, x.t as isize - 1) as usize;
478                    acc += x.data[c * x.t + s] * self.down[a];
479                }
480                out.data[c * t2 + p] = acc;
481            }
482        }
483        out
484    }
485}
486
487/// `x + sin(αx)² / β`, with α and β kept in log space.
488struct SnakeBeta {
489    alpha: Vec<f32>,
490    beta: Vec<f32>,
491    aa: Aliasing,
492}
493
494impl SnakeBeta {
495    fn load(model: &Arc<CmfModel>, p: &str) -> Result<SnakeBeta, String> {
496        Ok(SnakeBeta {
497            alpha: tensor_f32(model, &format!("{p}.act.alpha"))?.0,
498            beta: tensor_f32(model, &format!("{p}.act.beta"))?.0,
499            aa: Aliasing::load(model, p)?,
500        })
501    }
502
503    fn forward(&self, x: &Sig) -> Sig {
504        let mut up = self.aa.upsample(x);
505        for c in 0..up.c {
506            let a = self.alpha[c].exp();
507            let b = self.beta[c].exp();
508            for v in up.data[c * up.t..(c + 1) * up.t].iter_mut() {
509                let s = (*v * a).sin();
510                *v += s * s / (b + 1e-9);
511            }
512        }
513        self.aa.downsample(&up)
514    }
515}
516
517/// One multi-receptive-field block: three dilated conv pairs, each wrapped
518/// in its own anti-aliased activation, summed into the residual.
519struct AmpBlock {
520    convs1: Vec<Conv1d>,
521    convs2: Vec<Conv1d>,
522    acts1: Vec<SnakeBeta>,
523    acts2: Vec<SnakeBeta>,
524}
525
526impl AmpBlock {
527    fn load(model: &Arc<CmfModel>, p: &str, dil: &[usize]) -> Result<AmpBlock, String> {
528        let mut convs1 = Vec::new();
529        let mut convs2 = Vec::new();
530        let mut acts1 = Vec::new();
531        let mut acts2 = Vec::new();
532        for (i, &d) in dil.iter().enumerate() {
533            convs1.push(Conv1d::load(model, &format!("{p}.convs1.{i}"), d)?);
534            convs2.push(Conv1d::load(model, &format!("{p}.convs2.{i}"), 1)?);
535            acts1.push(SnakeBeta::load(model, &format!("{p}.acts1.{i}"))?);
536            acts2.push(SnakeBeta::load(model, &format!("{p}.acts2.{i}"))?);
537        }
538        Ok(AmpBlock { convs1, convs2, acts1, acts2 })
539    }
540
541    fn forward(&self, x: &Sig, pool: Option<&Pool>) -> Sig {
542        let mut x = x.clone();
543        for i in 0..self.convs1.len() {
544            let h = self.acts1[i].forward(&x);
545            let h = self.convs1[i].forward(&h, pool);
546            let h = self.acts2[i].forward(&h);
547            let h = self.convs2[i].forward(&h, pool);
548            for (v, &y) in x.data.iter_mut().zip(&h.data) {
549                *v += y;
550            }
551        }
552        x
553    }
554}
555
556/// BigVGAN v2: `conv_pre`, then per level a transposed convolution and three
557/// receptive-field blocks whose outputs are averaged, then `conv_post`.
558pub struct Vocoder {
559    conv_pre: Conv1d,
560    ups: Vec<ConvT1d>,
561    blocks: Vec<AmpBlock>,
562    per_level: usize,
563    act_post: SnakeBeta,
564    conv_post: Conv1d,
565    tanh_final: bool,
566    apply_final: bool,
567}
568
569impl Vocoder {
570    fn from_cmf(
571        model: &Arc<CmfModel>,
572        p: &str,
573        rates: &[usize],
574        dils: &[Vec<usize>],
575        apply_final: bool,
576    ) -> Result<Vocoder, String> {
577        let mut ups = Vec::new();
578        for (i, &r) in rates.iter().enumerate() {
579            ups.push(ConvT1d::load(model, &format!("{p}.ups.{i}"), r)?);
580        }
581        let mut blocks = Vec::new();
582        let mut i = 0usize;
583        while model.tensor(&format!("{p}.resblocks.{i}.convs1.0.weight")).is_some() {
584            blocks.push(AmpBlock::load(model, &format!("{p}.resblocks.{i}"), &dils[i % dils.len()])?);
585            i += 1;
586        }
587        let per_level = blocks.len() / rates.len().max(1);
588        Ok(Vocoder {
589            conv_pre: Conv1d::load(model, &format!("{p}.conv_pre"), 1)?,
590            ups,
591            blocks,
592            per_level,
593            act_post: SnakeBeta::load(model, &format!("{p}.act_post"))?,
594            conv_post: Conv1d::load(model, &format!("{p}.conv_post"), 1)?,
595            tanh_final: false,
596            apply_final,
597        })
598    }
599
600    /// `[2, T, mel]` log-mel → `[2, T·∏rates]` waveform.
601    fn forward(&self, mel: &Grid, pool: Option<&Pool>) -> Sig {
602        // (channels, time, mel) → (channels·mel, time)
603        let mut x = Sig::zeros(mel.c * mel.w, mel.h);
604        for s in 0..mel.c {
605            for m in 0..mel.w {
606                let c = s * mel.w + m;
607                for t in 0..mel.h {
608                    x.data[c * mel.h + t] = mel.data[(s * mel.h + t) * mel.w + m];
609                }
610            }
611        }
612        let dbg = std::env::var("CMF_LTX_VOC_DBG").is_ok();
613        let rms = |name: &str, s: &Sig| {
614            let n = s.data.len().max(1) as f64;
615            let r = (s.data.iter().map(|&x| (x as f64) * (x as f64)).sum::<f64>() / n).sqrt();
616            println!("  voc {name:<10} [{}, {}] rms {r:.6}", s.c, s.t);
617        };
618        let mut h = self.conv_pre.forward(&x, pool);
619        if dbg {
620            rms("in", &x);
621            rms("conv_pre", &h);
622        }
623        for (i, up) in self.ups.iter().enumerate() {
624            h = up.forward(&h);
625            let mut acc: Option<Sig> = None;
626            for j in 0..self.per_level {
627                let b = &self.blocks[i * self.per_level + j];
628                let o = b.forward(&h, pool);
629                match &mut acc {
630                    None => acc = Some(o),
631                    Some(a) => {
632                        for (v, &y) in a.data.iter_mut().zip(&o.data) {
633                            *v += y;
634                        }
635                    }
636                }
637            }
638            h = acc.unwrap();
639            let inv = 1.0 / self.per_level as f32;
640            h.data.iter_mut().for_each(|v| *v *= inv);
641            if dbg {
642                rms(&format!("level{i}"), &h);
643            }
644        }
645        let h = self.act_post.forward(&h);
646        let mut out = self.conv_post.forward(&h, pool);
647        if self.apply_final {
648            out.data.iter_mut().for_each(|v| {
649                *v = if self.tanh_final { v.tanh() } else { v.clamp(-1.0, 1.0) }
650            });
651        }
652        out
653    }
654}
655
656/// The causal log-mel the bandwidth extender is conditioned on: an STFT
657/// carried out as a convolution with the checkpoint's own DFT × Hann bases,
658/// so the numbers match what the extender was trained against.
659struct MelStft {
660    forward_basis: Vec<f32>,
661    mel_basis: Vec<f32>,
662    n_freqs: usize,
663    filter_len: usize,
664    hop: usize,
665    n_mels: usize,
666}
667
668impl MelStft {
669    fn load(model: &Arc<CmfModel>, p: &str, hop: usize) -> Result<MelStft, String> {
670        let (fb, fs) = tensor_f32(model, &format!("{p}.stft_fn.forward_basis"))?;
671        let (mb, ms) = tensor_f32(model, &format!("{p}.mel_basis"))?;
672        Ok(MelStft {
673            n_freqs: fs[0] / 2,
674            filter_len: fs[2],
675            forward_basis: fb,
676            mel_basis: mb,
677            n_mels: ms[0],
678            hop,
679        })
680    }
681
682    /// `[C, T]` waveform → `[C, frames, n_mels]` log-mel.
683    fn forward(&self, x: &Sig) -> Grid {
684        let left = self.filter_len.saturating_sub(self.hop);
685        let padded = x.t + left;
686        let frames = if padded >= self.filter_len {
687            (padded - self.filter_len) / self.hop + 1
688        } else {
689            0
690        };
691        let mut out = Grid::zeros(x.c, frames, self.n_mels);
692        let rows = 2 * self.n_freqs;
693        let mut mag = vec![0f32; self.n_freqs];
694        for c in 0..x.c {
695            for f in 0..frames {
696                let start = f * self.hop;
697                for r in 0..self.n_freqs {
698                    let mut re = 0f32;
699                    let mut im = 0f32;
700                    for j in 0..self.filter_len {
701                        let idx = start + j;
702                        let v = if idx < left {
703                            0.0
704                        } else {
705                            let s = idx - left;
706                            if s < x.t { x.data[c * x.t + s] } else { 0.0 }
707                        };
708                        re += v * self.forward_basis[r * self.filter_len + j];
709                        im += v * self.forward_basis[(self.n_freqs + r) * self.filter_len + j];
710                    }
711                    mag[r] = (re * re + im * im).sqrt();
712                }
713                for m in 0..self.n_mels {
714                    let mut acc = 0f32;
715                    for r in 0..self.n_freqs {
716                        acc += self.mel_basis[m * self.n_freqs + r] * mag[r];
717                    }
718                    out.data[(c * frames + f) * self.n_mels + m] = acc.max(1e-5).ln();
719                }
720            }
721        }
722        let _ = rows;
723        out
724    }
725}
726
727/// A Hann-windowed sinc resampler, the ×3 skip path from 16 kHz to 48 kHz.
728/// The reference does not store this filter, so it is rebuilt here.
729fn hann_sinc_upsample(x: &Sig, ratio: usize) -> Sig {
730    let rolloff = 0.99f64;
731    let lpw = 6f64;
732    let width = (lpw / rolloff).ceil() as usize;
733    let k = 2 * width * ratio + 1;
734    let pad = width;
735    let pad_left = 2 * width * ratio;
736    let pad_right = k - ratio;
737    let filt: Vec<f32> = (0..k)
738        .map(|i| {
739            let ta = (i as f64 / ratio as f64 - width as f64) * rolloff;
740            let tc = ta.clamp(-lpw, lpw);
741            let win = (tc * std::f64::consts::PI / lpw / 2.0).cos().powi(2);
742            let s = if ta == 0.0 {
743                1.0
744            } else {
745                (std::f64::consts::PI * ta).sin() / (std::f64::consts::PI * ta)
746            };
747            (s * win * rolloff / ratio as f64) as f32
748        })
749        .collect();
750    let tp = x.t + 2 * pad;
751    let full = (tp - 1) * ratio + k;
752    let mut acc = Sig::zeros(x.c, full);
753    for c in 0..x.c {
754        for p in 0..tp {
755            let src = (p as isize - pad as isize).clamp(0, x.t as isize - 1) as usize;
756            let v = x.data[c * x.t + src] * ratio as f32;
757            if v == 0.0 {
758                continue;
759            }
760            for a in 0..k {
761                acc.data[c * full + p * ratio + a] += v * filt[a];
762            }
763        }
764    }
765    let (lo, hi) = (pad_left, full - pad_right);
766    let t2 = hi - lo;
767    let mut out = Sig::zeros(x.c, t2);
768    for c in 0..x.c {
769        out.data[c * t2..(c + 1) * t2].copy_from_slice(&acc.data[c * full + lo..c * full + hi]);
770    }
771    out
772}
773
774/// The whole audio tail: latent → spectrogram → 16 kHz → 48 kHz.
775pub struct AudioStack {
776    pub decoder: AudioVaeDecoder,
777    vocoder: Vocoder,
778    bwe: Vocoder,
779    mel: MelStft,
780    hop: usize,
781    in_rate: usize,
782    pub out_rate: usize,
783}
784
785impl AudioStack {
786    pub fn from_cmf(model: &Arc<CmfModel>) -> Result<AudioStack, String> {
787        let cfg: serde_json::Value = ["avae.config_json"]
788            .iter()
789            .filter_map(|n| model.tensor(n).map(|e| model.entry_bytes(e)))
790            .filter_map(|b| serde_json::from_slice(b).ok())
791            .next()
792            .unwrap_or(serde_json::Value::Null);
793        let voc = cfg.pointer("/vocoder/vocoder").cloned().unwrap_or_default();
794        let bwe = cfg.pointer("/vocoder/bwe").cloned().unwrap_or_default();
795        let rates = |v: &serde_json::Value, d: Vec<usize>| -> Vec<usize> {
796            v.get("upsample_rates")
797                .and_then(|a| a.as_array())
798                .map(|a| a.iter().filter_map(|x| x.as_u64()).map(|x| x as usize).collect())
799                .unwrap_or(d)
800        };
801        let dils = vec![vec![1usize, 3, 5], vec![1, 3, 5], vec![1, 3, 5]];
802        let hop = bwe.get("hop_length").and_then(|v| v.as_u64()).unwrap_or(80) as usize;
803        Ok(AudioStack {
804            decoder: AudioVaeDecoder::from_cmf(model)?,
805            // The first generator clamps its output; the bandwidth extender
806            // does not, because its output is a residual that is added to a
807            // resampled copy of the first one and clamped after the sum.
808            vocoder: Vocoder::from_cmf(
809                model,
810                "avae.vocoder.vocoder",
811                &rates(&voc, vec![5, 2, 2, 2, 2, 2]),
812                &dils,
813                voc.get("apply_final_activation").and_then(|v| v.as_bool()).unwrap_or(true),
814            )?,
815            bwe: Vocoder::from_cmf(
816                model,
817                "avae.vocoder.bwe_generator",
818                &rates(&bwe, vec![6, 5, 2, 2, 2]),
819                &dils,
820                bwe.get("apply_final_activation").and_then(|v| v.as_bool()).unwrap_or(true),
821            )?,
822            mel: MelStft::load(model, "avae.vocoder.mel_stft", hop)?,
823            hop,
824            in_rate: bwe.get("input_sampling_rate").and_then(|v| v.as_u64()).unwrap_or(16000) as usize,
825            out_rate: bwe.get("output_sampling_rate").and_then(|v| v.as_u64()).unwrap_or(48000) as usize,
826        })
827    }
828
829    /// `[8, T, 16]` latent → stereo waveform at `out_rate`.
830    pub fn decode(&self, latent: &Grid, pool: Option<&Pool>) -> Sig {
831        let mel = self.decoder.decode(latent, pool);
832        self.decode_from_mel(&mel, pool)
833    }
834
835    /// The vocoder half on its own, from a `[2, frames, mel]` log-mel.
836    pub fn decode_from_mel(&self, mel: &Grid, pool: Option<&Pool>) -> Sig {
837        let low = self.vocoder.forward(mel, pool);
838        // the 16 kHz stage on its own, for bisecting the two generators
839        if let Ok(p) = std::env::var("CMF_LTX_LOW_WAV") {
840            let _ = write_wav(std::path::Path::new(&p), &low, self.in_rate);
841        }
842        let out_len = low.t * self.out_rate / self.in_rate;
843        // pad to a whole number of hops so the mel frame count is exact
844        let rem = low.t % self.hop;
845        let padded = if rem == 0 {
846            low.clone()
847        } else {
848            let t2 = low.t + self.hop - rem;
849            let mut p = Sig::zeros(low.c, t2);
850            for c in 0..low.c {
851                p.data[c * t2..c * t2 + low.t].copy_from_slice(&low.data[c * low.t..(c + 1) * low.t]);
852            }
853            p
854        };
855        let m = self.mel.forward(&padded);
856        let residual = self.bwe.forward(&m, pool);
857        let skip = hann_sinc_upsample(&padded, self.out_rate / self.in_rate);
858        let t = residual.t.min(skip.t).min(out_len);
859        let mut out = Sig::zeros(skip.c, t);
860        for c in 0..skip.c {
861            for i in 0..t {
862                out.data[c * t + i] =
863                    (residual.data[c * residual.t + i] + skip.data[c * skip.t + i]).clamp(-1.0, 1.0);
864            }
865        }
866        out
867    }
868}
869
870/// 16-bit PCM WAV — the one container every tool and browser reads.
871pub fn write_wav(path: &std::path::Path, sig: &Sig, rate: usize) -> std::io::Result<()> {
872    use std::io::Write;
873    let n = sig.t;
874    let ch = sig.c as u16;
875    let bytes = (n * sig.c * 2) as u32;
876    let mut f = std::io::BufWriter::new(std::fs::File::create(path)?);
877    f.write_all(b"RIFF")?;
878    f.write_all(&(36 + bytes).to_le_bytes())?;
879    f.write_all(b"WAVEfmt ")?;
880    f.write_all(&16u32.to_le_bytes())?;
881    f.write_all(&1u16.to_le_bytes())?;
882    f.write_all(&ch.to_le_bytes())?;
883    f.write_all(&(rate as u32).to_le_bytes())?;
884    f.write_all(&((rate * sig.c * 2) as u32).to_le_bytes())?;
885    f.write_all(&((sig.c * 2) as u16).to_le_bytes())?;
886    f.write_all(&16u16.to_le_bytes())?;
887    f.write_all(b"data")?;
888    f.write_all(&bytes.to_le_bytes())?;
889    for i in 0..n {
890        for c in 0..sig.c {
891            let v = (sig.data[c * n + i].clamp(-1.0, 1.0) * 32767.0) as i16;
892            f.write_all(&v.to_le_bytes())?;
893        }
894    }
895    Ok(())
896}
897
898// ------------------------------------------------------------ the encoder
899
900/// A slaney mel filterbank — the one `torchaudio.transforms.MelSpectrogram`
901/// builds with `mel_scale="slaney", norm="slaney"`. Not in the checkpoint,
902/// so it is rebuilt from the same formula the reference's preprocessing used.
903fn mel_filterbank(sr: f64, n_fft: usize, n_mels: usize, fmin: f64, fmax: f64) -> Vec<f32> {
904    let n_freqs = n_fft / 2 + 1;
905    let hz_to_mel = |f: f64| 3.0 * f / 200.0;
906    let mel_to_hz = |m: f64| m * 200.0 / 3.0;
907    // slaney is linear below 1 kHz and logarithmic above it
908    let (f_min_log, min_log_mel) = (1000.0f64, 15.0f64);
909    let logstep = (6.4f64).ln() / 27.0;
910    let hz_to_mel_s = |f: f64| {
911        if f >= f_min_log {
912            min_log_mel + (f / f_min_log).ln() / logstep
913        } else {
914            hz_to_mel(f)
915        }
916    };
917    let mel_to_hz_s = |m: f64| {
918        if m >= min_log_mel {
919            f_min_log * ((m - min_log_mel) * logstep).exp()
920        } else {
921            mel_to_hz(m)
922        }
923    };
924    let (m0, m1) = (hz_to_mel_s(fmin), hz_to_mel_s(fmax));
925    let pts: Vec<f64> = (0..n_mels + 2)
926        .map(|i| mel_to_hz_s(m0 + (m1 - m0) * i as f64 / (n_mels + 1) as f64))
927        .collect();
928    let freqs: Vec<f64> = (0..n_freqs).map(|i| sr * i as f64 / n_fft as f64).collect();
929    let mut fb = vec![0f32; n_mels * n_freqs];
930    for m in 0..n_mels {
931        let (lo, ctr, hi) = (pts[m], pts[m + 1], pts[m + 2]);
932        // slaney normalization: unit area per filter
933        let enorm = 2.0 / (hi - lo);
934        for (k, &f) in freqs.iter().enumerate() {
935            let v = if f >= lo && f <= ctr {
936                (f - lo) / (ctr - lo).max(1e-12)
937            } else if f > ctr && f <= hi {
938                (hi - f) / (hi - ctr).max(1e-12)
939            } else {
940                0.0
941            };
942            fb[m * n_freqs + k] = (v * enorm) as f32;
943        }
944    }
945    fb
946}
947
948/// Waveform → log-mel in the layout the audio VAE encodes: `[C, frames, mel]`.
949/// Centered STFT with a Hann window and reflect padding, magnitude (not
950/// power), then the mel projection and a log with the reference's floor.
951pub fn waveform_to_mel(x: &Sig, sr: usize, n_fft: usize, hop: usize, n_mels: usize) -> Grid {
952    let n_freqs = n_fft / 2 + 1;
953    let fb = mel_filterbank(sr as f64, n_fft, n_mels, 0.0, sr as f64 / 2.0);
954    let win: Vec<f32> = (0..n_fft)
955        .map(|i| {
956            let a = std::f64::consts::PI * 2.0 * i as f64 / n_fft as f64;
957            (0.5 - 0.5 * a.cos()) as f32
958        })
959        .collect();
960    let pad = n_fft / 2;
961    let frames = x.t / hop + 1;
962    let mut out = Grid::zeros(x.c, frames, n_mels);
963    let mut re = vec![0f32; n_freqs];
964    let mut im = vec![0f32; n_freqs];
965    for c in 0..x.c {
966        for f in 0..frames {
967            let start = f as isize * hop as isize - pad as isize;
968            re.iter_mut().for_each(|v| *v = 0.0);
969            im.iter_mut().for_each(|v| *v = 0.0);
970            for j in 0..n_fft {
971                // reflect padding at both ends
972                let mut s = start + j as isize;
973                if s < 0 {
974                    s = -s;
975                }
976                if s >= x.t as isize {
977                    s = 2 * (x.t as isize - 1) - s;
978                }
979                let v = if s >= 0 && s < x.t as isize { x.data[c * x.t + s as usize] } else { 0.0 };
980                let v = v * win[j];
981                if v == 0.0 {
982                    continue;
983                }
984                for (k, (rr, ii)) in re.iter_mut().zip(im.iter_mut()).enumerate() {
985                    let a = -2.0 * std::f64::consts::PI * (k * j) as f64 / n_fft as f64;
986                    *rr += v * a.cos() as f32;
987                    *ii += v * a.sin() as f32;
988                }
989            }
990            for m in 0..n_mels {
991                let mut acc = 0f32;
992                for k in 0..n_freqs {
993                    acc += fb[m * n_freqs + k] * (re[k] * re[k] + im[k] * im[k]).sqrt();
994                }
995                out.data[(c * frames + f) * n_mels + m] = acc.max(1e-5).ln();
996            }
997        }
998    }
999    out
1000}
1001
1002struct Downsample2 {
1003    conv: Conv2d,
1004}
1005
1006impl Downsample2 {
1007    /// Stride-2 convolution with the encoder's asymmetric padding: two rows
1008    /// of history on the causal (time) axis, one column on the right.
1009    fn forward(&self, x: &Grid, pool: Option<&Pool>) -> Grid {
1010        let (h, w) = (x.h + 2, x.w + 1);
1011        let mut p = Grid::zeros(x.c, h, w);
1012        for c in 0..x.c {
1013            for y in 0..x.h {
1014                for z in 0..x.w {
1015                    p.data[(c * h + y + 2) * w + z] = x.data[(c * x.h + y) * x.w + z];
1016                }
1017            }
1018        }
1019        self.conv.forward_strided(&p, 2, pool)
1020    }
1021}
1022
1023impl Conv2d {
1024    /// The same convolution with an explicit stride and *no* padding of its
1025    /// own — the caller has already padded.
1026    fn forward_strided(&self, x: &Grid, stride: usize, pool: Option<&Pool>) -> Grid {
1027        let (oh, ow) = ((x.h - self.kh) / stride + 1, (x.w - self.kw) / stride + 1);
1028        let npos = oh * ow;
1029        let k = self.c_in * self.kh * self.kw;
1030        let mut patches = vec![0f32; npos * k];
1031        for i in 0..npos {
1032            let (pw, ph) = (i % ow, i / ow);
1033            for ci in 0..self.c_in {
1034                for a in 0..self.kh {
1035                    for b in 0..self.kw {
1036                        patches[i * k + (ci * self.kh + a) * self.kw + b] =
1037                            x.data[(ci * x.h + ph * stride + a) * x.w + pw * stride + b];
1038                    }
1039                }
1040            }
1041        }
1042        let mut ys = vec![0f32; npos * self.c_out];
1043        crate::fcd_ops::gemm_nt(&patches, &self.w, &mut ys, npos, k, self.c_out, pool);
1044        let mut out = Grid::zeros(self.c_out, oh, ow);
1045        for i in 0..npos {
1046            for co in 0..self.c_out {
1047                out.data[co * npos + i] = ys[i * self.c_out + co] + self.b[co];
1048            }
1049        }
1050        out
1051    }
1052}
1053
1054/// The audio VAE's encoder half: log-mel in, latent out.
1055pub struct AudioVaeEncoder {
1056    conv_in: Conv2d,
1057    levels: Vec<(Vec<ResnetBlock>, Option<Downsample2>)>,
1058    mid: Vec<ResnetBlock>,
1059    conv_out: Conv2d,
1060    mean: Vec<f32>,
1061    std: Vec<f32>,
1062    z: usize,
1063}
1064
1065impl AudioVaeEncoder {
1066    pub fn from_cmf(model: &Arc<CmfModel>) -> Result<AudioVaeEncoder, String> {
1067        let mut levels = Vec::new();
1068        let mut lv = 0usize;
1069        while model
1070            .tensor(&format!("avae.encoder.down.{lv}.block.0.conv1.conv.weight"))
1071            .is_some()
1072        {
1073            let mut blocks = Vec::new();
1074            let mut bi = 0usize;
1075            while model
1076                .tensor(&format!("avae.encoder.down.{lv}.block.{bi}.conv1.conv.weight"))
1077                .is_some()
1078            {
1079                blocks.push(ResnetBlock::load(model, &format!("avae.encoder.down.{lv}.block.{bi}"))?);
1080                bi += 1;
1081            }
1082            let down = match model.tensor(&format!("avae.encoder.down.{lv}.downsample.conv.weight")) {
1083                // the downsample holds a plain Conv2d, not the causal wrapper
1084                // the residual blocks use, so it is one `.conv` shallower
1085                Some(_) => Some(Downsample2 {
1086                    conv: Conv2d::load(model, &format!("avae.encoder.down.{lv}.downsample.conv"))?,
1087                }),
1088                None => None,
1089            };
1090            levels.push((blocks, down));
1091            lv += 1;
1092        }
1093        let out = Conv2d::load(model, "avae.encoder.conv_out.conv")?;
1094        let z = out.c_out / 2;
1095        Ok(AudioVaeEncoder {
1096            conv_in: Conv2d::load(model, "avae.encoder.conv_in.conv")?,
1097            levels,
1098            mid: vec![
1099                ResnetBlock::load(model, "avae.encoder.mid.block_1")?,
1100                ResnetBlock::load(model, "avae.encoder.mid.block_2")?,
1101            ],
1102            conv_out: out,
1103            mean: tensor_f32(model, "avae.per_channel_statistics.mean-of-means")?.0,
1104            std: tensor_f32(model, "avae.per_channel_statistics.std-of-means")?.0,
1105            z,
1106        })
1107    }
1108
1109    /// `[2, frames, 64]` log-mel → `[8, frames/4, 16]` latent.
1110    pub fn encode(&self, mel: &Grid, pool: Option<&Pool>) -> Grid {
1111        let mut h = self.conv_in.forward(mel, pool);
1112        for (blocks, down) in &self.levels {
1113            for b in blocks {
1114                h = b.forward(&h, pool);
1115            }
1116            if let Some(d) = down {
1117                h = d.forward(&h, pool);
1118            }
1119        }
1120        for b in &self.mid {
1121            h = b.forward(&h, pool);
1122        }
1123        pixel_norm(&mut h);
1124        h.data.iter_mut().for_each(|v| *v = silu(*v));
1125        let out = self.conv_out.forward(&h, pool);
1126        // means only, then the per-channel statistics of the *patchified*
1127        // layout (channel-major over mel bins)
1128        let npos = out.n();
1129        let mut lat = Grid::zeros(self.z, out.h, out.w);
1130        for c in 0..self.z {
1131            for hi in 0..out.h {
1132                for wi in 0..out.w {
1133                    let idx = c * out.w + wi;
1134                    let v = out.data[(c * out.h + hi) * out.w + wi];
1135                    lat.data[(c * out.h + hi) * out.w + wi] = (v - self.mean[idx]) / self.std[idx];
1136                }
1137            }
1138        }
1139        lat
1140    }
1141}
1142
1143/// A 16-bit PCM WAV back into a signal in `[-1, 1]`.
1144pub fn read_wav(path: &std::path::Path) -> Result<(Sig, usize), String> {
1145    let raw = std::fs::read(path).map_err(|e| format!("{}: {e}", path.display()))?;
1146    if raw.len() < 44 || &raw[..4] != b"RIFF" || &raw[8..12] != b"WAVE" {
1147        return Err(format!("{}: not a RIFF/WAVE file", path.display()));
1148    }
1149    let mut i = 12usize;
1150    let (mut ch, mut rate, mut bits) = (2usize, 48000usize, 16usize);
1151    let mut data: Option<(usize, usize)> = None;
1152    while i + 8 <= raw.len() {
1153        let id = &raw[i..i + 4];
1154        let len = u32::from_le_bytes(raw[i + 4..i + 8].try_into().unwrap()) as usize;
1155        let body = i + 8;
1156        if id == b"fmt " && body + 16 <= raw.len() {
1157            ch = u16::from_le_bytes(raw[body + 2..body + 4].try_into().unwrap()) as usize;
1158            rate = u32::from_le_bytes(raw[body + 4..body + 8].try_into().unwrap()) as usize;
1159            bits = u16::from_le_bytes(raw[body + 14..body + 16].try_into().unwrap()) as usize;
1160        } else if id == b"data" {
1161            data = Some((body, len.min(raw.len() - body)));
1162            break;
1163        }
1164        i = body + len + (len & 1);
1165    }
1166    let (off, len) = data.ok_or_else(|| format!("{}: no data chunk", path.display()))?;
1167    if bits != 16 {
1168        return Err(format!("{}: only 16-bit PCM is read", path.display()));
1169    }
1170    let n = len / 2 / ch.max(1);
1171    let mut sig = Sig::zeros(ch, n);
1172    for i in 0..n {
1173        for c in 0..ch {
1174            let o = off + (i * ch + c) * 2;
1175            let v = i16::from_le_bytes([raw[o], raw[o + 1]]) as f32 / 32768.0;
1176            sig.data[c * n + i] = v;
1177        }
1178    }
1179    Ok((sig, rate))
1180}
1181
1182/// Resample by linear interpolation — good enough for conditioning input,
1183/// which the mel transform is about to smear across 64 bands anyway.
1184pub fn resample(x: &Sig, from: usize, to: usize) -> Sig {
1185    if from == to {
1186        return x.clone();
1187    }
1188    let n = (x.t as f64 * to as f64 / from as f64).round() as usize;
1189    let mut out = Sig::zeros(x.c, n);
1190    for c in 0..x.c {
1191        for i in 0..n {
1192            let p = i as f64 * from as f64 / to as f64;
1193            let j = p.floor() as usize;
1194            let f = (p - j as f64) as f32;
1195            let a = x.data[c * x.t + j.min(x.t - 1)];
1196            let b = x.data[c * x.t + (j + 1).min(x.t - 1)];
1197            out.data[c * n + i] = a + (b - a) * f;
1198        }
1199    }
1200    out
1201}