use std::path::Path;
use anyhow::{Context, Result};
use crate::cpu_gemm::{PackedWeight, gemm_packed};
use crate::weights::LazySt;
const N_FFT: usize = 4096;
const N_HOP: usize = 1024;
const NB_FFT_BINS: usize = N_FFT / 2 + 1;
const NB_CHANNELS: usize = 2;
const BN_EPS: f32 = 1e-5;
pub const MODEL_SR: u32 = 44100;
pub const PIPELINE_SR: u32 = 16000;
struct Fft {
n: usize,
rev: Vec<u32>,
tw_re: Vec<f64>,
tw_im: Vec<f64>,
}
impl Fft {
fn new(n: usize) -> Self {
assert!(n.is_power_of_two(), "radix-2 needs a power of two");
let bits = n.trailing_zeros();
let rev = (0..n)
.map(|i| (i as u32).reverse_bits() >> (32 - bits))
.collect();
let (mut tw_re, mut tw_im) = (vec![0f64; n / 2], vec![0f64; n / 2]);
for j in 0..n / 2 {
let a = -2.0 * std::f64::consts::PI * j as f64 / n as f64;
tw_re[j] = a.cos();
tw_im[j] = a.sin();
}
Self { n, rev, tw_re, tw_im }
}
fn run(&self, re: &mut [f64], im: &mut [f64], inverse: bool) {
let n = self.n;
for i in 0..n {
let j = self.rev[i] as usize;
if j > i {
re.swap(i, j);
im.swap(i, j);
}
}
let sign = if inverse { -1.0 } else { 1.0 };
let mut len = 2;
while len <= n {
let step = n / len;
let half = len / 2;
for base in (0..n).step_by(len) {
for j in 0..half {
let (wr, wi) = (self.tw_re[j * step], sign * self.tw_im[j * step]);
let (a, b) = (base + j, base + j + half);
let (vr, vi) = (re[b] * wr - im[b] * wi, re[b] * wi + im[b] * wr);
let (ur, ui) = (re[a], im[a]);
re[a] = ur + vr;
im[a] = ui + vi;
re[b] = ur - vr;
im[b] = ui - vi;
}
}
len <<= 1;
}
}
}
struct Linear {
w: Vec<f32>,
b: Option<Vec<f32>>,
n: usize,
k: usize,
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}: expected [{n}, {k}], got {}", w.len());
Ok(Self { w, b: None, 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, self.b.as_deref());
out
}
}
struct BatchNorm {
scale: Vec<f32>,
shift: Vec<f32>,
}
impl BatchNorm {
fn load(st: &LazySt, prefix: &str, c: usize) -> Result<Self> {
let w = 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"))?;
anyhow::ensure!(w.len() == c, "{prefix}: expected {c} channels, got {}", w.len());
let scale: Vec<f32> = (0..c).map(|i| w[i] / (var[i] + BN_EPS).sqrt()).collect();
let shift: Vec<f32> = (0..c).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, v) in row.iter_mut().enumerate() {
*v = *v * self.scale[i] + self.shift[i];
}
}
}
}
fn sigmoid(v: f32) -> f32 {
1.0 / (1.0 + (-v).exp())
}
struct LstmDir {
w: Vec<f32>, r: Vec<f32>, b: Vec<f32>, hidden: usize,
input: usize,
packed_w: std::sync::OnceLock<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(|| PackedWeight::new(&self.w, g4, self.input));
let mut gates_all = vec![0f32; t * g4];
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 fg = sigmoid(gates[hn + j]);
let cg = gates[2 * hn + j].tanh();
let og = sigmoid(gates[3 * hn + j]);
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, layer: usize, input: usize, hidden: usize) -> Result<Self> {
let g4 = 4 * hidden;
let dir = |suffix: &str| -> Result<LstmDir> {
let w = st.tensor_f32(&format!("lstm.weight_ih_l{layer}{suffix}"))?;
let r = st.tensor_f32(&format!("lstm.weight_hh_l{layer}{suffix}"))?;
let bi = st.tensor_f32(&format!("lstm.bias_ih_l{layer}{suffix}"))?;
let bh = st.tensor_f32(&format!("lstm.bias_hh_l{layer}{suffix}"))?;
anyhow::ensure!(
w.len() == g4 * input && r.len() == g4 * hidden && bi.len() == g4,
"lstm l{layer}{suffix}: shape mismatch"
);
Ok(LstmDir {
w,
r,
b: (0..g4).map(|i| bi[i] + bh[i]).collect(),
hidden,
input,
packed_w: std::sync::OnceLock::new(),
})
};
Ok(Self { fwd: dir("")?, bwd: dir("_reverse")? })
}
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, b) = 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(&b[(t - 1 - i) * hn..][..hn]);
}
out
}
}
pub struct Pass {
pub frames: usize,
pub mag: Vec<f32>,
pub fc1_tanh: Vec<f32>,
pub lstm_out: Vec<f32>,
pub fc2_relu: Vec<f32>,
pub mask: Vec<f32>,
pub est_mag: Vec<f32>,
re: Vec<f32>,
im: Vec<f32>,
}
pub struct Separator {
fft: Fft,
window: Vec<f64>,
input_mean: Vec<f32>,
input_scale: Vec<f32>,
output_mean: Vec<f32>,
output_scale: Vec<f32>,
fc1: Linear,
bn1: BatchNorm,
lstm: Vec<BiLstm>,
fc2: Linear,
bn2: BatchNorm,
fc3: Linear,
bn3: BatchNorm,
nb_bins: usize,
hidden: usize,
}
impl Separator {
pub fn load(dir: &Path) -> Result<Self> {
let bytes = std::fs::read(dir.join("umxhq_vocals.safetensors"))
.with_context(|| format!("umxhq_vocals.safetensors in {}", dir.display()))?;
let st = LazySt::from_bytes(vec![bytes])?;
let input_mean = st.tensor_f32("input_mean")?;
let input_scale = st.tensor_f32("input_scale")?;
let output_mean = st.tensor_f32("output_mean")?;
let output_scale = st.tensor_f32("output_scale")?;
let nb_bins = input_mean.len();
anyhow::ensure!(
output_mean.len() == NB_FFT_BINS && nb_bins <= NB_FFT_BINS,
"unexpected umx bandwidth: {nb_bins} in, {} out",
output_mean.len()
);
let hidden = st.shape("bn1.weight")?[0];
let fc1 = Linear::load(&st, "fc1.weight", hidden, NB_CHANNELS * nb_bins)?;
let fc2 = Linear::load(&st, "fc2.weight", hidden, 2 * hidden)?;
let fc3 = Linear::load(&st, "fc3.weight", NB_CHANNELS * NB_FFT_BINS, hidden)?;
let half = hidden / 2;
let lstm = (0..3)
.map(|l| BiLstm::load(&st, l, hidden, half))
.collect::<Result<Vec<_>>>()?;
let window = (0..N_FFT)
.map(|i| 0.5 * (1.0 - (2.0 * std::f64::consts::PI * i as f64 / N_FFT as f64).cos()))
.collect();
Ok(Self {
fft: Fft::new(N_FFT),
window,
input_mean,
input_scale,
output_mean,
output_scale,
fc1,
bn1: BatchNorm::load(&st, "bn1", hidden)?,
lstm,
fc2,
bn2: BatchNorm::load(&st, "bn2", hidden)?,
fc3,
bn3: BatchNorm::load(&st, "bn3", NB_CHANNELS * NB_FFT_BINS)?,
nb_bins,
hidden,
})
}
fn frames_for(len: usize) -> usize {
len / N_HOP + 1
}
fn stft(&self, channels: &[&[f32]]) -> (Vec<f32>, Vec<f32>, usize) {
let len = channels[0].len();
let frames = Self::frames_for(len);
let nch = channels.len();
let mut re = vec![0f32; nch * NB_FFT_BINS * frames];
let mut im = vec![0f32; nch * NB_FFT_BINS * frames];
let l = len as i64;
for (c, ch) in channels.iter().enumerate() {
let at = |i: i64| -> f32 {
let j = if i < 0 {
-i
} else if i >= l {
2 * l - 2 - i
} else {
i
};
ch[j.clamp(0, l - 1) as usize]
};
let mut buf_re = vec![0f64; N_FFT];
let mut buf_im = vec![0f64; N_FFT];
for f in 0..frames {
let start = f as i64 * N_HOP as i64 - (N_FFT / 2) as i64;
for j in 0..N_FFT {
buf_re[j] = at(start + j as i64) as f64 * self.window[j];
buf_im[j] = 0.0;
}
self.fft.run(&mut buf_re, &mut buf_im, false);
let base = c * NB_FFT_BINS * frames;
for b in 0..NB_FFT_BINS {
re[base + b * frames + f] = buf_re[b] as f32;
im[base + b * frames + f] = buf_im[b] as f32;
}
}
}
(re, im, frames)
}
fn istft(&self, re: &[f32], im: &[f32], frames: usize, length: usize) -> Vec<f32> {
let full = (frames - 1) * N_HOP + N_FFT;
let mut y = vec![0f64; full];
let mut env = vec![0f64; full];
let inv_n = 1.0 / N_FFT as f64;
let mut buf_re = vec![0f64; N_FFT];
let mut buf_im = vec![0f64; N_FFT];
for f in 0..frames {
for b in 0..NB_FFT_BINS {
buf_re[b] = re[b * frames + f] as f64;
buf_im[b] = im[b * frames + f] as f64;
}
for b in NB_FFT_BINS..N_FFT {
buf_re[b] = buf_re[N_FFT - b];
buf_im[b] = -buf_im[N_FFT - b];
}
self.fft.run(&mut buf_re, &mut buf_im, true);
let off = f * N_HOP;
for j in 0..N_FFT {
y[off + j] += buf_re[j] * inv_n * self.window[j];
env[off + j] += self.window[j] * self.window[j];
}
}
let start = N_FFT / 2;
(0..length)
.map(|i| {
let k = start + i;
if k < full && env[k] > 1e-11 {
(y[k] / env[k]) as f32
} else {
0.0
}
})
.collect()
}
fn forward(&self, mag: &[f32], frames: usize) -> (Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>) {
let nb = self.nb_bins;
let hidden = self.hidden;
let mut x = vec![0f32; frames * NB_CHANNELS * nb];
for f in 0..frames {
for c in 0..NB_CHANNELS {
let src = c * NB_FFT_BINS * frames;
let dst = f * NB_CHANNELS * nb + c * nb;
for b in 0..nb {
x[dst + b] = (mag[src + b * frames + f] + self.input_mean[b])
* self.input_scale[b];
}
}
}
let mut h = self.fc1.forward(&x);
self.bn1.apply(&mut h);
for v in h.iter_mut() {
*v = v.tanh();
}
let fc1_tanh = h.clone();
let mut lstm_out = fc1_tanh.clone();
for layer in &self.lstm {
lstm_out = layer.run(&lstm_out, frames);
}
let mut cat = vec![0f32; frames * 2 * hidden];
for f in 0..frames {
cat[f * 2 * hidden..][..hidden].copy_from_slice(&fc1_tanh[f * hidden..][..hidden]);
cat[f * 2 * hidden + hidden..][..hidden]
.copy_from_slice(&lstm_out[f * hidden..][..hidden]);
}
let mut d = self.fc2.forward(&cat);
self.bn2.apply(&mut d);
for v in d.iter_mut() {
*v = v.max(0.0);
}
let fc2_relu = d.clone();
let mut o = self.fc3.forward(&fc2_relu);
self.bn3.apply(&mut o);
let mut mask = vec![0f32; NB_CHANNELS * NB_FFT_BINS * frames];
for f in 0..frames {
for c in 0..NB_CHANNELS {
let src = f * NB_CHANNELS * NB_FFT_BINS + c * NB_FFT_BINS;
let dst = c * NB_FFT_BINS * frames;
for b in 0..NB_FFT_BINS {
let v = o[src + b] * self.output_scale[b] + self.output_mean[b];
mask[dst + b * frames + f] = v.max(0.0);
}
}
}
(fc1_tanh, lstm_out, fc2_relu, mask)
}
pub fn pass(&self, left: &[f32], right: &[f32]) -> Pass {
assert_eq!(left.len(), right.len(), "channels must be the same length");
let (re, im, frames) = self.stft(&[left, right]);
let mag: Vec<f32> = re
.iter()
.zip(&im)
.map(|(r, i)| (*r as f64).hypot(*i as f64) as f32)
.collect();
let (fc1_tanh, lstm_out, fc2_relu, mask) = self.forward(&mag, frames);
let est_mag: Vec<f32> = mask.iter().zip(&mag).map(|(m, x)| m * x).collect();
Pass { frames, mag, fc1_tanh, lstm_out, fc2_relu, mask, est_mag, re, im }
}
pub fn stem(&self, pass: &Pass, length: usize) -> (Vec<f32>, Vec<f32>) {
let n = NB_FFT_BINS * pass.frames;
let mut out = Vec::with_capacity(NB_CHANNELS);
for c in 0..NB_CHANNELS {
let (mut yr, mut yi) = (vec![0f32; n], vec![0f32; n]);
for k in 0..n {
let (r, i) = (pass.re[c * n + k], pass.im[c * n + k]);
let m = (r as f64).hypot(i as f64);
let est = pass.est_mag[c * n + k] as f64;
let (cos, sin) = if m > 0.0 { (r as f64 / m, i as f64 / m) } else { (1.0, 0.0) };
yr[k] = (est * cos) as f32;
yi[k] = (est * sin) as f32;
}
out.push(self.istft(&yr, &yi, pass.frames, length));
}
let right = out.pop().unwrap();
(out.pop().unwrap(), right)
}
pub fn separate(&self, left: &[f32], right: &[f32]) -> (Vec<f32>, Vec<f32>) {
let pass = self.pass(left, right);
self.stem(&pass, left.len())
}
pub fn vocals(&self, mono_16k: &[f32]) -> Vec<f32> {
if mono_16k.len() < N_FFT {
return mono_16k.to_vec();
}
let up = crate::chatterbox::resample_sinc(mono_16k, PIPELINE_SR, MODEL_SR);
let (l, r) = self.separate(&up, &up);
let mono: Vec<f32> = l.iter().zip(&r).map(|(a, b)| 0.5 * (a + b)).collect();
let mut out = crate::chatterbox::resample_sinc(&mono, MODEL_SR, PIPELINE_SR);
out.truncate(mono_16k.len());
out
}
}