use crate::qwen3tts::chatterbox_st::St;
use crate::qwen3tts::{Qwen3TtsConfig, cfg, matvec, rms_norm, silu, softmax_inplace};
use anyhow::Result;
use rayon::prelude::*;
use std::path::Path;
fn gelu_erf(x: f32) -> f32 {
0.5 * x * (1.0 + crate::simd_math::erf_f32(x * std::f32::consts::FRAC_1_SQRT_2))
}
fn snake_beta(x: &mut [f32], alpha: &[f32], beta: &[f32], t_len: usize) {
for (c, (a, b)) in alpha.iter().zip(beta).enumerate() {
let ea = a.exp();
let inv_b = 1.0 / (b.exp() + 1e-9);
for v in &mut x[c * t_len..(c + 1) * t_len] {
let s = (*v * ea).sin();
*v += inv_b * s * s;
}
}
}
pub(crate) struct Conv1d {
w: Vec<f32>, b: Vec<f32>,
c_in: usize,
c_out: usize,
k: usize,
dilation: usize,
stride: usize,
groups: usize,
}
impl Conv1d {
#[allow(clippy::too_many_arguments)]
fn load(
st: &St,
prefix: &str,
c_in: usize,
c_out: usize,
k: usize,
dilation: usize,
stride: usize,
groups: usize,
) -> Result<Self> {
let w = st.f32(&format!("{prefix}.weight"))?;
anyhow::ensure!(
w.len() == c_out * (c_in / groups) * k,
"{prefix}: weight len {} != {c_out}×{}×{k}",
w.len(),
c_in / groups
);
Ok(Self {
w,
b: st.f32(&format!("{prefix}.bias"))?,
c_in,
c_out,
k,
dilation,
stride,
groups,
})
}
#[allow(clippy::too_many_arguments)]
fn load_no_bias(
st: &St,
prefix: &str,
c_in: usize,
c_out: usize,
k: usize,
dilation: usize,
stride: usize,
groups: usize,
) -> Result<Self> {
let w = st.f32(&format!("{prefix}.weight"))?;
anyhow::ensure!(
w.len() == c_out * (c_in / groups) * k,
"{prefix}: weight len {} != {c_out}×{}×{k}",
w.len(),
c_in / groups
);
Ok(Self {
w,
b: vec![0f32; c_out],
c_in,
c_out,
k,
dilation,
stride,
groups,
})
}
fn forward(&self, x: &[f32], t_len: usize) -> (Vec<f32>, usize) {
let eff_k = (self.k - 1) * self.dilation + 1;
let left = eff_k - self.stride;
let padded_min = t_len + left;
let n_out = padded_min
.saturating_sub(eff_k)
.div_ceil(self.stride.max(1))
+ 1;
let cig = self.c_in / self.groups;
let cog = self.c_out / self.groups;
let mut out = vec![0f32; self.c_out * n_out];
out.par_chunks_mut(n_out)
.enumerate()
.for_each(|(oc, orow)| {
let g = oc / cog;
for ci in 0..cig {
let ic = g * cig + ci;
let xrow = &x[ic * t_len..(ic + 1) * t_len];
let wrow = &self.w[(oc * cig + ci) * self.k..(oc * cig + ci + 1) * self.k];
for (o, out_v) in orow.iter_mut().enumerate() {
let base = o * self.stride; let mut acc = 0f32;
for (j, &wv) in wrow.iter().enumerate() {
let p = base + j * self.dilation;
if p >= left {
let xi = p - left;
if xi < t_len {
acc += wv * xrow[xi];
}
}
}
*out_v += acc;
}
}
for v in orow.iter_mut() {
*v += self.b[oc];
}
});
(out, n_out)
}
fn forward_stream(&self, state: &mut Vec<f32>, x: &[f32], t_len: usize) -> Vec<f32> {
debug_assert_eq!(self.stride, 1, "streaming path is stride-1 only");
let eff_k = (self.k - 1) * self.dilation + 1;
let left = eff_k - 1;
if state.is_empty() {
*state = vec![0f32; self.c_in * left];
}
let ext_t = left + t_len;
let mut ext = vec![0f32; self.c_in * ext_t];
for c in 0..self.c_in {
ext[c * ext_t..c * ext_t + left].copy_from_slice(&state[c * left..(c + 1) * left]);
ext[c * ext_t + left..(c + 1) * ext_t].copy_from_slice(&x[c * t_len..(c + 1) * t_len]);
}
let cig = self.c_in / self.groups;
let cog = self.c_out / self.groups;
let mut out = vec![0f32; self.c_out * t_len];
out.par_chunks_mut(t_len)
.enumerate()
.for_each(|(oc, orow)| {
let g = oc / cog;
for ci in 0..cig {
let ic = g * cig + ci;
let xrow = &ext[ic * ext_t..(ic + 1) * ext_t];
let wrow = &self.w[(oc * cig + ci) * self.k..(oc * cig + ci + 1) * self.k];
for (o, out_v) in orow.iter_mut().enumerate() {
let mut acc = 0f32;
for (j, &wv) in wrow.iter().enumerate() {
acc += wv * xrow[o + j * self.dilation];
}
*out_v += acc;
}
}
for v in orow.iter_mut() {
*v += self.b[oc];
}
});
if left > 0 {
for c in 0..self.c_in {
state[c * left..(c + 1) * left]
.copy_from_slice(&ext[(c + 1) * ext_t - left..(c + 1) * ext_t]);
}
}
out
}
}
struct ConvTr1d {
w: Vec<f32>,
b: Vec<f32>,
c_in: usize,
c_out: usize,
k: usize,
stride: usize,
}
impl ConvTr1d {
fn load(
st: &St,
prefix: &str,
c_in: usize,
c_out: usize,
k: usize,
stride: usize,
) -> Result<Self> {
let w = st.f32(&format!("{prefix}.weight"))?;
anyhow::ensure!(
w.len() == c_in * c_out * k,
"{prefix}: weight len {} != {c_in}×{c_out}×{k}",
w.len()
);
Ok(Self {
w,
b: st.f32(&format!("{prefix}.bias"))?,
c_in,
c_out,
k,
stride,
})
}
fn accumulate_raw(&self, raw: &mut [f32], x: &[f32], t_len: usize) {
let full = (t_len - 1) * self.stride + self.k;
debug_assert_eq!(raw.len(), self.c_out * full);
raw.par_chunks_mut(full).enumerate().for_each(|(co, orow)| {
for i in 0..t_len {
let base = i * self.stride;
for ci in 0..self.c_in {
let xv = x[ci * t_len + i];
if xv == 0.0 {
continue;
}
let wrow = &self.w
[(ci * self.c_out + co) * self.k..(ci * self.c_out + co + 1) * self.k];
let dst = &mut orow[base..base + self.k];
for (j, &wv) in wrow.iter().enumerate() {
dst[j] += xv * wv;
}
}
}
});
}
fn forward(&self, x: &[f32], t_len: usize) -> (Vec<f32>, usize) {
let full = (t_len - 1) * self.stride + self.k;
let n_out = t_len * self.stride; let mut raw = vec![0f32; self.c_out * full];
self.accumulate_raw(&mut raw, x, t_len);
let mut out = vec![0f32; self.c_out * n_out];
for co in 0..self.c_out {
for t in 0..n_out {
out[co * n_out + t] = raw[co * full + t] + self.b[co];
}
}
(out, n_out)
}
fn forward_stream(&self, carry: &mut Vec<f32>, x: &[f32], t_len: usize) -> Vec<f32> {
let trim = self.k - self.stride;
let full = (t_len - 1) * self.stride + self.k;
let n_out = t_len * self.stride;
if carry.is_empty() {
*carry = vec![0f32; self.c_out * trim];
}
let mut raw = vec![0f32; self.c_out * full];
for co in 0..self.c_out {
raw[co * full..co * full + trim].copy_from_slice(&carry[co * trim..(co + 1) * trim]);
}
self.accumulate_raw(&mut raw, x, t_len);
let mut out = vec![0f32; self.c_out * n_out];
for co in 0..self.c_out {
for t in 0..n_out {
out[co * n_out + t] = raw[co * full + t] + self.b[co];
}
for j in 0..trim {
carry[co * trim + j] = raw[co * full + n_out + j];
}
}
out
}
}
struct RvqDec {
codebooks: Vec<Vec<f32>>,
out_w: Vec<f32>, }
impl RvqDec {
fn load(st: &St, prefix: &str, n_layers: usize) -> Result<Self> {
let mut codebooks = Vec::with_capacity(n_layers);
for l in 0..n_layers {
let sum = st.f32(&format!("{prefix}.vq.layers.{l}._codebook.embedding_sum"))?;
let usage = st.f32(&format!("{prefix}.vq.layers.{l}._codebook.cluster_usage"))?;
let dim = cfg::ENC_CODEBOOK_DIM;
anyhow::ensure!(sum.len() == cfg::CP_VOCAB * dim && usage.len() == cfg::CP_VOCAB);
let mut eff = sum;
for (row, &u) in eff.chunks_exact_mut(dim).zip(&usage) {
let inv = 1.0 / u.max(1e-5);
for v in row {
*v *= inv;
}
}
codebooks.push(eff);
}
let out_w = st.f32(&format!("{prefix}.output_proj.weight"))?;
anyhow::ensure!(out_w.len() == cfg::DEC_HIDDEN * cfg::ENC_CODEBOOK_DIM);
Ok(Self { codebooks, out_w })
}
fn decode(&self, codes: &[Vec<u32>], t_len: usize) -> Vec<f32> {
let dim = cfg::ENC_CODEBOOK_DIM;
let mut acc = vec![0f32; t_len * dim]; for (layer, cb) in codes.iter().zip(&self.codebooks) {
for (t, &code) in layer.iter().enumerate() {
let row = &cb[code as usize * dim..(code as usize + 1) * dim];
for (a, v) in acc[t * dim..(t + 1) * dim].iter_mut().zip(row) {
*a += *v;
}
}
}
let mut out = vec![0f32; cfg::DEC_HIDDEN * t_len];
for t in 0..t_len {
let y = matvec(
&self.out_w,
&acc[t * dim..(t + 1) * dim],
cfg::DEC_HIDDEN,
dim,
);
for (c, v) in y.iter().enumerate() {
out[c * t_len + t] = *v;
}
}
out
}
}
struct TrLayer {
input_ln: Vec<f32>,
q: Vec<f32>,
k: Vec<f32>,
v: Vec<f32>,
o: Vec<f32>,
attn_scale: Vec<f32>,
post_ln: Vec<f32>,
gate: Vec<f32>,
up: Vec<f32>,
down: Vec<f32>,
mlp_scale: Vec<f32>,
}
struct PreTransformer {
in_w: Vec<f32>,
in_b: Vec<f32>,
out_w: Vec<f32>,
out_b: Vec<f32>,
norm: Vec<f32>,
layers: Vec<TrLayer>,
}
impl PreTransformer {
fn load(st: &St) -> Result<Self> {
let (h, aw, inter) = (
cfg::DEC_HIDDEN,
cfg::DEC_HEADS * cfg::DEC_HEAD_DIM,
cfg::DEC_LATENT,
);
let p = "decoder.pre_transformer";
let mut layers = Vec::with_capacity(cfg::DEC_LAYERS);
for l in 0..cfg::DEC_LAYERS {
let q = format!("{p}.layers.{l}");
layers.push(TrLayer {
input_ln: st.f32(&format!("{q}.input_layernorm.weight"))?,
q: st.mat(&format!("{q}.self_attn.q_proj.weight"), aw, h)?,
k: st.mat(&format!("{q}.self_attn.k_proj.weight"), aw, h)?,
v: st.mat(&format!("{q}.self_attn.v_proj.weight"), aw, h)?,
o: st.mat(&format!("{q}.self_attn.o_proj.weight"), h, aw)?,
attn_scale: st.f32(&format!("{q}.self_attn_layer_scale.scale"))?,
post_ln: st.f32(&format!("{q}.post_attention_layernorm.weight"))?,
gate: st.mat(&format!("{q}.mlp.gate_proj.weight"), inter, h)?,
up: st.mat(&format!("{q}.mlp.up_proj.weight"), inter, h)?,
down: st.mat(&format!("{q}.mlp.down_proj.weight"), h, inter)?,
mlp_scale: st.f32(&format!("{q}.mlp_layer_scale.scale"))?,
});
}
Ok(Self {
in_w: st.mat(&format!("{p}.input_proj.weight"), h, cfg::DEC_LATENT)?,
in_b: st.f32(&format!("{p}.input_proj.bias"))?,
out_w: st.mat(&format!("{p}.output_proj.weight"), cfg::DEC_LATENT, h)?,
out_b: st.f32(&format!("{p}.output_proj.bias"))?,
norm: st.f32(&format!("{p}.norm.weight"))?,
layers,
})
}
fn forward(&self, x_in: &[f32], t_len: usize) -> Vec<f32> {
let (h, hd, heads) = (cfg::DEC_HIDDEN, cfg::DEC_HEAD_DIM, cfg::DEC_HEADS);
let aw = heads * hd;
let eps = cfg::DEC_RMS_EPS;
let scale = 1.0 / (hd as f32).sqrt();
let win = cfg::DEC_SLIDING_WINDOW;
let mut x: Vec<f32> = Vec::with_capacity(t_len * h);
for t in 0..t_len {
let mut r = matvec(
&self.in_w,
&x_in[t * cfg::DEC_LATENT..(t + 1) * cfg::DEC_LATENT],
h,
cfg::DEC_LATENT,
);
for (v, b) in r.iter_mut().zip(&self.in_b) {
*v += *b;
}
x.extend(r);
}
let tables: Vec<(Vec<f32>, Vec<f32>)> = (0..t_len)
.map(|t| crate::qwen3tts::rope_cs_1d(hd, cfg::DEC_ROPE_THETA, t as u32))
.collect();
for layer in &self.layers {
let mut q = vec![0f32; t_len * aw];
let mut k = vec![0f32; t_len * aw];
let mut v = vec![0f32; t_len * aw];
for t in 0..t_len {
let xn = rms_norm(&x[t * h..(t + 1) * h], &layer.input_ln, eps);
q[t * aw..(t + 1) * aw].copy_from_slice(&matvec(&layer.q, &xn, aw, h));
k[t * aw..(t + 1) * aw].copy_from_slice(&matvec(&layer.k, &xn, aw, h));
v[t * aw..(t + 1) * aw].copy_from_slice(&matvec(&layer.v, &xn, aw, h));
let (cos, sin) = &tables[t];
crate::qwen3tts::apply_rope(&mut q[t * aw..(t + 1) * aw], cos, sin, heads, hd);
crate::qwen3tts::apply_rope(&mut k[t * aw..(t + 1) * aw], cos, sin, heads, hd);
}
for t in 0..t_len {
let lo = t.saturating_sub(win - 1); let mut attn = vec![0f32; aw];
for head in 0..heads {
let qh = &q[t * aw + head * hd..t * aw + (head + 1) * hd];
let mut scores: Vec<f32> = (lo..=t)
.map(|s| {
let kh = &k[s * aw + head * hd..s * aw + (head + 1) * hd];
qh.iter().zip(kh).map(|(a, b)| a * b).sum::<f32>() * scale
})
.collect();
softmax_inplace(&mut scores);
let out = &mut attn[head * hd..(head + 1) * hd];
for (idx, &w) in scores.iter().enumerate() {
let s = lo + idx;
let vh = &v[s * aw + head * hd..s * aw + (head + 1) * hd];
for (d, ov) in out.iter_mut().enumerate() {
*ov += w * vh[d];
}
}
}
let o = matvec(&layer.o, &attn, h, aw);
for (j, (ov, sc)) in o.iter().zip(&layer.attn_scale).enumerate() {
x[t * h + j] += ov * sc;
}
let xn = rms_norm(&x[t * h..(t + 1) * h], &layer.post_ln, eps);
let g = matvec(&layer.gate, &xn, cfg::DEC_LATENT, h);
let u = matvec(&layer.up, &xn, cfg::DEC_LATENT, h);
let act: Vec<f32> = g.iter().zip(&u).map(|(gv, uv)| silu(*gv) * uv).collect();
let d = matvec(&layer.down, &act, h, cfg::DEC_LATENT);
for (j, (dv, sc)) in d.iter().zip(&layer.mlp_scale).enumerate() {
x[t * h + j] += dv * sc;
}
}
}
let mut out = Vec::with_capacity(t_len * cfg::DEC_LATENT);
for t in 0..t_len {
let xn = rms_norm(&x[t * h..(t + 1) * h], &self.norm, eps);
let mut r = matvec(&self.out_w, &xn, cfg::DEC_LATENT, h);
for (v, b) in r.iter_mut().zip(&self.out_b) {
*v += *b;
}
out.extend(r);
}
out
}
}
struct ConvNeXt {
dw: Conv1d,
norm_w: Vec<f32>,
norm_b: Vec<f32>,
pw1_w: Vec<f32>,
pw1_b: Vec<f32>,
pw2_w: Vec<f32>,
pw2_b: Vec<f32>,
gamma: Vec<f32>,
}
impl ConvNeXt {
fn load(st: &St, prefix: &str, dim: usize) -> Result<Self> {
Ok(Self {
dw: Conv1d::load(st, &format!("{prefix}.dwconv.conv"), dim, dim, 7, 1, 1, dim)?,
norm_w: st.f32(&format!("{prefix}.norm.weight"))?,
norm_b: st.f32(&format!("{prefix}.norm.bias"))?,
pw1_w: st.mat(&format!("{prefix}.pwconv1.weight"), 4 * dim, dim)?,
pw1_b: st.f32(&format!("{prefix}.pwconv1.bias"))?,
pw2_w: st.mat(&format!("{prefix}.pwconv2.weight"), dim, 4 * dim)?,
pw2_b: st.f32(&format!("{prefix}.pwconv2.bias"))?,
gamma: st.f32(&format!("{prefix}.gamma"))?,
})
}
fn forward(&self, x: &[f32], t_len: usize) -> Vec<f32> {
let dim = self.gamma.len();
let (dwo, _) = self.dw.forward(x, t_len);
let mut out = x.to_vec();
for t in 0..t_len {
let mut col: Vec<f32> = (0..dim).map(|c| dwo[c * t_len + t]).collect();
let mean = col.iter().sum::<f32>() / dim as f32;
let var = col.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / dim as f32;
let inv = 1.0 / (var + 1e-6).sqrt();
for (v, (w, b)) in col.iter_mut().zip(self.norm_w.iter().zip(&self.norm_b)) {
*v = (*v - mean) * inv * w + b;
}
let mut a = matvec(&self.pw1_w, &col, 4 * dim, dim);
for (v, b) in a.iter_mut().zip(&self.pw1_b) {
*v = gelu_erf(*v + *b);
}
let mut y = matvec(&self.pw2_w, &a, dim, 4 * dim);
for ((v, b), g) in y.iter_mut().zip(&self.pw2_b).zip(&self.gamma) {
*v = (*v + *b) * g;
}
for (c, v) in y.iter().enumerate() {
out[c * t_len + t] += *v;
}
}
out
}
}
struct ResUnit {
a1: (Vec<f32>, Vec<f32>),
c1: Conv1d,
a2: (Vec<f32>, Vec<f32>),
c2: Conv1d,
}
struct DecBlock {
snake: (Vec<f32>, Vec<f32>),
tr: ConvTr1d,
units: Vec<ResUnit>,
}
pub struct CodecDecoder {
rvq_first: RvqDec,
rvq_rest: RvqDec,
pre_conv: Conv1d,
pre_tr: PreTransformer,
upsample: Vec<(ConvTr1d, ConvNeXt)>,
dec0: Conv1d,
blocks: Vec<DecBlock>,
final_snake: (Vec<f32>, Vec<f32>),
final_conv: Conv1d,
}
impl CodecDecoder {
pub fn load(dir: &Path, config: &Qwen3TtsConfig) -> Result<Self> {
let d = &config.codec.decoder_config;
let st = St::open(&dir.join(cfg::DIR_CODEC).join(cfg::FILE_MODEL))?;
let lat = d.latent_dim;
let mut upsample = Vec::new();
for (s, &factor) in d.upsampling_ratios.iter().enumerate() {
upsample.push((
ConvTr1d::load(
&st,
&format!("decoder.upsample.{s}.0.conv"),
lat,
lat,
factor,
factor,
)?,
ConvNeXt::load(&st, &format!("decoder.upsample.{s}.1"), lat)?,
));
}
let mut blocks = Vec::new();
for (i, &rate) in d.upsample_rates.iter().enumerate() {
let in_dim = d.decoder_dim / (1 << i);
let out_dim = d.decoder_dim / (1 << (i + 1));
let p = format!("decoder.decoder.{}", i + 1);
let mut units = Vec::new();
for (u, dil) in [1usize, 3, 9].iter().enumerate() {
let q = format!("{p}.block.{}", u + 2);
units.push(ResUnit {
a1: (
st.f32(&format!("{q}.act1.alpha"))?,
st.f32(&format!("{q}.act1.beta"))?,
),
c1: Conv1d::load(
&st,
&format!("{q}.conv1.conv"),
out_dim,
out_dim,
7,
*dil,
1,
1,
)?,
a2: (
st.f32(&format!("{q}.act2.alpha"))?,
st.f32(&format!("{q}.act2.beta"))?,
),
c2: Conv1d::load(
&st,
&format!("{q}.conv2.conv"),
out_dim,
out_dim,
1,
1,
1,
1,
)?,
});
}
blocks.push(DecBlock {
snake: (
st.f32(&format!("{p}.block.0.alpha"))?,
st.f32(&format!("{p}.block.0.beta"))?,
),
tr: ConvTr1d::load(
&st,
&format!("{p}.block.1.conv"),
in_dim,
out_dim,
2 * rate,
rate,
)?,
units,
});
}
Ok(Self {
rvq_first: RvqDec::load(&st, "decoder.quantizer.rvq_first", 1)?,
rvq_rest: RvqDec::load(&st, "decoder.quantizer.rvq_rest", cfg::NUM_CODE_GROUPS - 1)?,
pre_conv: Conv1d::load(
&st,
"decoder.pre_conv.conv",
cfg::DEC_HIDDEN,
lat,
3,
1,
1,
1,
)?,
pre_tr: PreTransformer::load(&st)?,
upsample,
dec0: Conv1d::load(
&st,
"decoder.decoder.0.conv",
lat,
d.decoder_dim,
7,
1,
1,
1,
)?,
blocks,
final_snake: (
st.f32("decoder.decoder.5.alpha")?,
st.f32("decoder.decoder.5.beta")?,
),
final_conv: Conv1d::load(&st, "decoder.decoder.6.conv", 96, 1, 7, 1, 1, 1)?,
})
}
pub fn decode(&self, frames: &[[u32; cfg::NUM_CODE_GROUPS]]) -> Vec<f32> {
let t_len = frames.len();
if t_len == 0 {
return Vec::new();
}
let sem: Vec<Vec<u32>> = vec![frames.iter().map(|f| f[0]).collect()];
let ac: Vec<Vec<u32>> = (1..cfg::NUM_CODE_GROUPS)
.map(|g| frames.iter().map(|f| f[g]).collect())
.collect();
let mut latent = self.rvq_first.decode(&sem, t_len); let rest = self.rvq_rest.decode(&ac, t_len);
for (a, b) in latent.iter_mut().zip(&rest) {
*a += *b;
}
let (pc, t1) = self.pre_conv.forward(&latent, t_len);
debug_assert_eq!(t1, t_len);
let lat = cfg::DEC_LATENT;
let mut rows = vec![0f32; t_len * lat];
for c in 0..lat {
for t in 0..t_len {
rows[t * lat + c] = pc[c * t_len + t];
}
}
let tr_out = self.pre_tr.forward(&rows, t_len);
let mut x = vec![0f32; lat * t_len];
for t in 0..t_len {
for c in 0..lat {
x[c * t_len + t] = tr_out[t * lat + c];
}
}
let mut cur_t = t_len;
for (tr, cnx) in &self.upsample {
let (y, t2) = tr.forward(&x, cur_t);
x = cnx.forward(&y, t2);
cur_t = t2;
}
let (mut wav, mut wt) = self.dec0.forward(&x, cur_t);
for b in &self.blocks {
snake_beta(&mut wav, &b.snake.0, &b.snake.1, wt);
let (y, t2) = b.tr.forward(&wav, wt);
wav = y;
wt = t2;
for u in &b.units {
let residual = wav.clone();
snake_beta(&mut wav, &u.a1.0, &u.a1.1, wt);
let (y1, _) = u.c1.forward(&wav, wt);
wav = y1;
snake_beta(&mut wav, &u.a2.0, &u.a2.1, wt);
let (y2, _) = u.c2.forward(&wav, wt);
wav = y2;
for (w, r) in wav.iter_mut().zip(&residual) {
*w += *r;
}
}
}
snake_beta(&mut wav, &self.final_snake.0, &self.final_snake.1, wt);
let (y, t_out) = self.final_conv.forward(&wav, wt);
debug_assert_eq!(t_out, frames.len() * cfg::SAMPLES_PER_FRAME);
y.into_iter().map(|v| v.clamp(-1.0, 1.0)).collect()
}
}
fn elu(x: f32) -> f32 {
if x > 0.0 { x } else { x.exp() - 1.0 }
}
struct EncResBlock {
c3: Conv1d, c1: Conv1d, }
impl EncResBlock {
fn forward(&self, x: &[f32], t_len: usize) -> Vec<f32> {
let mut h: Vec<f32> = x.iter().map(|&v| elu(v)).collect();
let (y, _) = self.c3.forward(&h, t_len);
h = y.into_iter().map(elu).collect();
let (y, _) = self.c1.forward(&h, t_len);
y.iter().zip(x).map(|(a, b)| a + b).collect()
}
}
struct EncTrLayer {
ln1: (Vec<f32>, Vec<f32>),
q: Vec<f32>,
k: Vec<f32>,
v: Vec<f32>,
o: Vec<f32>,
ls1: Vec<f32>,
ln2: (Vec<f32>, Vec<f32>),
fc1: Vec<f32>,
fc2: Vec<f32>,
ls2: Vec<f32>,
}
struct RvqEnc {
in_w: Vec<f32>, codebooks: Vec<(Vec<f32>, Vec<f32>)>,
}
impl RvqEnc {
fn load(st: &St, prefix: &str, n_layers: usize) -> Result<Self> {
let dim = cfg::ENC_CODEBOOK_DIM;
let in_w = st.f32(&format!("{prefix}.input_proj.weight"))?;
anyhow::ensure!(in_w.len() == dim * cfg::ENC_HIDDEN);
let mut codebooks = Vec::with_capacity(n_layers);
for l in 0..n_layers {
let sum = st.f32(&format!("{prefix}.layers.{l}.codebook.embed_sum"))?;
let usage = st.f32(&format!("{prefix}.layers.{l}.codebook.cluster_usage"))?;
anyhow::ensure!(sum.len() == cfg::CP_VOCAB * dim && usage.len() == cfg::CP_VOCAB);
let mut eff = sum;
for (row, &u) in eff.chunks_exact_mut(dim).zip(&usage) {
let inv = 1.0 / u.max(1e-5);
for v in row {
*v *= inv;
}
}
let norms: Vec<f32> = eff
.chunks_exact(dim)
.map(|r| r.iter().map(|v| v * v).sum())
.collect();
codebooks.push((eff, norms));
}
Ok(Self { in_w, codebooks })
}
fn encode_col(&self, x: &[f32]) -> Vec<u32> {
let dim = cfg::ENC_CODEBOOK_DIM;
let mut r = matvec(&self.in_w, x, dim, cfg::ENC_HIDDEN);
let mut codes = Vec::with_capacity(self.codebooks.len());
for (cb, norms) in &self.codebooks {
let mut best = 0usize;
let mut best_d = f32::INFINITY;
for (j, (row, &n2)) in cb.chunks_exact(dim).zip(norms).enumerate() {
let dot: f32 = row.iter().zip(&r).map(|(a, b)| a * b).sum();
let d = n2 - 2.0 * dot;
if d < best_d {
best_d = d;
best = j;
}
}
codes.push(best as u32);
let row = &cb[best * dim..(best + 1) * dim];
for (rv, ev) in r.iter_mut().zip(row) {
*rv -= *ev;
}
}
codes
}
}
pub struct CodecEncoder {
conv0: Conv1d,
stages: Vec<(EncResBlock, Conv1d)>, final_conv: Conv1d,
layers: Vec<EncTrLayer>,
downsample: Conv1d,
rvq_sem: RvqEnc,
rvq_ac: RvqEnc,
}
impl CodecEncoder {
pub fn load(dir: &Path, config: &Qwen3TtsConfig) -> Result<Self> {
let e = &config.codec.encoder_config;
let st = St::open(&dir.join(cfg::DIR_CODEC).join(cfg::FILE_MODEL))?;
let nf = e.num_filters; let mut ratios: Vec<usize> = e.upsampling_ratios.clone();
ratios.reverse();
let mut stages = Vec::new();
let mut width = nf;
for (s, &r) in ratios.iter().enumerate() {
let res_idx = 1 + 3 * s;
let conv_idx = 3 + 3 * s;
let p = format!("encoder.encoder.layers.{res_idx}");
let res = EncResBlock {
c3: Conv1d::load(
&st,
&format!("{p}.block.1.conv"),
width,
width / 2,
e.residual_kernel_size,
1,
1,
1,
)?,
c1: Conv1d::load(
&st,
&format!("{p}.block.3.conv"),
width / 2,
width,
1,
1,
1,
1,
)?,
};
let strided = Conv1d::load(
&st,
&format!("encoder.encoder.layers.{conv_idx}.conv"),
width,
width * 2,
2 * r,
1,
r,
1,
)?;
stages.push((res, strided));
width *= 2;
}
let mut layers = Vec::with_capacity(e.num_hidden_layers);
let (h, ffn) = (e.hidden_size, e.intermediate_size);
for l in 0..e.num_hidden_layers {
let p = format!("encoder.encoder_transformer.layers.{l}");
layers.push(EncTrLayer {
ln1: (
st.f32(&format!("{p}.input_layernorm.weight"))?,
st.f32(&format!("{p}.input_layernorm.bias"))?,
),
q: st.mat(&format!("{p}.self_attn.q_proj.weight"), h, h)?,
k: st.mat(&format!("{p}.self_attn.k_proj.weight"), h, h)?,
v: st.mat(&format!("{p}.self_attn.v_proj.weight"), h, h)?,
o: st.mat(&format!("{p}.self_attn.o_proj.weight"), h, h)?,
ls1: st.f32(&format!("{p}.self_attn_layer_scale.scale"))?,
ln2: (
st.f32(&format!("{p}.post_attention_layernorm.weight"))?,
st.f32(&format!("{p}.post_attention_layernorm.bias"))?,
),
fc1: st.mat(&format!("{p}.mlp.fc1.weight"), ffn, h)?,
fc2: st.mat(&format!("{p}.mlp.fc2.weight"), h, ffn)?,
ls2: st.f32(&format!("{p}.mlp_layer_scale.scale"))?,
});
}
Ok(Self {
conv0: Conv1d::load(
&st,
"encoder.encoder.layers.0.conv",
1,
nf,
e.kernel_size,
1,
1,
1,
)?,
stages,
final_conv: Conv1d::load(
&st,
"encoder.encoder.layers.14.conv",
width,
h,
e.last_kernel_size,
1,
1,
1,
)?,
layers,
downsample: Conv1d::load_no_bias(&st, "encoder.downsample.conv", h, h, 4, 1, 2, 1)?,
rvq_sem: RvqEnc::load(
&st,
"encoder.quantizer.semantic_residual_vector_quantizer",
1,
)?,
rvq_ac: RvqEnc::load(
&st,
"encoder.quantizer.acoustic_residual_vector_quantizer",
cfg::ENC_QUANTIZERS_VALID - 1,
)?,
})
}
fn transformer(&self, x: &mut [f32], t_len: usize) {
let h = cfg::ENC_HIDDEN;
let (heads, hd) = (cfg::ENC_HEADS, h / cfg::ENC_HEADS);
let scale = 1.0 / (hd as f32).sqrt();
let win = cfg::ENC_SLIDING_WINDOW;
let ln = |row: &[f32], w: &[f32], b: &[f32]| -> Vec<f32> {
let mean = row.iter().sum::<f32>() / h as f32;
let var = row.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / h as f32;
let inv = 1.0 / (var + 1e-5).sqrt();
row.iter()
.enumerate()
.map(|(i, v)| (*v - mean) * inv * w[i] + b[i])
.collect()
};
let tables: Vec<(Vec<f32>, Vec<f32>)> = (0..t_len)
.map(|t| crate::qwen3tts::rope_cs_1d(hd, cfg::DEC_ROPE_THETA, t as u32))
.collect();
for layer in &self.layers {
let mut q = vec![0f32; t_len * h];
let mut k = vec![0f32; t_len * h];
let mut v = vec![0f32; t_len * h];
for t in 0..t_len {
let xn = ln(&x[t * h..(t + 1) * h], &layer.ln1.0, &layer.ln1.1);
q[t * h..(t + 1) * h].copy_from_slice(&matvec(&layer.q, &xn, h, h));
k[t * h..(t + 1) * h].copy_from_slice(&matvec(&layer.k, &xn, h, h));
v[t * h..(t + 1) * h].copy_from_slice(&matvec(&layer.v, &xn, h, h));
let (cos, sin) = &tables[t];
crate::qwen3tts::apply_rope(&mut q[t * h..(t + 1) * h], cos, sin, heads, hd);
crate::qwen3tts::apply_rope(&mut k[t * h..(t + 1) * h], cos, sin, heads, hd);
}
for t in 0..t_len {
let lo = t.saturating_sub(win - 1);
let mut attn = vec![0f32; h];
for head in 0..heads {
let qh = &q[t * h + head * hd..t * h + (head + 1) * hd];
let mut scores: Vec<f32> = (lo..=t)
.map(|s| {
let kh = &k[s * h + head * hd..s * h + (head + 1) * hd];
qh.iter().zip(kh).map(|(a, b)| a * b).sum::<f32>() * scale
})
.collect();
softmax_inplace(&mut scores);
let out = &mut attn[head * hd..(head + 1) * hd];
for (idx, &w) in scores.iter().enumerate() {
let s = lo + idx;
let vh = &v[s * h + head * hd..s * h + (head + 1) * hd];
for (d, ov) in out.iter_mut().enumerate() {
*ov += w * vh[d];
}
}
}
let o = matvec(&layer.o, &attn, h, h);
for (j, (ov, sc)) in o.iter().zip(&layer.ls1).enumerate() {
x[t * h + j] += ov * sc;
}
let xn = ln(&x[t * h..(t + 1) * h], &layer.ln2.0, &layer.ln2.1);
let mut ff = matvec(&layer.fc1, &xn, self.layers[0].fc1.len() / h, h);
for f in ff.iter_mut() {
*f = gelu_erf(*f);
}
let d = matvec(&layer.fc2, &ff, h, ff.len());
for (j, (dv, sc)) in d.iter().zip(&layer.ls2).enumerate() {
x[t * h + j] += dv * sc;
}
}
}
}
pub fn encode(&self, pcm: &[f32]) -> Vec<[u32; cfg::NUM_CODE_GROUPS]> {
if pcm.is_empty() {
return Vec::new();
}
let (mut x, mut t) = self.conv0.forward(pcm, pcm.len());
for (res, strided) in &self.stages {
let y: Vec<f32> = res.forward(&x, t).into_iter().map(elu).collect();
let (y2, t2) = strided.forward(&y, t);
x = y2;
t = t2;
}
let x: Vec<f32> = x.into_iter().map(elu).collect();
let (fx, ft) = self.final_conv.forward(&x, t);
let h = cfg::ENC_HIDDEN;
let mut rows = vec![0f32; ft * h];
for c in 0..h {
for ti in 0..ft {
rows[ti * h + c] = fx[c * ft + ti];
}
}
self.transformer(&mut rows, ft);
let mut cm = vec![0f32; h * ft];
for ti in 0..ft {
for c in 0..h {
cm[c * ft + ti] = rows[ti * h + c];
}
}
let (dx, dt) = self.downsample.forward(&cm, ft);
let mut out = Vec::with_capacity(dt);
for ti in 0..dt {
let col: Vec<f32> = (0..h).map(|c| dx[c * dt + ti]).collect();
let sem = self.rvq_sem.encode_col(&col);
let ac = self.rvq_ac.encode_col(&col);
let mut f = [0u32; cfg::NUM_CODE_GROUPS];
f[0] = sem[0];
f[1..].copy_from_slice(&ac);
out.push(f);
}
out
}
}
struct TrRing {
k: std::collections::VecDeque<Vec<f32>>,
v: std::collections::VecDeque<Vec<f32>>,
}
struct PreTrStream {
rings: Vec<TrRing>,
pos: usize,
}
impl PreTransformer {
fn stream_state(&self) -> PreTrStream {
PreTrStream {
rings: (0..self.layers.len())
.map(|_| TrRing {
k: Default::default(),
v: Default::default(),
})
.collect(),
pos: 0,
}
}
fn forward_stream(&self, st: &mut PreTrStream, row: &[f32]) -> Vec<f32> {
let (h, hd, heads) = (cfg::DEC_HIDDEN, cfg::DEC_HEAD_DIM, cfg::DEC_HEADS);
let aw = heads * hd;
let eps = cfg::DEC_RMS_EPS;
let scale = 1.0 / (hd as f32).sqrt();
let win = cfg::DEC_SLIDING_WINDOW;
let mut x = matvec(&self.in_w, row, h, cfg::DEC_LATENT);
for (v, b) in x.iter_mut().zip(&self.in_b) {
*v += *b;
}
let (cos, sin) = crate::qwen3tts::rope_cs_1d(hd, cfg::DEC_ROPE_THETA, st.pos as u32);
for (layer, ring) in self.layers.iter().zip(&mut st.rings) {
let xn = rms_norm(&x, &layer.input_ln, eps);
let mut q = matvec(&layer.q, &xn, aw, h);
let mut k = matvec(&layer.k, &xn, aw, h);
let v = matvec(&layer.v, &xn, aw, h);
crate::qwen3tts::apply_rope(&mut q, &cos, &sin, heads, hd);
crate::qwen3tts::apply_rope(&mut k, &cos, &sin, heads, hd);
ring.k.push_back(k);
ring.v.push_back(v);
while ring.k.len() > win {
ring.k.pop_front();
ring.v.pop_front();
}
let t_ctx = ring.k.len();
let mut attn = vec![0f32; aw];
for head in 0..heads {
let qh = &q[head * hd..(head + 1) * hd];
let mut scores: Vec<f32> = (0..t_ctx)
.map(|s| {
let kh = &ring.k[s][head * hd..(head + 1) * hd];
qh.iter().zip(kh).map(|(a, b)| a * b).sum::<f32>() * scale
})
.collect();
softmax_inplace(&mut scores);
let out = &mut attn[head * hd..(head + 1) * hd];
for (s, &w) in scores.iter().enumerate() {
let vh = &ring.v[s][head * hd..(head + 1) * hd];
for (d, ov) in out.iter_mut().enumerate() {
*ov += w * vh[d];
}
}
}
let o = matvec(&layer.o, &attn, h, aw);
for (j, (ov, sc)) in o.iter().zip(&layer.attn_scale).enumerate() {
x[j] += ov * sc;
}
let xn = rms_norm(&x, &layer.post_ln, eps);
let g = matvec(&layer.gate, &xn, cfg::DEC_LATENT, h);
let u = matvec(&layer.up, &xn, cfg::DEC_LATENT, h);
let act: Vec<f32> = g.iter().zip(&u).map(|(gv, uv)| silu(*gv) * uv).collect();
let d = matvec(&layer.down, &act, h, cfg::DEC_LATENT);
for (j, (dv, sc)) in d.iter().zip(&layer.mlp_scale).enumerate() {
x[j] += dv * sc;
}
}
st.pos += 1;
let xn = rms_norm(&x, &self.norm, eps);
let mut out = matvec(&self.out_w, &xn, cfg::DEC_LATENT, h);
for (v, b) in out.iter_mut().zip(&self.out_b) {
*v += *b;
}
out
}
}
impl ConvNeXt {
fn forward_stream(&self, dw_state: &mut Vec<f32>, x: &[f32], t_len: usize) -> Vec<f32> {
let dim = self.gamma.len();
let dwo = self.dw.forward_stream(dw_state, x, t_len);
let mut out = x.to_vec();
for t in 0..t_len {
let mut col: Vec<f32> = (0..dim).map(|c| dwo[c * t_len + t]).collect();
let mean = col.iter().sum::<f32>() / dim as f32;
let var = col.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / dim as f32;
let inv = 1.0 / (var + 1e-6).sqrt();
for (v, (w, b)) in col.iter_mut().zip(self.norm_w.iter().zip(&self.norm_b)) {
*v = (*v - mean) * inv * w + b;
}
let mut a = matvec(&self.pw1_w, &col, 4 * dim, dim);
for (v, b) in a.iter_mut().zip(&self.pw1_b) {
*v = gelu_erf(*v + *b);
}
let mut y = matvec(&self.pw2_w, &a, dim, 4 * dim);
for ((v, b), g) in y.iter_mut().zip(&self.pw2_b).zip(&self.gamma) {
*v = (*v + *b) * g;
}
for (c, v) in y.iter().enumerate() {
out[c * t_len + t] += *v;
}
}
out
}
}
pub struct CodecDecoderStream<'a> {
d: &'a CodecDecoder,
pre_conv_st: Vec<f32>,
tr_st: PreTrStream,
up_tr_carry: Vec<Vec<f32>>, up_cnx_st: Vec<Vec<f32>>, dec0_st: Vec<f32>,
blk_tr_carry: Vec<Vec<f32>>, blk_c1_st: Vec<Vec<f32>>, blk_c2_st: Vec<Vec<f32>>, final_st: Vec<f32>,
}
impl CodecDecoder {
pub fn stream(&self) -> CodecDecoderStream<'_> {
let n_units = self.blocks.len() * 3;
CodecDecoderStream {
d: self,
pre_conv_st: Vec::new(),
tr_st: self.pre_tr.stream_state(),
up_tr_carry: vec![Vec::new(); self.upsample.len()],
up_cnx_st: vec![Vec::new(); self.upsample.len()],
dec0_st: Vec::new(),
blk_tr_carry: vec![Vec::new(); self.blocks.len()],
blk_c1_st: vec![Vec::new(); n_units],
blk_c2_st: vec![Vec::new(); n_units],
final_st: Vec::new(),
}
}
}
impl CodecDecoderStream<'_> {
pub fn decode_frame(&mut self, codes: &[u32; cfg::NUM_CODE_GROUPS]) -> Vec<f32> {
let d = self.d;
let sem: Vec<Vec<u32>> = vec![vec![codes[0]]];
let ac: Vec<Vec<u32>> = (1..cfg::NUM_CODE_GROUPS).map(|g| vec![codes[g]]).collect();
let mut latent = d.rvq_first.decode(&sem, 1);
let rest = d.rvq_rest.decode(&ac, 1);
for (a, b) in latent.iter_mut().zip(&rest) {
*a += *b;
}
let pc = d.pre_conv.forward_stream(&mut self.pre_conv_st, &latent, 1);
let row: Vec<f32> = pc.clone();
let tr_out = d.pre_tr.forward_stream(&mut self.tr_st, &row);
let mut x = tr_out;
let mut cur_t = 1usize;
for (s, (tr, cnx)) in d.upsample.iter().enumerate() {
let y = tr.forward_stream(&mut self.up_tr_carry[s], &x, cur_t);
cur_t *= tr.stride;
x = cnx.forward_stream(&mut self.up_cnx_st[s], &y, cur_t);
}
let mut wav = d.dec0.forward_stream(&mut self.dec0_st, &x, cur_t);
let mut wt = cur_t;
for (bi, b) in d.blocks.iter().enumerate() {
snake_beta(&mut wav, &b.snake.0, &b.snake.1, wt);
wav = b.tr.forward_stream(&mut self.blk_tr_carry[bi], &wav, wt);
wt *= b.tr.stride;
for (ui, u) in b.units.iter().enumerate() {
let residual = wav.clone();
snake_beta(&mut wav, &u.a1.0, &u.a1.1, wt);
wav =
u.c1.forward_stream(&mut self.blk_c1_st[bi * 3 + ui], &wav, wt);
snake_beta(&mut wav, &u.a2.0, &u.a2.1, wt);
wav =
u.c2.forward_stream(&mut self.blk_c2_st[bi * 3 + ui], &wav, wt);
for (w, r) in wav.iter_mut().zip(&residual) {
*w += *r;
}
}
}
snake_beta(&mut wav, &d.final_snake.0, &d.final_snake.1, wt);
let y = d.final_conv.forward_stream(&mut self.final_st, &wav, wt);
debug_assert_eq!(y.len(), cfg::SAMPLES_PER_FRAME);
y.into_iter().map(|v| v.clamp(-1.0, 1.0)).collect()
}
}