1use 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#[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
59struct 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; 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
132fn 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
185fn 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 pub fn decode(&self, latent: &Grid, pool: Option<&Pool>) -> Grid {
257 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 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#[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
364struct 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
416struct 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 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
487struct 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
517struct 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
556pub 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 fn forward(&self, mel: &Grid, pool: Option<&Pool>) -> Sig {
602 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
656struct 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 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
727fn 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
774pub 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 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 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 pub fn decode_from_mel(&self, mel: &Grid, pool: Option<&Pool>) -> Sig {
837 let low = self.vocoder.forward(mel, pool);
838 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 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
870pub 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
898fn 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 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 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
948pub 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 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 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 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
1054pub 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 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 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 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
1143pub 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
1182pub 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}