use std::path::Path;
use anyhow::{Context, Result};
use crate::cpu_gemm::{PackedWeight, gemm_packed};
use crate::weights::LazySt;
fn sigmoid(v: f32) -> f32 {
1.0 / (1.0 + (-v).exp())
}
struct Linear {
w: Vec<f32>,
b: Vec<f32>,
n: usize,
k: usize,
packed: std::sync::OnceLock<PackedWeight>,
}
impl Linear {
fn new(w: Vec<f32>, b: Vec<f32>, n: usize, k: usize) -> Self {
Self {
w,
b,
n,
k,
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
}
}
const SR: usize = 16000;
const FRAME_LEN: usize = 400; const FRAME_SHIFT: usize = 160; const N_FFT: usize = 512;
const N_MELS: usize = 80;
pub struct Fbank {
window: Vec<f32>, dft_cos: Vec<f32>, dft_sin: Vec<f32>,
mel: Vec<(usize, Vec<f32>)>, }
fn mel_scale_slaney(f: f64) -> f64 {
if f < 1000.0 {
f * 3.0 / 200.0
} else {
15.0 + (f / 1000.0).ln() * 27.0 / 6.4f64.ln()
}
}
fn inverse_mel_scale_slaney(m: f64) -> f64 {
if m < 15.0 {
m * 200.0 / 3.0
} else {
1000.0 * ((6.4f64.ln() / 27.0) * (m - 15.0)).exp()
}
}
impl Default for Fbank {
fn default() -> Self {
Self::new()
}
}
impl Fbank {
pub fn new() -> Self {
let window: Vec<f32> = (0..FRAME_LEN)
.map(|i| {
(0.5 - 0.5 * (2.0 * std::f64::consts::PI * i as f64 / FRAME_LEN as f64).cos())
as f32
})
.collect();
let bins = N_FFT / 2 + 1;
let mut dft_cos = vec![0f32; bins * FRAME_LEN];
let mut dft_sin = vec![0f32; bins * FRAME_LEN];
for b in 0..bins {
for j in 0..FRAME_LEN {
let ang = 2.0 * std::f64::consts::PI * b as f64 * j as f64 / N_FFT as f64;
dft_cos[b * FRAME_LEN + j] = ang.cos() as f32;
dft_sin[b * FRAME_LEN + j] = ang.sin() as f32;
}
}
let (low_hz, high_hz) = (0.0f64, (SR as f64 / 2.0) - 400.0);
let mel_low = mel_scale_slaney(low_hz);
let mel_delta = (mel_scale_slaney(high_hz) - mel_low) / (N_MELS as f64 + 1.0);
let fft_bin_width = SR as f64 / N_FFT as f64;
let mut mel = Vec::with_capacity(N_MELS);
for bin in 0..N_MELS {
let left = inverse_mel_scale_slaney(mel_low + bin as f64 * mel_delta) as f32;
let center = inverse_mel_scale_slaney(mel_low + (bin as f64 + 1.0) * mel_delta) as f32;
let right = inverse_mel_scale_slaney(mel_low + (bin as f64 + 2.0) * mel_delta) as f32;
let mut first = None;
let mut weights = Vec::new();
for i in 0..=(N_FFT / 2) {
let hz = (fft_bin_width * i as f64) as f32;
if hz > left && hz < right {
let mut w = if hz <= center {
(hz - left) / (center - left)
} else {
(right - hz) / (right - center)
};
w *= 2.0 / (right - left); if first.is_none() {
first = Some(i);
}
weights.push(w);
} else if first.is_some() {
break;
}
}
mel.push((first.unwrap_or(0), weights));
}
Self {
window,
dft_cos,
dft_sin,
mel,
}
}
pub fn compute(&self, samples: &[f32]) -> (Vec<f32>, usize) {
if samples.len() < FRAME_LEN {
return (Vec::new(), 0);
}
let frames = 1 + (samples.len() - FRAME_LEN) / FRAME_SHIFT;
let bins = N_FFT / 2 + 1;
let mut out = vec![0f32; frames * N_MELS];
let mut frame = vec![0f32; FRAME_LEN];
for t in 0..frames {
frame.copy_from_slice(&samples[t * FRAME_SHIFT..][..FRAME_LEN]);
for i in (1..FRAME_LEN).rev() {
frame[i] -= 0.97 * frame[i - 1];
}
frame[0] -= 0.97 * frame[0];
for i in 0..FRAME_LEN {
frame[i] *= self.window[i];
}
let mut spec = vec![0f32; bins];
for b in 0..bins {
let cos = &self.dft_cos[b * FRAME_LEN..][..FRAME_LEN];
let sin = &self.dft_sin[b * FRAME_LEN..][..FRAME_LEN];
let mut re = 0f64;
let mut im = 0f64;
for j in 0..FRAME_LEN {
re += frame[j] as f64 * cos[j] as f64;
im -= frame[j] as f64 * sin[j] as f64;
}
spec[b] = (re * re + im * im) as f32;
}
let orow = &mut out[t * N_MELS..][..N_MELS];
for (m, (first, weights)) in self.mel.iter().enumerate() {
let mut acc = 0f32;
for (i, w) in weights.iter().enumerate() {
acc += w * spec[first + i];
}
orow[m] = acc.max(f32::EPSILON).ln();
}
}
(out, frames)
}
}
fn normalize_per_feature(feats: &mut [f32], frames: usize) {
let d = N_MELS;
for c in 0..d {
let mut sum = 0f32;
let mut sum2 = 0f32;
for t in 0..frames {
let v = feats[t * d + c];
sum += v;
sum2 += v * v;
}
let ex = sum / frames as f32;
let ex2 = sum2 / frames as f32;
let var = (ex2 - ex * ex).max(1e-5);
let inv = 1.0 / (var.sqrt() + 1e-5);
for t in 0..frames {
feats[t * d + c] = (feats[t * d + c] - ex) * inv;
}
}
}
struct DwConv {
w: Vec<f32>,
c: usize,
k: usize,
}
impl DwConv {
fn forward(&self, x: &[f32], t: usize) -> Vec<f32> {
let (c, k) = (self.c, self.k);
let pad = (k - 1) / 2;
let mut out = vec![0f32; t * c];
for i in 0..t {
let orow = &mut out[i * c..][..c];
for tap in 0..k {
let j = i as isize + tap as isize - pad as isize;
if j < 0 || j as usize >= t {
continue;
}
let xrow = &x[(j as usize) * c..][..c];
for ch in 0..c {
orow[ch] += self.w[ch * k + tap] * xrow[ch];
}
}
}
out
}
}
struct Se {
down: Linear,
up: Linear,
}
impl Se {
fn apply(&self, x: &mut [f32], t: usize, c: usize) {
let mut mean = vec![0f32; c];
for row in x.chunks_exact(c) {
for (m, v) in mean.iter_mut().zip(row) {
*m += v;
}
}
for m in mean.iter_mut() {
*m /= t as f32;
}
let mut h = self.down.forward(&mean);
for v in h.iter_mut() {
*v = v.max(0.0);
}
let gate = self.up.forward(&h);
for row in x.chunks_exact_mut(c) {
for (v, g) in row.iter_mut().zip(&gate) {
*v *= sigmoid(*g);
}
}
}
}
struct TitaBlock {
subs: Vec<(DwConv, Linear)>, se: Se,
res: Option<Linear>, }
impl TitaBlock {
fn forward(&self, x: &[f32], t: usize) -> Vec<f32> {
let mut h = x.to_vec();
for (i, (dw, pw)) in self.subs.iter().enumerate() {
h = dw.forward(&h, t);
h = pw.forward(&h);
if i + 1 < self.subs.len() {
for v in h.iter_mut() {
*v = v.max(0.0);
}
}
}
self.se.apply(&mut h, t, self.subs.last().unwrap().1.n);
if let Some(res) = &self.res {
let r = res.forward(x);
for (a, b) in h.iter_mut().zip(r) {
*a += b;
}
}
for v in h.iter_mut() {
*v = v.max(0.0);
}
h
}
}
struct BnFold {
scale: Vec<f32>,
shift: Vec<f32>,
}
impl BnFold {
fn load(st: &LazySt, prefix: &str) -> Result<Self> {
let g = st.tensor_f32(&format!("{prefix}.weight"))?;
let b = st.tensor_f32(&format!("{prefix}.bias"))?;
let mean = st.tensor_f32(&format!("{prefix}.running_mean"))?;
let var = st.tensor_f32(&format!("{prefix}.running_var"))?;
let scale: Vec<f32> = (0..g.len())
.map(|i| g[i] / (var[i] + 1e-5).sqrt())
.collect();
let shift: Vec<f32> = (0..g.len()).map(|i| b[i] - mean[i] * scale[i]).collect();
Ok(Self { scale, shift })
}
fn apply(&self, x: &mut [f32]) {
let c = self.scale.len();
for row in x.chunks_exact_mut(c) {
for i in 0..c {
row[i] = row[i] * self.scale[i] + self.shift[i];
}
}
}
}
pub struct TitaNet {
fbank: Fbank,
blocks: Vec<TitaBlock>,
attn_pre: Linear, attn_bn: BnFold,
attn_out: Linear, emb_bn: BnFold, emb_out: Linear, }
impl TitaNet {
fn load(st: &LazySt) -> Result<Self> {
let lin = |wname: &str, bname: Option<&str>, n: usize, k: usize| -> Result<Linear> {
let w = st.tensor_f32(wname)?;
anyhow::ensure!(w.len() == n * k, "{wname}: {} != {n}·{k}", w.len());
let b = match bname {
Some(b) => st.tensor_f32(b)?,
None => vec![0.0; n],
};
Ok(Linear::new(w, b, n, k))
};
let mm = |wname: &str, n: usize, k: usize| -> Result<Linear> {
let w = st.tensor_f32(wname)?;
anyhow::ensure!(w.len() == n * k, "{wname}: {} != {n}·{k}", w.len());
let mut wt = vec![0f32; n * k];
for i in 0..k {
for j in 0..n {
wt[j * k + i] = w[i * n + j];
}
}
Ok(Linear::new(wt, vec![0.0; n], n, k))
};
let dw = |name: &str, c: usize, k: usize| -> Result<DwConv> {
let w = st.tensor_f32(name)?;
anyhow::ensure!(w.len() == c * k, "{name}: {} != {c}·{k}", w.len());
Ok(DwConv { w, c, k })
};
let block = |dwn: &[(&str, usize, usize)],
pws: &[(&str, &str, usize, usize)],
se_d: (&str, usize, usize),
se_u: (&str, usize, usize),
res: Option<(&str, &str, usize, usize)>|
-> Result<TitaBlock> {
let mut subs = Vec::new();
for ((dname, c, k), (pw_w, pw_b, n, kin)) in dwn.iter().zip(pws) {
subs.push((dw(dname, *c, *k)?, lin(pw_w, Some(pw_b), *n, *kin)?));
}
Ok(TitaBlock {
subs,
se: Se {
down: mm(se_d.0, se_d.2, se_d.1)?,
up: mm(se_u.0, se_u.2, se_u.1)?,
},
res: match res {
Some((w, b, n, k)) => Some(lin(w, Some(b), n, k)?),
None => None,
},
})
};
let blocks = vec![
block(
&[("encoder.encoder.0.mconv.0.conv.weight", 80, 3)],
&[("onnx::Conv_575", "onnx::Conv_576", 1024, 80)],
("onnx::MatMul_620", 1024, 128),
("onnx::MatMul_621", 128, 1024),
None,
)?,
block(
&[
("encoder.encoder.1.mconv.0.conv.weight", 1024, 7),
("encoder.encoder.1.mconv.5.conv.weight", 1024, 7),
("encoder.encoder.1.mconv.10.conv.weight", 1024, 7),
],
&[
("onnx::Conv_578", "onnx::Conv_579", 1024, 1024),
("onnx::Conv_581", "onnx::Conv_582", 1024, 1024),
("onnx::Conv_584", "onnx::Conv_585", 1024, 1024),
],
("onnx::MatMul_626", 1024, 128),
("onnx::MatMul_627", 128, 1024),
Some(("onnx::Conv_587", "onnx::Conv_588", 1024, 1024)),
)?,
block(
&[
("encoder.encoder.2.mconv.0.conv.weight", 1024, 11),
("encoder.encoder.2.mconv.5.conv.weight", 1024, 11),
("encoder.encoder.2.mconv.10.conv.weight", 1024, 11),
],
&[
("onnx::Conv_590", "onnx::Conv_591", 1024, 1024),
("onnx::Conv_593", "onnx::Conv_594", 1024, 1024),
("onnx::Conv_596", "onnx::Conv_597", 1024, 1024),
],
("onnx::MatMul_632", 1024, 128),
("onnx::MatMul_633", 128, 1024),
Some(("onnx::Conv_599", "onnx::Conv_600", 1024, 1024)),
)?,
block(
&[
("encoder.encoder.3.mconv.0.conv.weight", 1024, 15),
("encoder.encoder.3.mconv.5.conv.weight", 1024, 15),
("encoder.encoder.3.mconv.10.conv.weight", 1024, 15),
],
&[
("onnx::Conv_602", "onnx::Conv_603", 1024, 1024),
("onnx::Conv_605", "onnx::Conv_606", 1024, 1024),
("onnx::Conv_608", "onnx::Conv_609", 1024, 1024),
],
("onnx::MatMul_638", 1024, 128),
("onnx::MatMul_639", 128, 1024),
Some(("onnx::Conv_611", "onnx::Conv_612", 1024, 1024)),
)?,
block(
&[("encoder.encoder.4.mconv.0.conv.weight", 1024, 1)],
&[("onnx::Conv_614", "onnx::Conv_615", 3072, 1024)],
("onnx::MatMul_644", 3072, 384),
("onnx::MatMul_645", 384, 3072),
None,
)?,
];
Ok(Self {
fbank: Fbank::new(),
blocks,
attn_pre: lin(
"decoder._pooling.attention_layer.0.conv_layer.weight",
Some("decoder._pooling.attention_layer.0.conv_layer.bias"),
128,
9216,
)?,
attn_bn: BnFold::load(st, "decoder._pooling.attention_layer.0.bn")?,
attn_out: lin(
"decoder._pooling.attention_layer.2.weight",
Some("decoder._pooling.attention_layer.2.bias"),
3072,
128,
)?,
emb_bn: BnFold::load(st, "decoder.emb_layers.0.0")?,
emb_out: lin(
"decoder.emb_layers.0.1.weight",
Some("decoder.emb_layers.0.1.bias"),
192,
6144,
)?,
})
}
fn forward(&self, feats: &[f32], t: usize) -> Vec<f32> {
let mut x = feats.to_vec();
for b in &self.blocks {
x = b.forward(&x, t);
}
let c = 3072;
let mut mean = vec![0f32; c];
let mut mean2 = vec![0f32; c];
for row in x.chunks_exact(c) {
for i in 0..c {
mean[i] += row[i];
mean2[i] += row[i] * row[i];
}
}
for i in 0..c {
mean[i] /= t as f32;
mean2[i] /= t as f32;
}
let std: Vec<f32> = (0..c)
.map(|i| (mean2[i] - mean[i] * mean[i]).max(1e-10).sqrt())
.collect();
let mut attn_in = vec![0f32; t * 3 * c];
for i in 0..t {
attn_in[i * 3 * c..][..c].copy_from_slice(&x[i * c..][..c]);
attn_in[i * 3 * c + c..][..c].copy_from_slice(&mean);
attn_in[i * 3 * c + 2 * c..][..c].copy_from_slice(&std);
}
let mut h = self.attn_pre.forward(&attn_in);
for v in h.iter_mut() {
*v = v.max(0.0);
}
self.attn_bn.apply(&mut h);
for v in h.iter_mut() {
*v = v.tanh();
}
let scores = self.attn_out.forward(&h); let mut alpha = vec![0f32; t * c];
for ch in 0..c {
let mut max = f32::NEG_INFINITY;
for i in 0..t {
max = max.max(scores[i * c + ch]);
}
let mut sum = 0f32;
for i in 0..t {
let e = (scores[i * c + ch] - max).exp();
alpha[i * c + ch] = e;
sum += e;
}
for i in 0..t {
alpha[i * c + ch] /= sum;
}
}
let mut mu = vec![0f32; c];
let mut mu2 = vec![0f32; c];
for i in 0..t {
for ch in 0..c {
let a = alpha[i * c + ch];
let v = x[i * c + ch];
mu[ch] += a * v;
mu2[ch] += a * v * v;
}
}
let mut pooled = vec![0f32; 2 * c];
pooled[..c].copy_from_slice(&mu);
for ch in 0..c {
pooled[c + ch] = (mu2[ch] - mu[ch] * mu[ch]).max(1e-10).sqrt();
}
self.emb_bn.apply(&mut pooled);
self.emb_out.forward(&pooled)
}
pub fn embed(&self, samples: &[f32]) -> Option<Vec<f32>> {
let (mut feats, frames) = self.fbank.compute(samples);
if frames == 0 {
return None;
}
normalize_per_feature(&mut feats, frames);
let e = self.forward(&feats, frames);
if e.iter().any(|v| v.is_nan()) {
None
} else {
Some(e)
}
}
}
const WINDOW: usize = 160_000;
const WINDOW_SHIFT: usize = 16_000; const RF_SIZE: usize = 991;
const RF_SHIFT: usize = 270;
const NUM_SPEAKERS: usize = 3;
const NUM_CLASSES: usize = 7;
struct LstmDir {
w: Vec<f32>, r: Vec<f32>, b: Vec<f32>, hidden: usize,
input: usize,
packed_w: std::sync::OnceLock<crate::cpu_gemm::PackedWeight>,
}
impl LstmDir {
fn run(&self, x: &[f32], t: usize) -> Vec<f32> {
let hn = self.hidden;
let g4 = 4 * hn;
let packed = self
.packed_w
.get_or_init(|| crate::cpu_gemm::PackedWeight::new(&self.w, g4, self.input));
let mut gates_all = vec![0f32; t * g4];
crate::cpu_gemm::gemm_packed(&mut gates_all, x, packed, t, Some(&self.b));
let mut h = vec![0f32; hn];
let mut c = vec![0f32; hn];
let mut out = vec![0f32; t * hn];
for i in 0..t {
let gates = &mut gates_all[i * g4..][..g4];
for (rix, g) in gates.iter_mut().enumerate() {
let rrow = &self.r[rix * hn..][..hn];
let mut acc = 0f32;
for k in 0..hn {
acc += rrow[k] * h[k];
}
*g += acc;
}
for j in 0..hn {
let ig = sigmoid(gates[j]);
let og = sigmoid(gates[hn + j]);
let fg = sigmoid(gates[2 * hn + j]);
let cg = gates[3 * hn + j].tanh();
c[j] = fg * c[j] + ig * cg;
h[j] = og * c[j].tanh();
}
out[i * hn..][..hn].copy_from_slice(&h);
}
out
}
}
struct BiLstm {
fwd: LstmDir,
bwd: LstmDir,
}
impl BiLstm {
fn load(st: &LazySt, w: &str, r: &str, b: &str, input: usize, hidden: usize) -> Result<Self> {
let wt = st.tensor_f32(w)?; let rt = st.tensor_f32(r)?; let bt = st.tensor_f32(b)?; let g4 = 4 * hidden;
let dir = |d: usize| -> LstmDir {
let b: Vec<f32> = (0..g4)
.map(|i| bt[d * 8 * hidden + i] + bt[d * 8 * hidden + g4 + i])
.collect();
LstmDir {
w: wt[d * g4 * input..][..g4 * input].to_vec(),
r: rt[d * g4 * hidden..][..g4 * hidden].to_vec(),
b,
hidden,
input,
packed_w: std::sync::OnceLock::new(),
}
};
Ok(Self {
fwd: dir(0),
bwd: dir(1),
})
}
fn run(&self, x: &[f32], t: usize) -> Vec<f32> {
let hn = self.fwd.hidden;
let d = self.fwd.input;
let mut rev = vec![0f32; x.len()];
for i in 0..t {
rev[i * d..][..d].copy_from_slice(&x[(t - 1 - i) * d..][..d]);
}
let (f, bwd) = rayon::join(|| self.fwd.run(x, t), || self.bwd.run(&rev, t));
let mut out = vec![0f32; t * 2 * hn];
for i in 0..t {
out[i * 2 * hn..][..hn].copy_from_slice(&f[i * hn..][..hn]);
out[i * 2 * hn + hn..][..hn].copy_from_slice(&bwd[(t - 1 - i) * hn..][..hn]);
}
out
}
}
struct Conv1d {
w: Vec<f32>,
b: Vec<f32>,
out_c: usize,
in_c: usize,
k: usize,
stride: usize,
}
impl Conv1d {
fn forward(&self, x: &[f32], t: usize) -> (Vec<f32>, usize) {
use rayon::prelude::*;
let to = (t - self.k) / self.stride + 1;
let mut out = vec![0f32; to * self.out_c];
out.par_chunks_mut(self.out_c)
.enumerate()
.for_each(|(i, orow)| {
for oc in 0..self.out_c {
let mut acc = self.b[oc];
let wbase = oc * self.in_c * self.k;
for ic in 0..self.in_c {
let wrow = &self.w[wbase + ic * self.k..][..self.k];
for tap in 0..self.k {
acc += wrow[tap] * x[(i * self.stride + tap) * self.in_c + ic];
}
}
orow[oc] = acc;
}
});
(out, to)
}
}
fn instance_norm(x: &mut [f32], t: usize, w: &[f32], b: &[f32]) {
let c = w.len();
for ch in 0..c {
let mut sum = 0f64;
for i in 0..t {
sum += x[i * c + ch] as f64;
}
let mean = (sum / t as f64) as f32;
let mut var = 0f64;
for i in 0..t {
let d = x[i * c + ch] - mean;
var += (d * d) as f64;
}
let inv = 1.0 / ((var / t as f64) as f32 + 1e-5).sqrt();
for i in 0..t {
x[i * c + ch] = (x[i * c + ch] - mean) * inv * w[ch] + b[ch];
}
}
}
fn leaky_relu(x: &mut [f32]) {
for v in x.iter_mut() {
if *v < 0.0 {
*v *= 0.01;
}
}
}
fn max_pool3(x: &[f32], t: usize, c: usize) -> (Vec<f32>, usize) {
let to = (t - 3) / 3 + 1;
let mut out = vec![f32::NEG_INFINITY; to * c];
for i in 0..to {
for tap in 0..3 {
let row = &x[(i * 3 + tap) * c..][..c];
let orow = &mut out[i * c..][..c];
for ch in 0..c {
orow[ch] = orow[ch].max(row[ch]);
}
}
}
(out, to)
}
pub struct Pyannote {
wav_norm: (f32, f32), conv0: Conv1d, norms: [(Vec<f32>, Vec<f32>); 3],
conv1: Conv1d,
conv2: Conv1d,
lstms: Vec<BiLstm>,
lin0: Linear,
lin1: Linear,
classifier: Linear,
}
impl Pyannote {
fn load(st: &LazySt) -> Result<Self> {
let t = |n: &str| st.tensor_f32(n);
let conv =
|w: Vec<f32>, b: Vec<f32>, out_c: usize, in_c: usize, k: usize, s: usize| Conv1d {
w,
b,
out_c,
in_c,
k,
stride: s,
};
let mm = |wname: &str, bias: Vec<f32>, n: usize, k: usize| -> Result<Linear> {
let w = t(wname)?; let mut wt = vec![0f32; n * k];
for i in 0..k {
for j in 0..n {
wt[j * k + i] = w[i * n + j];
}
}
Ok(Linear::new(wt, bias, n, k))
};
Ok(Self {
wav_norm: (
t("ortshared_1_1_1_1_token_110")?[0],
t("ortshared_1_1_1_0_token_107")?[0],
),
conv0: conv(
t("/sincnet/conv1d.0/Concat_2_output_0")?,
vec![0.0; 80],
80,
1,
251,
10,
),
norms: [
(t("sincnet.norm1d.0.weight")?, t("sincnet.norm1d.0.bias")?),
(t("sincnet.norm1d.1.weight")?, t("sincnet.norm1d.1.bias")?),
(t("sincnet.norm1d.2.weight")?, t("sincnet.norm1d.2.bias")?),
],
conv1: conv(
t("sincnet.conv1d.1.weight")?,
t("sincnet.conv1d.1.bias")?,
60,
80,
5,
1,
),
conv2: conv(
t("sincnet.conv1d.2.weight")?,
t("sincnet.conv1d.2.bias")?,
60,
60,
5,
1,
),
lstms: vec![
BiLstm::load(
st,
"onnx::LSTM_784",
"onnx::LSTM_785",
"onnx::LSTM_783",
60,
128,
)?,
BiLstm::load(
st,
"onnx::LSTM_827",
"onnx::LSTM_828",
"onnx::LSTM_826",
256,
128,
)?,
BiLstm::load(
st,
"onnx::LSTM_870",
"onnx::LSTM_871",
"onnx::LSTM_869",
256,
128,
)?,
BiLstm::load(
st,
"onnx::LSTM_913",
"onnx::LSTM_914",
"onnx::LSTM_912",
256,
128,
)?,
],
lin0: mm("onnx::MatMul_915", t("linear.0.bias")?, 128, 256)?,
lin1: mm("onnx::MatMul_916", t("linear.1.bias")?, 128, 128)?,
classifier: mm(
"onnx::MatMul_917",
t("ortshared_1_1_7_0_token_109")?,
7,
128,
)?,
})
}
pub fn forward(&self, window: &[f32]) -> (Vec<f32>, usize) {
assert_eq!(window.len(), WINDOW);
let prof = std::env::var("OSFKB_DIAR_PROFILE").as_deref() == Ok("1");
let mut clk = prof.then(std::time::Instant::now);
let mut tick = |name: &str| {
if let Some(c) = clk.as_mut() {
eprintln!(
" pyannote/{name}: {:.1} ms",
c.elapsed().as_secs_f64() * 1000.0
);
*c = std::time::Instant::now();
}
};
let n = window.len();
let mean = window.iter().map(|v| *v as f64).sum::<f64>() / n as f64;
let var = window
.iter()
.map(|v| (*v as f64 - mean) * (*v as f64 - mean))
.sum::<f64>()
/ n as f64;
let inv = 1.0 / ((var as f32) + 1e-5).sqrt();
let x: Vec<f32> = window
.iter()
.map(|v| (*v - mean as f32) * inv * self.wav_norm.0 + self.wav_norm.1)
.collect();
let (mut h, t) = self.conv0.forward(&x, n); for v in h.iter_mut() {
*v = v.abs();
}
let (mut h, mut t) = {
let (p, tp) = max_pool3(&h, t, 80);
(p, tp)
};
instance_norm(&mut h, t, &self.norms[0].0, &self.norms[0].1);
leaky_relu(&mut h);
let (h2, t2) = self.conv1.forward(&h, t);
let (mut h2, t2) = max_pool3(&h2, t2, 60);
instance_norm(&mut h2, t2, &self.norms[1].0, &self.norms[1].1);
leaky_relu(&mut h2);
let (h3, t3) = self.conv2.forward(&h2, t2);
let (mut h3, t3) = max_pool3(&h3, t3, 60);
instance_norm(&mut h3, t3, &self.norms[2].0, &self.norms[2].1);
leaky_relu(&mut h3);
h = h3;
t = t3;
tick("sincnet_convs");
for lstm in &self.lstms {
h = lstm.run(&h, t);
}
tick("lstms");
let mut h = self.lin0.forward(&h);
leaky_relu(&mut h);
let mut h = self.lin1.forward(&h);
leaky_relu(&mut h);
let mut logits = self.classifier.forward(&h); for row in logits.chunks_exact_mut(NUM_CLASSES) {
let max = row.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let lse = row.iter().map(|v| (v - max).exp()).sum::<f32>().ln() + max;
for v in row.iter_mut() {
*v -= lse;
}
}
tick("linears");
(logits, t)
}
}
fn ahc_complete_cutree_k(dist: &[f64], n: usize, k: usize) -> Vec<i32> {
if n == 1 {
return vec![0];
}
let d = |i: usize, j: usize| -> f64 {
let (a, b) = if i < j { (i, j) } else { (j, i) };
dist[a * n - a * (a + 1) / 2 + b - a - 1]
};
let mut nodes: Vec<(i32, Vec<usize>)> = (0..n).map(|i| (-(i as i32) - 1, vec![i])).collect();
let mut merge = vec![0i32; 2 * (n - 1)];
for step in 1..n {
let mut best = (f64::INFINITY, 0usize, 0usize);
for a in 0..nodes.len() {
for b in a + 1..nodes.len() {
let mut dmax = 0f64;
for &x in &nodes[a].1 {
for &y in &nodes[b].1 {
dmax = dmax.max(d(x, y));
}
}
if dmax < best.0 {
best = (dmax, a, b);
}
}
}
let (_, a, b) = best;
merge[step - 1] = nodes[a].0;
merge[n - 1 + step - 1] = nodes[b].0;
let mut members = std::mem::take(&mut nodes[a].1);
members.extend_from_slice(&nodes[b].1);
nodes[a] = (step as i32, members);
nodes.remove(b);
}
let nclust = k.min(n).max(1);
let mut labels = vec![0i32; n];
if nclust < 2 {
return labels;
}
let mut last_merge = vec![0i32; n];
for kk in 1..=(n - nclust) as i32 {
let m1 = merge[(kk - 1) as usize];
let m2 = merge[n - 1 + (kk - 1) as usize];
if m1 < 0 && m2 < 0 {
last_merge[(-m1 - 1) as usize] = kk;
last_merge[(-m2 - 1) as usize] = kk;
} else if m1 < 0 || m2 < 0 {
let (j, mc) = if m1 < 0 { (-m1, m2) } else { (-m2, m1) };
for l in 0..n {
if last_merge[l] == mc {
last_merge[l] = kk;
}
}
last_merge[(j - 1) as usize] = kk;
} else {
for l in 0..n {
if last_merge[l] == m1 || last_merge[l] == m2 {
last_merge[l] = kk;
}
}
}
}
let mut label = 0i32;
let mut z = vec![-1i32; n];
for j in 0..n {
if last_merge[j] == 0 {
labels[j] = label;
label += 1;
} else {
if z[last_merge[j] as usize] < 0 {
z[last_merge[j] as usize] = label;
label += 1;
}
labels[j] = z[last_merge[j] as usize];
}
}
labels
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct DiarSegment {
pub start: f32,
pub end: f32,
pub speaker: i32,
}
pub struct Diarizer {
pub pyannote: Pyannote,
pub titanet: TitaNet,
pub num_clusters: usize,
pub min_duration_on: f32,
pub min_duration_off: f32,
powerset: [[i32; NUM_SPEAKERS]; NUM_CLASSES],
}
impl Diarizer {
pub fn load(dir: &Path) -> Result<Self> {
let read = |name: &str| -> Result<Vec<u8>> {
std::fs::read(dir.join(name)).with_context(|| format!("{name} in {}", dir.display()))
};
Self::from_bytes(read("pyannote.safetensors")?, read("titanet.safetensors")?)
}
pub fn from_bytes(pyannote: Vec<u8>, titanet: Vec<u8>) -> Result<Self> {
let py = LazySt::from_bytes(vec![pyannote])?;
let ti = LazySt::from_bytes(vec![titanet])?;
let mut powerset = [[0i32; NUM_SPEAKERS]; NUM_CLASSES];
let mut kk = 1;
for j in 0..NUM_SPEAKERS {
powerset[kk][j] = 1;
kk += 1;
}
for j in 0..NUM_SPEAKERS {
for m in j + 1..NUM_SPEAKERS {
powerset[kk][j] = 1;
powerset[kk][m] = 1;
kk += 1;
}
}
Ok(Self {
pyannote: Pyannote::load(&py)?,
titanet: TitaNet::load(&ti)?,
num_clusters: 2,
min_duration_on: 0.2,
min_duration_off: 0.3,
powerset,
})
}
pub fn process(&self, audio: &[f32]) -> Vec<DiarSegment> {
self.process_with_clusters(audio, self.num_clusters)
}
pub fn process_with_clusters(&self, audio: &[f32], num_clusters: usize) -> Vec<DiarSegment> {
let n = audio.len();
if n == 0 {
return Vec::new();
}
let mut labels: Vec<(Vec<i32>, usize)> = Vec::new(); let chunk_starts: Vec<usize>;
let has_last_chunk;
if n <= WINDOW {
chunk_starts = vec![0];
has_last_chunk = false;
} else {
let num_chunks = (n - WINDOW) / WINDOW_SHIFT + 1;
has_last_chunk = !(n - WINDOW).is_multiple_of(WINDOW_SHIFT);
let mut s: Vec<usize> = (0..num_chunks).map(|i| i * WINDOW_SHIFT).collect();
if has_last_chunk {
s.push(num_chunks * WINDOW_SHIFT);
}
chunk_starts = s;
}
use rayon::prelude::*;
labels = chunk_starts
.par_iter()
.map(|&start| {
let mut buf = vec![0f32; WINDOW];
let end = (start + WINDOW).min(n);
buf[..end - start].copy_from_slice(&audio[start..end]);
let (logp, frames) = self.pyannote.forward(&buf);
let mut lab = vec![0i32; frames * NUM_SPEAKERS];
for (i, row) in logp.chunks_exact(NUM_CLASSES).enumerate() {
let mut best = 0;
for c in 1..NUM_CLASSES {
if row[c] > row[best] {
best = c;
}
}
lab[i * NUM_SPEAKERS..][..NUM_SPEAKERS].copy_from_slice(&self.powerset[best]);
}
(lab, frames)
})
.collect();
if labels.len() == 1 {
let (lab, frames) = &labels[0];
let keep = if (n as isize - WINDOW as isize) % WINDOW_SHIFT as isize > 0 && n > WINDOW {
(n / RF_SHIFT).min(*frames)
} else if n < WINDOW {
*frames
} else {
*frames
};
return self.compute_result(lab, keep, self.num_speaker_columns(&labels));
}
let num_chunks = labels.len();
let num_frames_total = (WINDOW + (num_chunks - 1) * WINDOW_SHIFT) / RF_SHIFT + 1;
let mut count = vec![0f32; num_frames_total];
let mut weight = vec![0f32; num_frames_total];
for (i, (lab, frames)) in labels.iter().enumerate() {
let start = (i as f32 * WINDOW_SHIFT as f32 / RF_SHIFT as f32 + 0.5) as usize;
for f in 0..*frames {
let s: i32 = lab[f * NUM_SPEAKERS..][..NUM_SPEAKERS].iter().sum();
count[start + f] += s as f32;
weight[start + f] += 1.0;
}
}
let speakers_per_frame: Vec<i32> = (0..num_frames_total)
.map(|i| (count[i] / (weight[i] + 1e-12) + 0.5) as i32)
.collect();
if speakers_per_frame.iter().all(|&s| s == 0) {
return Vec::new();
}
let mut chunk_speaker: Vec<(usize, usize)> = Vec::new();
let mut ranges: Vec<Vec<(usize, usize)>> = Vec::new();
for (ci, (lab, frames)) in labels.iter().enumerate() {
let sample_offset = ci * WINDOW_SHIFT;
for spk in 0..NUM_SPEAKERS {
let active = |f: usize| -> bool {
let row = &lab[f * NUM_SPEAKERS..][..NUM_SPEAKERS];
let total: i32 = row.iter().sum();
row[spk] == 1 && total < 2
};
let total_active = (0..*frames).filter(|&f| active(f)).count();
if total_active < 10 {
continue;
}
let mut segs = Vec::new();
let mut is_active = false;
let mut start_f = 0usize;
for f in 0..*frames {
if active(f) {
if !is_active {
is_active = true;
start_f = f;
}
} else if is_active {
is_active = false;
segs.push((
(start_f as f32 / *frames as f32 * WINDOW as f32) as usize
+ sample_offset,
(f as f32 / *frames as f32 * WINDOW as f32) as usize + sample_offset,
));
}
}
if is_active {
segs.push((
(start_f as f32 / *frames as f32 * WINDOW as f32) as usize + sample_offset,
((*frames as f32 - 1.0) / *frames as f32 * WINDOW as f32) as usize
+ sample_offset,
));
}
chunk_speaker.push((ci, spk));
ranges.push(segs);
}
}
let computed: Vec<Option<Vec<f32>>> = ranges
.par_iter()
.map(|segs| {
let mut samples = Vec::new();
for &(s, e) in segs {
let e = e.min(n);
if e > s {
samples.extend_from_slice(&audio[s..e]);
}
}
self.titanet.embed(&samples)
})
.collect();
let mut embeddings: Vec<Vec<f32>> = Vec::new();
let mut kept: Vec<usize> = Vec::new();
for (idx, e) in computed.into_iter().enumerate() {
if let Some(e) = e {
embeddings.push(e);
kept.push(idx);
}
}
if embeddings.is_empty() {
return Vec::new();
}
let chunk_speaker: Vec<(usize, usize)> = kept.iter().map(|&i| chunk_speaker[i]).collect();
let rows = embeddings.len();
let cluster_labels: Vec<i32> = if rows == 1 {
vec![0]
} else {
for e in embeddings.iter_mut() {
let norm = e.iter().map(|v| v * v).sum::<f32>().sqrt();
if norm > 0.0 {
for v in e.iter_mut() {
*v /= norm;
}
}
}
let mut dist = Vec::with_capacity(rows * (rows - 1) / 2);
for i in 0..rows {
for j in i + 1..rows {
let cs: f64 = embeddings[i]
.iter()
.zip(&embeddings[j])
.map(|(a, b)| (*a as f64) * (*b as f64))
.sum();
dist.push((1.0 - cs).max(0.0));
}
}
ahc_complete_cutree_k(&dist, rows, num_clusters)
};
let max_cluster = *cluster_labels.iter().max().unwrap() as usize;
let mut new_labels: Vec<(Vec<i32>, usize)> = Vec::with_capacity(labels.len());
for (ci, (lab, frames)) in labels.iter().enumerate() {
let mut nl = vec![0i32; frames * (max_cluster + 1)];
for spk in 0..NUM_SPEAKERS {
let Some(pos) = chunk_speaker.iter().position(|&(c, s)| c == ci && s == spk) else {
continue;
};
let cl = cluster_labels[pos] as usize;
for f in 0..*frames {
if lab[f * NUM_SPEAKERS + spk] == 1 {
nl[f * (max_cluster + 1) + cl] = 1;
}
}
}
new_labels.push((nl, *frames));
}
let ncols = max_cluster + 1;
let mut speaker_count = vec![0i32; num_frames_total * ncols];
for (i, (nl, frames)) in new_labels.iter().enumerate() {
let start = (i as f32 * WINDOW_SHIFT as f32 / RF_SHIFT as f32 + 0.5) as usize;
for f in 0..*frames {
for c in 0..ncols {
speaker_count[(start + f) * ncols + c] += nl[f * ncols + c];
}
}
}
let mut rows_kept = num_frames_total;
if has_last_chunk {
rows_kept = (n / RF_SHIFT + 1).min(num_frames_total);
}
let mut final_labels = vec![0i32; rows_kept * ncols];
for f in 0..rows_kept {
let k = speakers_per_frame[f].max(0) as usize;
if k == 0 {
continue;
}
let counts = &speaker_count[f * ncols..][..ncols];
let mut idx: Vec<usize> = (0..ncols).collect();
idx.sort_by(|&a, &b| counts[b].cmp(&counts[a]).then(a.cmp(&b)));
for &m in idx.iter().take(k.min(ncols)) {
final_labels[f * ncols + m] = 1;
}
}
self.segments_from_labels(&final_labels, rows_kept, ncols)
}
fn num_speaker_columns(&self, _labels: &[(Vec<i32>, usize)]) -> usize {
NUM_SPEAKERS
}
fn compute_result(&self, lab: &[i32], frames: usize, ncols: usize) -> Vec<DiarSegment> {
self.segments_from_labels(&lab[..frames * ncols], frames, ncols)
}
fn segments_from_labels(
&self,
labels: &[i32],
frames: usize,
ncols: usize,
) -> Vec<DiarSegment> {
let scale = RF_SHIFT as f32 / SR as f32;
let offset = 0.5 * RF_SIZE as f32 / SR as f32;
let mut all = Vec::new();
for spk in 0..ncols {
let mut segs: Vec<DiarSegment> = Vec::new();
let mut is_active = labels[spk] > 0;
let mut start_index = if is_active { 0usize } else { 0 };
for f in 1..frames {
if is_active {
if labels[f * ncols + spk] == 0 {
segs.push(DiarSegment {
start: start_index as f32 * scale + offset,
end: f as f32 * scale + offset,
speaker: spk as i32,
});
is_active = false;
}
} else if labels[f * ncols + spk] == 1 {
is_active = true;
start_index = f;
}
}
if is_active {
segs.push(DiarSegment {
start: start_index as f32 * scale + offset,
end: (frames - 1) as f32 * scale + offset,
speaker: spk as i32,
});
}
let mut changed = true;
while changed {
changed = false;
for i in 0..segs.len().saturating_sub(1) {
let (a, b) = (segs[i], segs[i + 1]);
if a.end < b.start && a.end + self.min_duration_off >= b.start {
segs[i] = DiarSegment {
start: a.start,
end: b.end,
speaker: a.speaker,
};
segs.remove(i + 1);
changed = true;
break;
}
}
}
all.extend(
segs.into_iter()
.filter(|s| s.end - s.start > self.min_duration_on),
);
}
all.sort_by(|a, b| {
a.start
.partial_cmp(&b.start)
.unwrap()
.then(a.speaker.cmp(&b.speaker))
});
all
}
}