inferencelayer 0.2.1

Kortexya's engine-native inference layer — LLM generation + embedding/encoder family on wgpu (WGSL kernels, any adapter) with a pure-Rust CPU fallback
Documentation
//! Qwen3-TTS speaker encoder — P6: reference audio → the x-vector that conditions the talker.
//!
//! ECAPA-TDNN, transcribed from `Qwen3TTSSpeakerEncoder` in the reference `modeling_qwen3_tts.py`
//! (2026-07-20). Distinctive facts pinned from source:
//!  - **no BatchNorm anywhere** (the checkpoint has zero BN tensors — TDNN block = Conv1d with
//!    `padding="same", padding_mode="reflect"` + ReLU);
//!  - mel frontend is the HiFi-GAN `mel_spectrogram`: 24 kHz, n_fft 1024, hop 256, win 1024
//!    (periodic Hann), `center=False` with `(n_fft−hop)/2 = 384` REFLECT pad both sides, slaney
//!    mel 128 bins fmin 0 / fmax 12000, `log(clamp(·, 1e-5))`;
//!  - ASP context = `cat[x, mean, std]` over channels (global UNWEIGHTED mean/std broadcast),
//!    then TDNN(4608→128) → tanh → conv k1 128→1536 → softmax over time → weighted mean‖std;
//!  - final `fc` conv k1 3072→2048; **no L2 norm** on the output.
//!
//! One micro-detail is NOT in the checkpoint or config: the SE-Res2Net dilations. Standard
//! ECAPA-TDNN uses (2, 3, 4) with kernel 3, scale 8 — adopted here and recorded as an assumption
//! in HANDOFF_qwen3tts.md §5; the P7 e2e speaker-similarity gate is the backstop (a wrong
//! dilation degrades cloning similarity, which that gate measures).
//!
//! Gating (§2): layer-2 self-consistency (determinism, distinct-speaker separation, finiteness)
//! — an honest gap until P7's e2e clone-similarity backstop.

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; // block width
const SCALE: usize = 8; // res2net scale → 64-ch sub-groups
const MFA: usize = 1536;

fn relu(x: f32) -> f32 {
    x.max(0.0)
}

/// Conv1d with torch `padding="same", padding_mode="reflect"` over channel-major `[C][T]`.
struct SameConv {
    w: Vec<f32>, // [c_out, c_in, k]
    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; // total padding (odd eff_k ⇒ symmetric halves)
        let left = eff / 2;
        // reflect index helper (torch reflect: mirror without repeating the edge sample)
        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
    }
}

/// SE-Res2Net block: tdnn1(k1) → res2net(k3, dil d, scale 8) → tdnn2(k1) → SE gate → +residual.
struct SeRes2Net {
    tdnn1: SameConv,
    res2: Vec<SameConv>, // 7 convs @64ch
    tdnn2: SameConv,
    se1: SameConv, // 512→128 k1
    se2: SameConv, // 128→512 k1
}

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();
        // res2net: part 0 identity; part i (1-indexed): out_i = tdnn(x_i + out_{i-1}).
        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();
        // SE gate: global average over time → conv/ReLU/conv/sigmoid → channel scale.
        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>, // [128, 513] slaney
    hann: Vec<f32>,   // periodic, 1024
    stem: SameConv,
    blocks: Vec<SeRes2Net>,
    mfa: SameConv,
    asp_tdnn: SameConv, // 4608→128 k1 (+ReLU inside TDNN, then tanh)
    asp_conv: SameConv, // 128→1536 k1
    fc: SameConv,       // 3072→2048 k1
    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() // periodic hann: sin²(πn/N)
            })
            .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,
        })
    }

    /// HiFi-GAN mel: reflect-pad 384, hann-1024 STFT hop 256 center=false, |·|, slaney mel,
    /// log-clamp. Returns `[128][frames]` channel-major.
    fn mel(&self, pcm: &[f32]) -> (Vec<f32>, usize) {
        let pad = (N_FFT - HOP) / 2;
        let n = pcm.len();
        let get = |i: isize| -> f32 {
            // torch reflect pad on the sample axis
            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();
            }
        }
        // mel: [128, 513] · mag[f] → log(clamp(., 1e-5)); output channel-major [128][frames]
        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)
    }

    /// 24 kHz mono reference → x-vector `[enc_dim]` (no L2 norm — raw fc output, as in reference).
    pub fn embed(&self, pcm: &[f32]) -> Result<Vec<f32>> {
        anyhow::ensure!(
            // The dilation-4 res2net conv needs ≥ 5 mel frames; below that, torch reflect-pad
            // ERRORS while our clamped reflect would silently produce a garbage x-vector
            // (review finding) — so the floor is the real minimum, not one frame.
            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);
        // MFA input: cat of the three block outputs (3×512 = 1536).
        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();
        // ASP: context = cat[x, mean, std] (global unweighted stats broadcast over time).
        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); // [1536][t] logits
        for c in 0..MFA {
            crate::qwen3tts::softmax_inplace(&mut attn[c * t..(c + 1) * t]);
        }
        // weighted mean ‖ weighted std per channel.
        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
    }
}