use crate::qwen3tts::chatterbox_st::St;
use crate::qwen3tts::{Qwen3TtsConfig, cfg, matvec};
use anyhow::Result;
use std::path::Path;
const N_MELS: usize = 128;
const N_FFT: usize = 1024;
const HOP: usize = 256;
const CH: usize = 512; const SCALE: usize = 8; const MFA: usize = 1536;
fn relu(x: f32) -> f32 {
x.max(0.0)
}
struct SameConv {
w: Vec<f32>, b: Vec<f32>,
c_in: usize,
c_out: usize,
k: usize,
dilation: usize,
}
impl SameConv {
fn load(
st: &St,
prefix: &str,
c_in: usize,
c_out: usize,
k: usize,
dilation: usize,
) -> Result<Self> {
let w = st.f32(&format!("{prefix}.weight"))?;
anyhow::ensure!(
w.len() == c_out * c_in * k,
"{prefix}: weight len {} != {c_out}×{c_in}×{k}",
w.len()
);
Ok(Self {
w,
b: st.f32(&format!("{prefix}.bias"))?,
c_in,
c_out,
k,
dilation,
})
}
fn forward(&self, x: &[f32], t_len: usize) -> Vec<f32> {
let eff = (self.k - 1) * self.dilation; let left = eff / 2;
let reflect = |i: isize| -> usize {
let n = t_len as isize;
let mut j = i;
if j < 0 {
j = -j;
}
if j >= n {
j = 2 * (n - 1) - j;
}
j.clamp(0, n - 1) as usize
};
let mut out = vec![0f32; self.c_out * t_len];
for co in 0..self.c_out {
let orow = &mut out[co * t_len..(co + 1) * t_len];
for ci in 0..self.c_in {
let xrow = &x[ci * t_len..(ci + 1) * t_len];
let wrow =
&self.w[(co * self.c_in + ci) * self.k..(co * self.c_in + ci + 1) * self.k];
for (t, ov) in orow.iter_mut().enumerate() {
let mut acc = 0f32;
for (j, &wv) in wrow.iter().enumerate() {
let p = t as isize + (j * self.dilation) as isize - left as isize;
acc += wv * xrow[reflect(p)];
}
*ov += acc;
}
}
for v in orow.iter_mut() {
*v += self.b[co];
}
}
out
}
}
struct SeRes2Net {
tdnn1: SameConv,
res2: Vec<SameConv>, tdnn2: SameConv,
se1: SameConv, se2: SameConv, }
impl SeRes2Net {
fn load(st: &St, prefix: &str, dilation: usize) -> Result<Self> {
let sub = CH / SCALE;
let mut res2 = Vec::with_capacity(SCALE - 1);
for i in 0..SCALE - 1 {
res2.push(SameConv::load(
st,
&format!("{prefix}.res2net_block.blocks.{i}.conv"),
sub,
sub,
3,
dilation,
)?);
}
Ok(Self {
tdnn1: SameConv::load(st, &format!("{prefix}.tdnn1.conv"), CH, CH, 1, 1)?,
res2,
tdnn2: SameConv::load(st, &format!("{prefix}.tdnn2.conv"), CH, CH, 1, 1)?,
se1: SameConv::load(st, &format!("{prefix}.se_block.conv1"), CH, CH / 4, 1, 1)?,
se2: SameConv::load(st, &format!("{prefix}.se_block.conv2"), CH / 4, CH, 1, 1)?,
})
}
fn forward(&self, x: &[f32], t_len: usize) -> Vec<f32> {
let sub = CH / SCALE;
let mut h: Vec<f32> = self.tdnn1.forward(x, t_len).into_iter().map(relu).collect();
let mut prev: Vec<f32> = Vec::new();
for i in 1..SCALE {
let seg = &mut h[i * sub * t_len..(i + 1) * sub * t_len];
if i > 1 {
for (s, p) in seg.iter_mut().zip(&prev) {
*s += *p;
}
}
let out: Vec<f32> = self.res2[i - 1]
.forward(seg, t_len)
.into_iter()
.map(relu)
.collect();
seg.copy_from_slice(&out);
prev = out;
}
let h: Vec<f32> = self
.tdnn2
.forward(&h, t_len)
.into_iter()
.map(relu)
.collect();
let mut pooled = vec![0f32; CH];
for c in 0..CH {
pooled[c] = h[c * t_len..(c + 1) * t_len].iter().sum::<f32>() / t_len as f32;
}
let s1: Vec<f32> = self.se1.forward(&pooled, 1).into_iter().map(relu).collect();
let gate: Vec<f32> = self
.se2
.forward(&s1, 1)
.into_iter()
.map(|v| 1.0 / (1.0 + (-v).exp()))
.collect();
let mut out = vec![0f32; CH * t_len];
for c in 0..CH {
for t in 0..t_len {
out[c * t_len + t] = h[c * t_len + t] * gate[c] + x[c * t_len + t];
}
}
out
}
}
pub struct SpeakerEncoder {
mel_fb: Vec<f32>, hann: Vec<f32>, stem: SameConv,
blocks: Vec<SeRes2Net>,
mfa: SameConv,
asp_tdnn: SameConv, asp_conv: SameConv, fc: SameConv, enc_dim: usize,
}
impl SpeakerEncoder {
pub fn load(dir: &Path, config: &Qwen3TtsConfig) -> Result<Self> {
let enc_dim = config
.top
.speaker_encoder_config
.as_ref()
.map(|s| s.enc_dim)
.unwrap_or(2048);
let st = St::open(&dir.join(cfg::FILE_MODEL))?;
let hann: Vec<f32> = (0..N_FFT)
.map(|i| {
let x = std::f32::consts::PI * i as f32 / N_FFT as f32;
x.sin() * x.sin() })
.collect();
Ok(Self {
mel_fb: crate::chatterbox::librosa_mel(24_000.0, N_FFT, N_MELS, 0.0, 12_000.0),
hann,
stem: SameConv::load(&st, "speaker_encoder.blocks.0.conv", N_MELS, CH, 5, 1)?,
blocks: vec![
SeRes2Net::load(&st, "speaker_encoder.blocks.1", 2)?,
SeRes2Net::load(&st, "speaker_encoder.blocks.2", 3)?,
SeRes2Net::load(&st, "speaker_encoder.blocks.3", 4)?,
],
mfa: SameConv::load(&st, "speaker_encoder.mfa.conv", MFA, MFA, 1, 1)?,
asp_tdnn: SameConv::load(&st, "speaker_encoder.asp.tdnn.conv", 3 * MFA, 128, 1, 1)?,
asp_conv: SameConv::load(&st, "speaker_encoder.asp.conv", 128, MFA, 1, 1)?,
fc: SameConv::load(&st, "speaker_encoder.fc", 2 * MFA, enc_dim, 1, 1)?,
enc_dim,
})
}
fn mel(&self, pcm: &[f32]) -> (Vec<f32>, usize) {
let pad = (N_FFT - HOP) / 2;
let n = pcm.len();
let get = |i: isize| -> f32 {
let mut j = i;
if j < 0 {
j = -j;
}
let nn = n as isize;
if j >= nn {
j = 2 * (nn - 1) - j;
}
pcm[j.clamp(0, nn - 1) as usize]
};
let padded_len = n + 2 * pad;
let frames = if padded_len >= N_FFT {
(padded_len - N_FFT) / HOP + 1
} else {
0
};
let n_freq = N_FFT / 2 + 1;
let mut mag = vec![0f32; frames * n_freq];
for f in 0..frames {
let start = f as isize * HOP as isize - pad as isize;
let frame: Vec<f32> = (0..N_FFT)
.map(|i| get(start + i as isize) * self.hann[i])
.collect();
for k in 0..n_freq {
let mut re = 0f32;
let mut im = 0f32;
for (i, &v) in frame.iter().enumerate() {
let a = -2.0 * std::f32::consts::PI * (k * i) as f32 / N_FFT as f32;
re += v * a.cos();
im += v * a.sin();
}
mag[f * n_freq + k] = (re * re + im * im + 1e-9).sqrt();
}
}
let mut out = vec![0f32; N_MELS * frames];
for f in 0..frames {
let m = matvec(
&self.mel_fb,
&mag[f * n_freq..(f + 1) * n_freq],
N_MELS,
n_freq,
);
for (c, v) in m.iter().enumerate() {
out[c * frames + f] = v.max(1e-5).ln();
}
}
(out, frames)
}
pub fn embed(&self, pcm: &[f32]) -> Result<Vec<f32>> {
anyhow::ensure!(
pcm.len() >= N_FFT + 4 * HOP,
"reference audio too short for the speaker encoder (need ≥ {} samples)",
N_FFT + 4 * HOP
);
let (mel, t) = self.mel(pcm);
let x: Vec<f32> = self.stem.forward(&mel, t).into_iter().map(relu).collect();
let b1 = self.blocks[0].forward(&x, t);
let b2 = self.blocks[1].forward(&b1, t);
let b3 = self.blocks[2].forward(&b2, t);
let mut cat = vec![0f32; 3 * CH * t];
cat[..CH * t].copy_from_slice(&b1);
cat[CH * t..2 * CH * t].copy_from_slice(&b2);
cat[2 * CH * t..].copy_from_slice(&b3);
let m: Vec<f32> = self.mfa.forward(&cat, t).into_iter().map(relu).collect();
let mut ctx = vec![0f32; 3 * MFA * t];
ctx[..MFA * t].copy_from_slice(&m);
for c in 0..MFA {
let row = &m[c * t..(c + 1) * t];
let mean = row.iter().sum::<f32>() / t as f32;
let var = row.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / t as f32;
let sd = var.max(1e-12).sqrt();
for ti in 0..t {
ctx[(MFA + c) * t + ti] = mean;
ctx[(2 * MFA + c) * t + ti] = sd;
}
}
let a: Vec<f32> = self
.asp_tdnn
.forward(&ctx, t)
.into_iter()
.map(|v| relu(v).tanh())
.collect();
let mut attn = self.asp_conv.forward(&a, t); for c in 0..MFA {
crate::qwen3tts::softmax_inplace(&mut attn[c * t..(c + 1) * t]);
}
let mut stats = vec![0f32; 2 * MFA];
for c in 0..MFA {
let (xr, ar) = (&m[c * t..(c + 1) * t], &attn[c * t..(c + 1) * t]);
let mean: f32 = xr.iter().zip(ar).map(|(x, w)| x * w).sum();
let ex2: f32 = xr.iter().zip(ar).map(|(x, w)| x * x * w).sum();
stats[c] = mean;
stats[MFA + c] = (ex2 - mean * mean).max(1e-12).sqrt();
}
let out = self.fc.forward(&stats, 1);
debug_assert_eq!(out.len(), self.enc_dim);
Ok(out)
}
pub fn enc_dim(&self) -> usize {
self.enc_dim
}
}