use std::path::Path;
use anyhow::{Context, Result};
use rayon::prelude::*;
use crate::cpu_gemm::{PackedWeight, gemm_packed};
use crate::weights::LazySt;
pub const FRAME_SIZE: usize = 1920;
pub const NUM_CODEBOOKS: usize = 8;
const DIM: usize = 512; const VQ_DIM: usize = 256; const BINS: usize = 2048; const HEADS: usize = 8;
const HEAD_DIM: usize = DIM / HEADS;
const CONTEXT: usize = 250; const LAYERS: usize = 8;
const FFN: usize = 2048;
const LAYER_EPS: f32 = 1e-5;
const RATIOS: [usize; 4] = [4, 5, 6, 8];
const N_FILTERS: usize = 64;
pub struct Ct {
pub c: usize,
pub t: usize,
pub d: Vec<f32>,
}
impl Ct {
fn zeros(c: usize, t: usize) -> Self {
Ct {
c,
t,
d: vec![0.0; c * t],
}
}
}
fn elu(x: f32) -> f32 {
if x > 0.0 { x } else { x.exp() - 1.0 }
}
fn gelu(x: f32) -> f32 {
0.5 * x * (1.0 + libm::erff(x * std::f32::consts::FRAC_1_SQRT_2))
}
pub(crate) struct Conv1d {
pub(crate) w: Vec<f32>, pub(crate) b: Option<Vec<f32>>,
pub(crate) in_c: usize,
pub(crate) out_c: usize,
pub(crate) k: usize,
pub(crate) stride: usize,
dilation: usize,
groups: usize,
pub(crate) replicate: bool,
packed: std::sync::OnceLock<PackedWeight>,
}
struct ConvState {
prev: Vec<f32>, first: bool,
}
impl Conv1d {
fn load(st: &LazySt, prefix: &str, in_c: usize, out_c: usize, k: usize) -> Result<Self> {
let w = st.tensor_f32(&format!("{prefix}.weight"))?;
anyhow::ensure!(
w.len() == out_c * in_c * k,
"{prefix}.weight len {}",
w.len()
);
let b = st.tensor_f32(&format!("{prefix}.bias")).ok();
if let Some(b) = &b {
anyhow::ensure!(b.len() == out_c, "{prefix}.bias len {}", b.len());
}
Ok(Self {
w,
b,
in_c,
out_c,
k,
stride: 1,
dilation: 1,
groups: 1,
replicate: false,
packed: std::sync::OnceLock::new(),
})
}
pub(crate) fn k_eff(&self) -> usize {
(self.k - 1) * self.dilation + 1
}
fn state(&self) -> ConvState {
ConvState {
prev: vec![0.0; self.in_c * (self.k_eff() - self.stride)],
first: true,
}
}
fn forward(&self, x: &Ct, st: &mut ConvState) -> Ct {
assert_eq!(x.c, self.in_c);
assert!(
x.t > 0 && x.t.is_multiple_of(self.stride),
"steps must be a multiple of stride"
);
let tp = self.k_eff() - self.stride;
if tp > 0 && self.replicate && st.first {
for c in 0..self.in_c {
let v = x.d[c * x.t];
st.prev[c * tp..(c + 1) * tp].fill(v);
}
}
let ta = x.t + tp;
let mut xa = vec![0.0f32; self.in_c * ta];
for c in 0..self.in_c {
xa[c * ta..c * ta + tp].copy_from_slice(&st.prev[c * tp..(c + 1) * tp]);
xa[c * ta + tp..(c + 1) * ta].copy_from_slice(&x.d[c * x.t..(c + 1) * x.t]);
}
let t_out = x.t / self.stride;
if self.groups == 1 {
let packed = self
.packed
.get_or_init(|| PackedWeight::new(&self.w, self.out_c, self.in_c * self.k));
let ck = self.in_c * self.k;
let mut cols = vec![0f32; t_out * ck];
cols.par_chunks_mut(ck).enumerate().for_each(|(ot, crow)| {
let base = ot * self.stride;
for ic in 0..self.in_c {
let xrow = &xa[ic * ta..];
for kk in 0..self.k {
crow[ic * self.k + kk] = xrow[base + kk * self.dilation];
}
}
});
let zeros;
let bias = match &self.b {
Some(b) => b.as_slice(),
None => {
zeros = vec![0f32; self.out_c];
zeros.as_slice()
}
};
let mut y = vec![0f32; t_out * self.out_c];
gemm_packed(&mut y, &cols, packed, t_out, Some(bias));
let mut out = Ct::zeros(self.out_c, t_out);
for ot in 0..t_out {
for oc in 0..self.out_c {
out.d[oc * t_out + ot] = y[ot * self.out_c + oc];
}
}
if tp > 0 {
for c in 0..self.in_c {
st.prev[c * tp..(c + 1) * tp]
.copy_from_slice(&xa[(c + 1) * ta - tp..(c + 1) * ta]);
}
st.first = false;
}
return out;
}
let icg = self.in_c / self.groups;
let ocg = self.out_c / self.groups;
let mut out = Ct::zeros(self.out_c, t_out);
out.d
.par_chunks_mut(t_out)
.enumerate()
.for_each(|(oc, orow)| {
let g = oc / ocg;
let bias = self.b.as_ref().map_or(0.0, |b| b[oc]);
for (ot, o) in orow.iter_mut().enumerate() {
let base = ot * self.stride;
let mut acc = bias;
for ic in 0..icg {
let xrow = &xa[(g * icg + ic) * ta..];
let wrow = &self.w[(oc * icg + ic) * self.k..(oc * icg + ic + 1) * self.k];
for (kk, w) in wrow.iter().enumerate() {
acc += w * xrow[base + kk * self.dilation];
}
}
*o = acc;
}
});
if tp > 0 {
for c in 0..self.in_c {
st.prev[c * tp..(c + 1) * tp].copy_from_slice(&xa[(c + 1) * ta - tp..(c + 1) * ta]);
}
st.first = false;
}
out
}
}
pub(crate) struct ConvTr1d {
pub(crate) w: Vec<f32>, pub(crate) b: Option<Vec<f32>>,
pub(crate) in_c: usize,
pub(crate) out_c: usize,
pub(crate) k: usize,
pub(crate) stride: usize,
pub(crate) groups: usize,
packed_t: std::sync::OnceLock<PackedWeight>,
}
struct ConvTrState {
partial: Vec<f32>, }
impl ConvTr1d {
fn load(
st: &LazySt,
prefix: &str,
in_c: usize,
out_c: usize,
k: usize,
stride: usize,
groups: usize,
) -> Result<Self> {
let w = st.tensor_f32(&format!("{prefix}.weight"))?;
anyhow::ensure!(
w.len() == in_c * (out_c / groups) * k,
"{prefix}.weight len {}",
w.len()
);
let b = st.tensor_f32(&format!("{prefix}.bias")).ok();
Ok(Self {
w,
b,
in_c,
out_c,
k,
stride,
groups,
packed_t: std::sync::OnceLock::new(),
})
}
fn state(&self) -> ConvTrState {
ConvTrState {
partial: vec![0.0; self.out_c * (self.k - self.stride)],
}
}
fn forward(&self, x: &Ct, st: &mut ConvTrState) -> Ct {
assert_eq!(x.c, self.in_c);
let (t, s, k) = (x.t, self.stride, self.k);
let tp = k - s;
let full_t = (t - 1) * s + k;
let icg = self.in_c / self.groups;
let ocg = self.out_c / self.groups;
let mut full = vec![0.0f32; self.out_c * full_t];
if self.groups == 1 {
let packed = self.packed_t.get_or_init(|| {
let mut wt = vec![0f32; self.out_c * self.k * self.in_c];
for ic in 0..self.in_c {
for oc in 0..self.out_c {
for kk in 0..self.k {
wt[(oc * self.k + kk) * self.in_c + ic] =
self.w[(ic * self.out_c + oc) * self.k + kk];
}
}
}
PackedWeight::new(&wt, self.out_c * self.k, self.in_c)
});
let mut a = vec![0f32; t * self.in_c];
for ic in 0..self.in_c {
for ti in 0..t {
a[ti * self.in_c + ic] = x.d[ic * t + ti];
}
}
let mut y = vec![0f32; t * self.out_c * self.k];
gemm_packed(&mut y, &a, packed, t, None);
full.par_chunks_mut(full_t).enumerate().for_each(|(oc, orow)| {
if let Some(b) = &self.b {
orow.fill(b[oc]);
}
for ti in 0..t {
let yrow = &y[ti * self.out_c * self.k + oc * self.k..][..self.k];
let base = ti * s;
for (kk, v) in yrow.iter().enumerate() {
orow[base + kk] += v;
}
}
});
} else {
full.par_chunks_mut(full_t)
.enumerate()
.for_each(|(oc, orow)| {
let g = oc / ocg;
let ocl = oc % ocg;
if let Some(b) = &self.b {
orow.fill(b[oc]);
}
for ic in 0..icg {
let icx = g * icg + ic;
let xrow = &x.d[icx * t..(icx + 1) * t];
let wrow = &self.w[(icx * ocg + ocl) * k..(icx * ocg + ocl + 1) * k];
for (ti, xv) in xrow.iter().enumerate() {
let base = ti * s;
for (kk, w) in wrow.iter().enumerate() {
orow[base + kk] += w * xv;
}
}
}
});
}
let mut out = Ct::zeros(self.out_c, t * s);
for oc in 0..self.out_c {
let orow = &mut full[oc * full_t..(oc + 1) * full_t];
for (i, p) in st.partial[oc * tp..(oc + 1) * tp].iter().enumerate() {
orow[i] += p;
}
let bias = self.b.as_ref().map_or(0.0, |b| b[oc]);
for (i, dst) in st.partial[oc * tp..(oc + 1) * tp].iter_mut().enumerate() {
*dst = orow[full_t - tp + i] - bias;
}
out.d[oc * t * s..(oc + 1) * t * s].copy_from_slice(&orow[..t * s]);
}
out
}
}
pub(crate) struct ResBlock {
pub(crate) c1: Conv1d, pub(crate) c2: Conv1d, }
impl ResBlock {
fn forward(&self, x: &Ct, s1: &mut ConvState, s2: &mut ConvState) -> Ct {
let mut h = Ct {
c: x.c,
t: x.t,
d: x.d.iter().map(|&v| elu(v)).collect(),
};
h = self.c1.forward(&h, s1);
h.d.iter_mut().for_each(|v| *v = elu(*v));
let mut h = self.c2.forward(&h, s2);
for (o, xi) in h.d.iter_mut().zip(&x.d) {
*o += xi; }
h
}
}
pub(crate) struct SeanetEnc {
pub(crate) init: Conv1d, pub(crate) blocks: Vec<(ResBlock, Conv1d)>, pub(crate) last: Conv1d, }
struct SeanetEncState(Vec<ConvState>);
impl SeanetEnc {
fn load(st: &LazySt) -> Result<Self> {
let init = Conv1d::load(st, "encoder.model.0.conv.conv", 1, N_FILTERS, 7)?;
let mut blocks = Vec::new();
let mut mult = 1usize;
for (i, &r) in RATIOS.iter().enumerate() {
let c = mult * N_FILTERS;
let b = 3 * i + 1;
let mut c1 = Conv1d::load(
st,
&format!("encoder.model.{b}.block.1.conv.conv"),
c,
c / 2,
3,
)?;
c1.dilation = 1;
let c2 = Conv1d::load(
st,
&format!("encoder.model.{b}.block.3.conv.conv"),
c / 2,
c,
1,
)?;
let mut down = Conv1d::load(
st,
&format!("encoder.model.{}.conv.conv", b + 2),
c,
c * 2,
2 * r,
)?;
down.stride = r;
blocks.push((ResBlock { c1, c2 }, down));
mult *= 2;
}
let last = Conv1d::load(st, "encoder.model.14.conv.conv", mult * N_FILTERS, DIM, 3)?;
Ok(Self { init, blocks, last })
}
fn state(&self) -> SeanetEncState {
let mut s = vec![self.init.state()];
for (rb, down) in &self.blocks {
s.push(rb.c1.state());
s.push(rb.c2.state());
s.push(down.state());
}
s.push(self.last.state());
SeanetEncState(s)
}
fn forward(&self, x: &Ct, st: &mut SeanetEncState) -> Ct {
let s = &mut st.0;
let mut h = self.init.forward(x, &mut s[0]);
for (i, (rb, down)) in self.blocks.iter().enumerate() {
let (a, rest) = s[3 * i + 1..].split_at_mut(1);
let (b, rest) = rest.split_at_mut(1);
h = rb.forward(&h, &mut a[0], &mut b[0]);
h.d.iter_mut().for_each(|v| *v = elu(*v));
h = down.forward(&h, &mut rest[0]);
}
h.d.iter_mut().for_each(|v| *v = elu(*v));
self.last.forward(&h, s.last_mut().unwrap())
}
}
pub(crate) struct SeanetDec {
pub(crate) init: Conv1d, pub(crate) blocks: Vec<(ConvTr1d, ResBlock)>, pub(crate) last: Conv1d, }
struct SeanetDecState {
convs: Vec<ConvState>,
trs: Vec<ConvTrState>,
}
impl SeanetDec {
fn load(st: &LazySt) -> Result<Self> {
let mut mult = 16usize; let init = Conv1d::load(st, "decoder.model.0.conv.conv", DIM, mult * N_FILTERS, 7)?;
let mut blocks = Vec::new();
for (i, &r) in [8usize, 6, 5, 4].iter().enumerate() {
let c = mult * N_FILTERS;
let b = 3 * i + 2;
let up = ConvTr1d::load(
st,
&format!("decoder.model.{b}.convtr.convtr"),
c,
c / 2,
2 * r,
r,
1,
)?;
let mut c1 = Conv1d::load(
st,
&format!("decoder.model.{}.block.1.conv.conv", b + 1),
c / 2,
c / 4,
3,
)?;
c1.dilation = 1;
let c2 = Conv1d::load(
st,
&format!("decoder.model.{}.block.3.conv.conv", b + 1),
c / 4,
c / 2,
1,
)?;
blocks.push((up, ResBlock { c1, c2 }));
mult /= 2;
}
let last = Conv1d::load(st, "decoder.model.14.conv.conv", N_FILTERS, 1, 3)?;
Ok(Self { init, blocks, last })
}
fn state(&self) -> SeanetDecState {
let mut convs = vec![self.init.state()];
let mut trs = Vec::new();
for (up, rb) in &self.blocks {
trs.push(up.state());
convs.push(rb.c1.state());
convs.push(rb.c2.state());
}
convs.push(self.last.state());
SeanetDecState { convs, trs }
}
fn forward(&self, x: &Ct, st: &mut SeanetDecState) -> Ct {
let mut h = self.init.forward(x, &mut st.convs[0]);
for (i, (up, rb)) in self.blocks.iter().enumerate() {
h.d.iter_mut().for_each(|v| *v = elu(*v));
h = up.forward(&h, &mut st.trs[i]);
let (a, rest) = st.convs[2 * i + 1..].split_at_mut(1);
h = rb.forward(&h, &mut a[0], &mut rest[0]);
}
h.d.iter_mut().for_each(|v| *v = elu(*v));
self.last.forward(&h, st.convs.last_mut().unwrap())
}
}
pub(crate) struct Linear {
n: usize,
k: usize,
pub(crate) w: Vec<f32>,
b: Vec<f32>, packed: std::sync::OnceLock<PackedWeight>,
}
impl Linear {
fn load(st: &LazySt, name: &str, n: usize, k: usize) -> Result<Self> {
let w = st.tensor_f32(name)?;
anyhow::ensure!(w.len() == n * k, "{name} {} != {n}x{k}", w.len());
Ok(Self {
n,
k,
w,
b: vec![0.0; n],
packed: std::sync::OnceLock::new(),
})
}
fn forward(&self, x: &[f32]) -> Vec<f32> {
let m = x.len() / self.k;
let mut out = vec![0f32; m * self.n];
let packed = self
.packed
.get_or_init(|| PackedWeight::new(&self.w, self.n, self.k));
gemm_packed(&mut out, x, packed, m, Some(&self.b));
out
}
}
pub(crate) struct TrLayer {
pub(crate) norm1: (Vec<f32>, Vec<f32>),
pub(crate) in_proj: Linear, pub(crate) out_proj: Linear, pub(crate) norm2: (Vec<f32>, Vec<f32>),
pub(crate) lin1: Linear, pub(crate) lin2: Linear, pub(crate) ls1: Vec<f32>,
pub(crate) ls2: Vec<f32>,
}
pub(crate) struct Transformer {
pub(crate) layers: Vec<TrLayer>,
}
struct LayerKv {
pos: std::collections::VecDeque<usize>,
k: std::collections::VecDeque<[f32; HEAD_DIM * HEADS]>,
v: std::collections::VecDeque<[f32; HEAD_DIM * HEADS]>,
}
pub struct TrState {
kv: Vec<LayerKv>,
offset: usize,
}
fn layer_norm(x: &mut [f32], w: &[f32], b: &[f32]) {
for row in x.chunks_exact_mut(DIM) {
let mean = row.iter().sum::<f32>() / DIM as f32;
let var = row.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / DIM as f32;
let inv = 1.0 / (var + LAYER_EPS).sqrt();
for (i, v) in row.iter_mut().enumerate() {
*v = (*v - mean) * inv * w[i] + b[i];
}
}
}
fn rope_inplace(qk: &mut [f32], t0: usize) {
for (t, row) in qk.chunks_exact_mut(DIM).enumerate() {
let ts = (t0 + t) as f32;
for h in 0..HEADS {
let head = &mut row[h * HEAD_DIM..(h + 1) * HEAD_DIM];
for i in 0..HEAD_DIM / 2 {
let freq = (-(10000f32).ln() * 2.0 * i as f32 / HEAD_DIM as f32).exp();
let (sin, cos) = (freq * ts).sin_cos();
let (r, im) = (head[2 * i], head[2 * i + 1]);
head[2 * i] = r * cos - im * sin;
head[2 * i + 1] = r * sin + im * cos;
}
}
}
}
impl Transformer {
fn load(st: &LazySt, prefix: &str) -> Result<Self> {
let mut layers = Vec::with_capacity(LAYERS);
for l in 0..LAYERS {
let p = format!("{prefix}.transformer.layers.{l}");
layers.push(TrLayer {
norm1: (
st.tensor_f32(&format!("{p}.norm1.weight"))?,
st.tensor_f32(&format!("{p}.norm1.bias"))?,
),
in_proj: Linear::load(
st,
&format!("{p}.self_attn.in_projs.0.weight"),
3 * DIM,
DIM,
)?,
out_proj: Linear::load(st, &format!("{p}.self_attn.out_projs.0.weight"), DIM, DIM)?,
norm2: (
st.tensor_f32(&format!("{p}.norm2.weight"))?,
st.tensor_f32(&format!("{p}.norm2.bias"))?,
),
lin1: Linear::load(st, &format!("{p}.linear1.weight"), FFN, DIM)?,
lin2: Linear::load(st, &format!("{p}.linear2.weight"), DIM, FFN)?,
ls1: st.tensor_f32(&format!("{p}.layer_scale_1.scale"))?,
ls2: st.tensor_f32(&format!("{p}.layer_scale_2.scale"))?,
});
}
Ok(Self { layers })
}
fn state(&self) -> TrState {
TrState {
kv: (0..LAYERS)
.map(|_| LayerKv {
pos: Default::default(),
k: Default::default(),
v: Default::default(),
})
.collect(),
offset: 0,
}
}
fn forward(&self, x: &mut Vec<f32>, st: &mut TrState) {
let t_new = x.len() / DIM;
let t0 = st.offset;
for (l, layer) in self.layers.iter().enumerate() {
let kv = &mut st.kv[l];
let mut h = x.clone();
layer_norm(&mut h, &layer.norm1.0, &layer.norm1.1);
let qkv = layer.in_proj.forward(&h);
let mut q = vec![0f32; t_new * DIM];
let mut k = vec![0f32; t_new * DIM];
let mut v = vec![0f32; t_new * DIM];
for t in 0..t_new {
q[t * DIM..(t + 1) * DIM].copy_from_slice(&qkv[t * 3 * DIM..t * 3 * DIM + DIM]);
k[t * DIM..(t + 1) * DIM]
.copy_from_slice(&qkv[t * 3 * DIM + DIM..t * 3 * DIM + 2 * DIM]);
v[t * DIM..(t + 1) * DIM]
.copy_from_slice(&qkv[t * 3 * DIM + 2 * DIM..(t + 1) * 3 * DIM]);
}
rope_inplace(&mut q, t0);
rope_inplace(&mut k, t0);
for t in 0..t_new {
kv.pos.push_back(t0 + t);
kv.k.push_back(k[t * DIM..(t + 1) * DIM].try_into().unwrap());
kv.v.push_back(v[t * DIM..(t + 1) * DIM].try_into().unwrap());
while kv.pos.len() > CONTEXT {
kv.pos.pop_front();
kv.k.pop_front();
kv.v.pop_front();
}
}
let scale = 1.0 / (HEAD_DIM as f32).sqrt();
let mut attn = vec![0f32; t_new * DIM];
attn.par_chunks_mut(DIM).enumerate().for_each(|(t, arow)| {
let pos_q = t0 + t;
let qrow = &q[t * DIM..(t + 1) * DIM];
for hh in 0..HEADS {
let qh = &qrow[hh * HEAD_DIM..(hh + 1) * HEAD_DIM];
let mut scores = Vec::with_capacity(kv.pos.len());
for (j, &pk) in kv.pos.iter().enumerate() {
if pk > pos_q || pos_q - pk >= CONTEXT {
continue;
}
let kh = &kv.k[j][hh * HEAD_DIM..(hh + 1) * HEAD_DIM];
let dot: f32 = qh.iter().zip(kh).map(|(a, b)| a * b).sum();
scores.push((j, dot * scale));
}
let mx = scores.iter().map(|s| s.1).fold(f32::NEG_INFINITY, f32::max);
let mut den = 0f32;
for s in scores.iter_mut() {
s.1 = (s.1 - mx).exp();
den += s.1;
}
let ah = &mut arow[hh * HEAD_DIM..(hh + 1) * HEAD_DIM];
for (j, wgt) in scores {
let vh = &kv.v[j][hh * HEAD_DIM..(hh + 1) * HEAD_DIM];
for (o, vv) in ah.iter_mut().zip(vh) {
*o += wgt / den * vv;
}
}
}
});
let upd = layer.out_proj.forward(&attn);
for (t, row) in upd.chunks_exact(DIM).enumerate() {
for (i, u) in row.iter().enumerate() {
x[t * DIM + i] += layer.ls1[i] * u;
}
}
let mut h = x.clone();
layer_norm(&mut h, &layer.norm2.0, &layer.norm2.1);
let mut ff = layer.lin1.forward(&h);
ff.iter_mut().for_each(|v| *v = gelu(*v));
let ff = layer.lin2.forward(&ff);
for (t, row) in ff.chunks_exact(DIM).enumerate() {
for (i, u) in row.iter().enumerate() {
x[t * DIM + i] += layer.ls2[i] * u;
}
}
}
st.offset += t_new;
}
fn forward_ct(&self, x: &Ct, st: &mut TrState) -> Ct {
let (c, t) = (x.c, x.t);
let mut tm = vec![0f32; t * c];
for ci in 0..c {
for ti in 0..t {
tm[ti * c + ci] = x.d[ci * t + ti];
}
}
self.forward(&mut tm, st);
let mut out = Ct::zeros(c, t);
for ci in 0..c {
for ti in 0..t {
out.d[ci * t + ti] = tm[ti * c + ci];
}
}
out
}
}
pub(crate) struct RvqHalf {
pub(crate) in_proj: Vec<f32>, pub(crate) out_proj: Vec<f32>, pub(crate) codebooks: Vec<Vec<f32>>, }
impl RvqHalf {
fn load(st: &LazySt, prefix: &str, n_q: usize) -> Result<Self> {
let in_proj = st.tensor_f32(&format!("{prefix}.input_proj.weight"))?;
let out_proj = st.tensor_f32(&format!("{prefix}.output_proj.weight"))?;
anyhow::ensure!(in_proj.len() == VQ_DIM * DIM && out_proj.len() == DIM * VQ_DIM);
let mut codebooks = Vec::with_capacity(n_q);
for l in 0..n_q {
let p = format!("{prefix}.vq.layers.{l}._codebook");
let sum = st.tensor_f32(&format!("{p}.embedding_sum"))?;
let usage = st.tensor_f32(&format!("{p}.cluster_usage"))?;
anyhow::ensure!(sum.len() == BINS * VQ_DIM && usage.len() == BINS);
let emb = sum
.chunks_exact(VQ_DIM)
.zip(&usage)
.flat_map(|(row, &u)| {
let inv = 1.0 / u.max(1e-5);
row.iter().map(move |&v| v * inv)
})
.collect();
codebooks.push(emb);
}
Ok(Self {
in_proj,
out_proj,
codebooks,
})
}
fn project_in(&self, x: &[f32]) -> [f32; VQ_DIM] {
let mut y = [0f32; VQ_DIM];
for (o, wrow) in y.iter_mut().zip(self.in_proj.chunks_exact(DIM)) {
*o = wrow.iter().zip(x).map(|(w, v)| w * v).sum();
}
y
}
fn nearest(codebook: &[f32], x: &[f32]) -> u32 {
let mut best = (f32::INFINITY, 0u32);
for (j, e) in codebook.chunks_exact(VQ_DIM).enumerate() {
let d: f32 = e.iter().zip(x).map(|(a, b)| (a - b) * (a - b)).sum();
if d < best.0 {
best = (d, j as u32);
}
}
best.1
}
fn encode(&self, frame: &[f32], out: &mut [u32]) {
let mut r = self.project_in(frame);
for (l, cb) in self.codebooks.iter().enumerate() {
let idx = Self::nearest(cb, &r);
out[l] = idx;
let e = &cb[idx as usize * VQ_DIM..(idx as usize + 1) * VQ_DIM];
for (rv, ev) in r.iter_mut().zip(e) {
*rv -= ev;
}
}
}
fn decode(&self, codes: &[u32]) -> [f32; DIM] {
let mut acc = [0f32; VQ_DIM];
for (l, &c) in codes.iter().enumerate() {
let e = &self.codebooks[l][c as usize * VQ_DIM..(c as usize + 1) * VQ_DIM];
for (a, v) in acc.iter_mut().zip(e) {
*a += v;
}
}
let mut y = [0f32; DIM];
for (o, wrow) in y.iter_mut().zip(self.out_proj.chunks_exact(VQ_DIM)) {
*o = wrow.iter().zip(&acc).map(|(w, v)| w * v).sum();
}
y
}
}
pub struct Mimi {
pub(crate) enc: SeanetEnc,
pub(crate) enc_tr: Transformer,
pub(crate) down: Conv1d,
pub(crate) rvq_first: RvqHalf, pub(crate) rvq_rest: RvqHalf, pub(crate) up: ConvTr1d,
pub(crate) dec_tr: Transformer,
pub(crate) dec: SeanetDec,
}
pub struct MimiStream<'a> {
m: &'a Mimi,
enc: SeanetEncState,
enc_tr: TrState,
down: ConvState,
up: ConvTrState,
dec_tr: TrState,
dec: SeanetDecState,
}
impl Mimi {
pub fn load(dir: &Path) -> Result<Self> {
let st = LazySt::open(dir).context("mimi checkpoint")?;
let enc = SeanetEnc::load(&st)?;
let enc_tr = Transformer::load(&st, "encoder_transformer")?;
let mut down = Conv1d::load(&st, "downsample.conv.conv.conv", DIM, DIM, 4)?;
down.stride = 2;
down.replicate = true;
let rvq_first = RvqHalf::load(&st, "quantizer.rvq_first", 1)?;
let rvq_rest = RvqHalf::load(&st, "quantizer.rvq_rest", NUM_CODEBOOKS - 1)?;
let up = ConvTr1d::load(&st, "upsample.convtr.convtr.convtr", DIM, DIM, 4, 2, DIM)?;
let dec_tr = Transformer::load(&st, "decoder_transformer")?;
let dec = SeanetDec::load(&st)?;
Ok(Self {
enc,
enc_tr,
down,
rvq_first,
rvq_rest,
up,
dec_tr,
dec,
})
}
pub fn seanet_encode(&self, pcm: &[f32]) -> Ct {
let x = Ct {
c: 1,
t: pcm.len(),
d: pcm.to_vec(),
};
self.enc.forward(&x, &mut self.enc.state())
}
pub fn enc_transformer(&self, x: &Ct) -> Ct {
self.enc_tr.forward_ct(x, &mut self.enc_tr.state())
}
pub fn downsample(&self, x: &Ct) -> Ct {
self.down.forward(x, &mut self.down.state())
}
pub fn rvq_encode(&self, x: &Ct) -> Vec<[u32; NUM_CODEBOOKS]> {
let mut frame = vec![0f32; DIM];
(0..x.t)
.map(|f| {
for c in 0..DIM {
frame[c] = x.d[c * x.t + f];
}
let mut codes = [0u32; NUM_CODEBOOKS];
self.rvq_first.encode(&frame, &mut codes[..1]);
self.rvq_rest.encode(&frame, &mut codes[1..]);
codes
})
.collect()
}
pub fn decode_latent(&self, codes: &[[u32; NUM_CODEBOOKS]]) -> Ct {
let f = codes.len();
let mut out = Ct::zeros(DIM, f);
for (fi, row) in codes.iter().enumerate() {
let a = self.rvq_first.decode(&row[..1]);
let b = self.rvq_rest.decode(&row[1..]);
for c in 0..DIM {
out.d[c * f + fi] = a[c] + b[c];
}
}
out
}
pub fn upsample(&self, x: &Ct) -> Ct {
self.up.forward(x, &mut self.up.state())
}
pub fn dec_transformer(&self, x: &Ct) -> Ct {
self.dec_tr.forward_ct(x, &mut self.dec_tr.state())
}
pub fn seanet_decode(&self, x: &Ct) -> Vec<f32> {
self.dec.forward(x, &mut self.dec.state()).d
}
pub fn encode(&self, pcm: &[f32]) -> Vec<[u32; NUM_CODEBOOKS]> {
let frames = pcm.len().div_ceil(FRAME_SIZE);
let mut padded = pcm.to_vec();
padded.resize(frames * FRAME_SIZE, 0.0);
let latent = self.downsample(&self.enc_transformer(&self.seanet_encode(&padded)));
self.rvq_encode(&latent)
}
pub fn decode(&self, codes: &[[u32; NUM_CODEBOOKS]]) -> Vec<f32> {
let latent = self.decode_latent(codes);
self.seanet_decode(&self.dec_transformer(&self.upsample(&latent)))
}
pub fn stream(&self) -> MimiStream<'_> {
MimiStream {
m: self,
enc: self.enc.state(),
enc_tr: self.enc_tr.state(),
down: self.down.state(),
up: self.up.state(),
dec_tr: self.dec_tr.state(),
dec: self.dec.state(),
}
}
}
impl MimiStream<'_> {
pub fn encode_frame(&mut self, frame: &[f32]) -> [u32; NUM_CODEBOOKS] {
assert_eq!(
frame.len(),
FRAME_SIZE,
"encode_frame wants exactly one 1920-sample frame"
);
let x = Ct {
c: 1,
t: FRAME_SIZE,
d: frame.to_vec(),
};
let h = self.m.enc.forward(&x, &mut self.enc);
let h = self.m.enc_tr.forward_ct(&h, &mut self.enc_tr);
let latent = self.m.down.forward(&h, &mut self.down);
self.m.rvq_encode(&latent).pop().expect("one latent frame")
}
pub fn decode_frame(&mut self, codes: &[u32; NUM_CODEBOOKS]) -> Vec<f32> {
let latent = self.m.decode_latent(std::slice::from_ref(codes));
let h = self.m.up.forward(&latent, &mut self.up);
let h = self.m.dec_tr.forward_ct(&h, &mut self.dec_tr);
self.m.dec.forward(&h, &mut self.dec).d
}
}