diff --git a/build.rs b/build.rs index 33d2ece..a5602d2 100644 --- a/build.rs +++ b/build.rs @@ -19,47 +19,100 @@ use std::env; use std::fs; use std::path::PathBuf; -/// All models documented in `ceva-ip/DPDFNet@main:package/src/dpdfnet/models.py`. -/// Entries: `(registry_name, sample_rate_hz, description)`. +/// All models documented in `ceva-ip/DPDFNet@main:package/src/dpdfnet/models.py`, +/// plus the split-band variants of the 16 kHz networks. /// -/// There is deliberately no per-model block size here any more. Inference -/// runs on a worker thread, so the audio callback does the FFT, the -/// overlap-add and two buffer copies and nothing else — its cost no longer -/// depends on the model. `tests/callback_deadline.rs` holds every model to -/// every block from 10 ms up. +/// Entries: `(registry_name, model_dir, model_rate_hz, stft_rate_hz, description)`. /// -/// What still varies per model is whether the worker keeps up: 3.3 ms of -/// inference per 10 ms hop for DPDFNet-8, a third of a core. That is a -/// CPU-load question for the quality tier, and the plugin reports it on the -/// `Hops Total` and `Hops Enhanced` control ports rather than promising it -/// here. -const REGISTRY: &[(&str, usize, &str)] = &[ +/// `model_dir` is where the IR files live under `model/`; it differs from +/// `registry_name` only for the split-band variants, which reuse the base +/// 16 kHz model's IR unchanged. `model_rate` is the rate the network was +/// trained at — it fixes the tensor geometry and the magnitude the model +/// expects. `stft_rate` is the rate the wrapper runs its STFT at, which is also +/// the host rate the plugin accepts. +/// +/// A split-band variant pairs a 48 kHz `stft_rate` with a 16 kHz `model_rate`: +/// the STFT runs at 48 kHz (nfft 960), the network is fed the first 161 bins +/// (0–8 kHz — the exact frequencies it was trained on, once scaled to 16 kHz +/// magnitudes), and the wrapper reconstructs the band above 8 kHz. The result +/// is full-band output with no resampler, instead of the telephone-band output +/// the pure 16 kHz plugin gives inside a 48 kHz graph. +/// +/// There is deliberately no per-model block size here. Inference runs on a +/// worker thread, so the audio callback does the FFT, the overlap-add and two +/// buffer copies and nothing else — its cost no longer depends on the model. +/// `tests/callback_deadline.rs` holds every model to every block from 10 ms up. +const REGISTRY: &[(&str, &str, usize, usize, &str)] = &[ ( + "baseline", "baseline", 16_000, + 16_000, "DPDFNet 16 kHz baseline (fastest, lowest compute)", ), ( + "dpdfnet2", "dpdfnet2", 16_000, + 16_000, "DPDFNet-2 16 kHz (balanced quality/speed)", ), - ("dpdfnet4", 16_000, "DPDFNet-4 16 kHz (higher quality)"), ( + "dpdfnet4", + "dpdfnet4", + 16_000, + 16_000, + "DPDFNet-4 16 kHz (higher quality)", + ), + ( + "dpdfnet8", "dpdfnet8", 16_000, + 16_000, "DPDFNet-8 16 kHz (highest quality 16 kHz, offline only)", ), ( "dpdfnet2_48khz_hr", + "dpdfnet2_48khz_hr", + 48_000, 48_000, "DPDFNet-2 48 kHz hi-res (full-band, balanced)", ), ( "dpdfnet8_48khz_hr", + "dpdfnet8_48khz_hr", + 48_000, 48_000, "DPDFNet-8 48 kHz hi-res (full-band, highest quality, offline only)", ), + ( + "baseline_sb", + "baseline", + 16_000, + 48_000, + "DPDFNet baseline split-band (full-band, lowest compute)", + ), + ( + "dpdfnet2_sb", + "dpdfnet2", + 16_000, + 48_000, + "DPDFNet-2 split-band (full-band, balanced, no resampler)", + ), + ( + "dpdfnet4_sb", + "dpdfnet4", + 16_000, + 48_000, + "DPDFNet-4 split-band (full-band, higher quality)", + ), + ( + "dpdfnet8_sb", + "dpdfnet8", + 16_000, + 48_000, + "DPDFNet-8 split-band (full-band, highest quality, offline only)", + ), ]; fn fnv1a64(s: &str) -> u64 { @@ -84,18 +137,21 @@ fn main() { let entry = REGISTRY .iter() - .find(|(name, _, _)| *name == model_name) + .find(|(name, ..)| *name == model_name) .unwrap_or_else(|| { - let known: Vec<&str> = REGISTRY.iter().map(|(n, _, _)| *n).collect(); + let known: Vec<&str> = REGISTRY.iter().map(|(n, ..)| *n).collect(); panic!( "unknown DPDFNet model `{model_name}`; valid choices: {}", known.join(", ") ); }); - let (name, sample_rate, description) = (entry.0, entry.1, entry.2); + let (name, model_subdir, model_rate, stft_rate, description) = + (entry.0, entry.1, entry.2, entry.3, entry.4); let manifest_dir = env::var("CARGO_MANIFEST_DIR").expect("CARGO_MANIFEST_DIR unset"); - let model_dir = PathBuf::from(&manifest_dir).join("model").join(name); + let model_dir = PathBuf::from(&manifest_dir) + .join("model") + .join(model_subdir); let xml_path = model_dir.join("model.xml"); let bin_path = model_dir.join("model.bin"); let state_path = model_dir.join("init_state.bin"); @@ -109,10 +165,32 @@ fn main() { println!("cargo:rerun-if-changed={}", path.display()); } - // STFT geometry follows directly from sample_rate at frame_ms=20.0. - let win_len = sample_rate / 50; + // Two geometries. The MODEL band is fixed by the rate the network was + // trained at (frame_ms = 20.0): its win/hop/bins set the ONNX tensor shape + // and never change for a given IR. The STFT the wrapper actually runs is at + // `stft_rate` — equal to the model rate for a normal model, 48 kHz for a + // split-band variant. When they differ, the audio callback runs the bigger + // FFT, feeds the model the low `freq_bins`, and reconstructs the rest. + let win_len = model_rate / 50; let hop_size = win_len / 2; let freq_bins = win_len / 2 + 1; + let split_band = stft_rate != model_rate; + let stft_win_len = stft_rate / 50; + let stft_hop = stft_win_len / 2; + let stft_freq_bins = stft_win_len / 2 + 1; + // `SPECTRUM_SCALE` is `win_len / stft_win_len`: the magnitude the model + // expects vs. what a `stft_win_len` FFT produces for the same tone, since + // the window sum scales with its length. The low bins are divided by it + // going in and multiplied coming out. Emitted as a literal `1.0` for a + // normal build (a `320.0 / 320.0` would trip `clippy::eq_op`) and as the + // exact division for a split-band build (1/3 has no finite decimal form). + // Verified against the ONNX: the scaled 48 kHz low bins match native 16 kHz + // to −39 dB. + let spectrum_scale = if split_band { + format!("{win_len}.0 / {stft_win_len}.0") + } else { + "1.0".to_string() + }; // Recurrent state size = init_state.bin length / 4 (f32 little-endian). let state_bytes = fs::metadata(&state_path) @@ -136,11 +214,20 @@ fn main() { r#"// AUTO-GENERATED by build.rs. Do not edit. pub const MODEL_NAME: &str = "{name}"; -pub const SAMPLE_RATE: usize = {sample_rate}; +pub const SAMPLE_RATE: usize = {model_rate}; pub const WIN_LEN: usize = {win_len}; pub const HOP_SIZE: usize = {hop_size}; pub const FREQ_BINS: usize = {freq_bins}; pub const STATE_SIZE: usize = {state_size}; + +// Host / STFT geometry. Equal to the model geometry for a normal build; for a +// split-band variant the host runs at 48 kHz while FREQ_BINS stays the model's. +pub const HOST_SAMPLE_RATE: usize = {stft_rate}; +pub const STFT_WIN_LEN: usize = {stft_win_len}; +pub const STFT_HOP: usize = {stft_hop}; +pub const STFT_FREQ_BINS: usize = {stft_freq_bins}; +pub const SPLIT_BAND: bool = {split_band}; +pub const SPECTRUM_SCALE: f32 = {spectrum_scale}; pub const LADSPA_LABEL: &str = "{label}"; pub const LADSPA_NAME: &str = "{display}"; pub const LADSPA_UNIQUE_ID: u64 = {unique_id}; diff --git a/pkgbuild/PKGBUILD b/pkgbuild/PKGBUILD index 297ff37..67d178e 100644 --- a/pkgbuild/PKGBUILD +++ b/pkgbuild/PKGBUILD @@ -14,6 +14,10 @@ pkgname=( 'dpdfnet-ladspa-dpdfnet8' 'dpdfnet-ladspa-dpdfnet2-48khz-hr' 'dpdfnet-ladspa-dpdfnet8-48khz-hr' + 'dpdfnet-ladspa-baseline-sb' + 'dpdfnet-ladspa-dpdfnet2-sb' + 'dpdfnet-ladspa-dpdfnet4-sb' + 'dpdfnet-ladspa-dpdfnet8-sb' ) pkgver=$(date +%y.%m.%d) pkgrel=$(date +%H%M) @@ -71,6 +75,17 @@ _models=( 'dpdfnet8_48khz_hr' ) +# Split-band variants. They reuse the 16 kHz models' IR, so they need no ONNX +# of their own and no prepare() step — only a build(), after the base models +# above have populated `model//`. Kept out of `_models` because that +# array is index-aligned with source=/sha256sums=. +_split_band=( + 'baseline_sb' + 'dpdfnet2_sb' + 'dpdfnet4_sb' + 'dpdfnet8_sb' +) + source=( # Repo itself. Build runners (BigLinux build-package, makechrootpkg) # only copy PKGBUILD into the chroot, so the crate sources must @@ -136,7 +151,7 @@ build() { # `cargo:rerun-if-env-changed=DPDFNET_MODEL`, but separate target # trees make the per-model artifact location unambiguous for the # package_*() install steps below). - for model in "${_models[@]}"; do + for model in "${_models[@]}" "${_split_band[@]}"; do msg "Building libdpdfnet_${model}_ladspa.so" DPDFNET_MODEL="${model}" \ CARGO_TARGET_DIR="target-${model}" \ @@ -213,3 +228,29 @@ package_dpdfnet-ladspa-dpdfnet8-48khz-hr() { depends=("${_split_depends[@]}") _install_model dpdfnet8_48khz_hr "${pkgname}" } + +# Split-band variants: the 16 kHz networks run at 48 kHz with the band above +# 8 kHz reconstructed, so they are full-band and need no resampler in the graph. +package_dpdfnet-ladspa-baseline-sb() { + pkgdesc='DPDFNet baseline split-band LADSPA plugin (full-band, lowest compute)' + depends=("${_split_depends[@]}") + _install_model baseline_sb "${pkgname}" +} + +package_dpdfnet-ladspa-dpdfnet2-sb() { + pkgdesc='DPDFNet-2 split-band LADSPA plugin (full-band, balanced, no resampler)' + depends=("${_split_depends[@]}") + _install_model dpdfnet2_sb "${pkgname}" +} + +package_dpdfnet-ladspa-dpdfnet4-sb() { + pkgdesc='DPDFNet-4 split-band LADSPA plugin (full-band, higher quality)' + depends=("${_split_depends[@]}") + _install_model dpdfnet4_sb "${pkgname}" +} + +package_dpdfnet-ladspa-dpdfnet8-sb() { + pkgdesc='DPDFNet-8 split-band LADSPA plugin (full-band, highest quality)' + depends=("${_split_depends[@]}") + _install_model dpdfnet8_sb "${pkgname}" +} diff --git a/src/highband.rs b/src/highband.rs new file mode 100644 index 0000000..8b75d79 --- /dev/null +++ b/src/highband.rs @@ -0,0 +1,198 @@ +//! High-band handling for the split-band variants. +//! +//! The 16 kHz DPDFNet networks only see 0–8 kHz. Run inside a 48 kHz graph +//! through this wrapper, the STFT is 48 kHz (nfft 960) and the network is fed +//! the low 161 bins; everything above 8 kHz (bins 161..481) never reaches the +//! model. This passes that captured band through, gated by the speech +//! probability and a per-bin SNR gate so background noise above 8 kHz is +//! attenuated while the real high band — fricatives, air — is kept. A short +//! bypass after a transient keeps consonants ("t", "s", "k") from being gated. +//! +//! It deliberately does **not** synthesize the high band. An earlier port +//! carried the GTCRN air exciter, which mirrors 4–8 kHz up and squares it for a +//! second harmonic when the high band's SNR is low. Measured on speech, that +//! fired on 68–89 % of frames — even on clean speech, because voice carries +//! almost no energy above 8 kHz outside fricatives, so the SNR is low most of +//! the time. That meant the "full-band" output was mostly fabricated, not the +//! captured voice, which is the opposite of what a natural-capture denoiser +//! should do. A synthetic highs effect, if ever wanted, belongs in its own +//! opt-in stage, not here. +//! +//! Ported from the GTCRN LADSPA wrapper (BigLinux, MIT). + +use realfft::num_complex::Complex; + +/// Passes and gates the band the model does not see. +pub struct HighBand { + /// Previous-frame high-band magnitudes, for spectral-flux transient + /// detection. + prev_mag: Vec, + has_prev: bool, + /// Frames of gate bypass left after a transient, and the full hold length. + transient_hold: usize, + transient_hold_frames: usize, + /// Adaptive per-bin noise floor, learned during silence, for the SNR gate. + noise_floor: Vec, + noise_initialised: bool, +} + +impl HighBand { + /// `hf_bins` is the number of bins above the model cutoff (STFT bins minus + /// model bins). + pub fn new(hf_bins: usize, hop: usize, host_rate: usize) -> Self { + let frame_dur_s = hop as f64 / host_rate as f64; + let hold_frames = (0.020 / frame_dur_s).ceil() as usize; + Self { + prev_mag: vec![0.0; hf_bins], + has_prev: false, + transient_hold: 0, + transient_hold_frames: hold_frames.max(1), + noise_floor: vec![1e-6; hf_bins], + noise_initialised: false, + } + } + + /// Gate the captured high band into `output`. `speech` is 0.0 (silence) .. + /// 1.0 (speech). `hf_original` and `output` are the STFT bins above the + /// model cutoff. + pub fn process( + &mut self, + hf_original: &[Complex], + speech: f32, + output: &mut [Complex], + ) { + self.update_noise_floor(hf_original, speech); + + if self.detect_transient(hf_original) { + self.transient_hold = self.transient_hold_frames; + } else if self.transient_hold > 0 { + self.transient_hold -= 1; + } + + self.spectral_gate(hf_original, speech, output); + } + + /// Half-wave-rectified spectral flux: a transient spikes positive flux + /// across many bins at once. + fn detect_transient(&mut self, hf: &[Complex]) -> bool { + let count = hf.len().min(self.prev_mag.len()); + if count == 0 { + return false; + } + let mut flux = 0.0_f32; + let mut avg_mag = 0.0_f32; + for (h, prev) in hf[..count].iter().zip(self.prev_mag[..count].iter_mut()) { + let mag = h.norm(); + avg_mag += mag; + let diff = mag - *prev; + if diff > 0.0 { + flux += diff; + } + *prev = mag; + } + avg_mag /= count as f32; + let was_valid = self.has_prev; + self.has_prev = true; + if !was_valid || avg_mag < 1e-10 { + return false; + } + flux / (avg_mag * count as f32) > 0.35 + } + + /// Pass the original high band through, gated by speech probability and a + /// per-bin SNR gate; bypass the gate briefly after a transient. + fn spectral_gate(&self, hf: &[Complex], speech: f32, output: &mut [Complex]) { + let bypassed = self.transient_hold > 0; + let bypass_floor = if bypassed { 0.5 } else { 0.0 }; + let envelope = if bypassed || speech > 0.7 { + 1.0 + } else if speech > 0.1 { + (speech - 0.1) / 0.6 + } else { + 0.0 + }; + let count = hf.len().min(output.len()); + let snr_threshold = 1.0; + let snr_range = 3.5; + + if self.noise_initialised { + for i in 0..count { + let snr = hf[i].norm() / (self.noise_floor[i] + 1e-10); + let snr_gain = ((snr - snr_threshold) / snr_range).clamp(0.0, 1.0); + let gain = if bypassed { + snr_gain.max(bypass_floor) + } else { + envelope * snr_gain + }; + output[i] = Complex::new(hf[i].re * gain, hf[i].im * gain); + } + } else { + let atten = envelope * 0.5; + for i in 0..count { + output[i] = Complex::new(hf[i].re * atten, hf[i].im * atten); + } + } + for o in output.iter_mut().skip(count) { + *o = Complex::new(0.0, 0.0); + } + } + + /// Learn the per-bin high-band noise floor while there is no speech. + fn update_noise_floor(&mut self, hf: &[Complex], speech: f32) { + if speech >= 0.1 { + return; + } + let count = hf.len().min(self.noise_floor.len()); + for (h, floor) in hf[..count].iter().zip(self.noise_floor[..count].iter_mut()) { + let mag = h.norm(); + *floor = if self.noise_initialised { + 0.99 * *floor + 0.01 * mag + } else { + mag + }; + } + self.noise_initialised = true; + } +} + +/// Smoothed speech probability from the low band's input vs. enhanced energy. +/// +/// The high-band gate needs a 0..1 speech gate; the DPDFNet network does not +/// emit one, so derive it the way the GTCRN wrapper does — from energy. The +/// denoiser keeps speech and removes noise, so the enhanced-to-input energy +/// ratio is high during voice and low in a pause; an onset over a tracked input +/// floor primes the gate so the first syllable is not gated out. +pub struct SpeechGate { + input_floor: f32, + gate: f32, +} + +impl SpeechGate { + pub fn new() -> Self { + Self { + input_floor: 1e-6, + gate: 0.0, + } + } + + /// One hop. `input_energy` and `enhanced_energy` are sums of squared + /// magnitude over the low band. Returns the 0..1 gate. + pub fn update(&mut self, input_energy: f32, enhanced_energy: f32) -> f32 { + let ratio = enhanced_energy / (input_energy + 1e-9); + let onset = input_energy > self.input_floor * 2.5; + let speaking = onset && ratio > 0.05; + if speaking { + self.gate = 1.0; + } else { + self.gate *= 0.95; + if self.gate < 0.01 { + self.gate = 0.0; + } + } + // Track the input floor only when confidently silent. + if self.gate < 0.1 { + self.input_floor = 0.95 * self.input_floor + 0.05 * input_energy; + } + self.gate + } +} diff --git a/src/lib.rs b/src/lib.rs index ad04346..77b02e3 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -38,6 +38,7 @@ use realfft::num_complex::Complex; use realfft::{ComplexToReal, RealFftPlanner, RealToComplex}; mod engine; +mod highband; use engine::{Inference, MAX_DEPTH}; @@ -154,7 +155,10 @@ pub unsafe extern "C" fn dpdfnet_hops(total: *mut u64, enhanced: *mut u64) { /// delivered and wrong again by however far the block sat from 20 ms. #[no_mangle] pub extern "C" fn dpdfnet_added_latency_frames() -> u32 { - ((MODEL_DELAY_HOPS + EMIT_LAG_HOPS) * model_const::HOP_SIZE + model_const::WIN_LEN) as u32 + // Expressed in host samples: on a split-band build the hop and window are + // the 48 kHz STFT's, not the 16 kHz model's, so the published figure is the + // audio the graph actually sees delayed. + ((MODEL_DELAY_HOPS + EMIT_LAG_HOPS) * model_const::STFT_HOP + model_const::STFT_WIN_LEN) as u32 } /// The host sample rate this build's model requires. A filter chain running @@ -200,6 +204,35 @@ fn vorbis_window(win_len: usize) -> Vec { .collect() } +/// Raised-cosine blend across the four bins straddling the model / high-band +/// boundary, so the enhanced low band and the reconstructed high band do not +/// meet at a hard edge. A no-op when there is no high band (`low_bins == +/// spectrum.len()`, the non-split build). +fn crossfade_boundary(spectrum: &mut [Complex], low_bins: usize) { + const WIDTH: usize = 4; + if low_bins < WIDTH / 2 || low_bins + WIDTH / 2 > spectrum.len() { + return; + } + let base = low_bins - WIDTH / 2; + let snapshot: [Complex; WIDTH] = std::array::from_fn(|k| spectrum[base + k]); + let low = [snapshot[0], snapshot[1]]; // pure model side + let high = [snapshot[2], snapshot[3]]; // pure high-band side + for k in 0..WIDTH { + let t = (k as f32 + 0.5) / WIDTH as f32; + let w = 0.5 * (1.0 - (std::f32::consts::PI * t).cos()); + let model_val = if k < WIDTH / 2 { low[k] } else { low[1] }; + let hf_val = if k >= WIDTH / 2 { + high[k - WIDTH / 2] + } else { + high[0] + }; + spectrum[base + k] = Complex::new( + model_val.re * (1.0 - w) + hf_val.re * w, + model_val.im * (1.0 - w) + hf_val.im * w, + ); + } +} + struct DpdfnetPlugin { /// Runs the model on a worker thread, `depth` hops behind the audio. inference: Inference, @@ -258,9 +291,20 @@ struct DpdfnetPlugin { fft_fwd: Arc>, fft_inv: Arc>, fft_real: Vec, + /// Forward-FFT result: the original noisy spectrum, `STFT_FREQ_BINS` long. + /// For a split-band build this is wider than the model band, and the bins + /// above the model cutoff are the high band the model never sees. fft_complex: Vec>, + /// Enhanced spectrum assembled for the inverse FFT: the model's low band + /// scaled back up, plus the reconstructed high band. Same width as + /// `fft_complex`; equal to it bin-for-bin on a non-split build. + out_complex: Vec>, spec_in: Vec, spec_out: Vec, + /// High-band reconstruction and its speech gate — only a split-band build + /// carries them; a normal build's high band is empty. + highband: Option, + speech_gate: highband::SpeechGate, } impl DpdfnetPlugin { @@ -268,16 +312,18 @@ impl DpdfnetPlugin { let sr = sample_rate as usize; // A wrong host rate used to abort the process. LADSPA cannot refuse // an instantiation, so refuse the model instead: `rate_ok` false - // means the engine is never asked for and audio passes through. - let rate_ok = sr == model_const::SAMPLE_RATE; + // means the engine is never asked for and audio passes through. The + // rate the plugin accepts is the STFT/host rate — 48 kHz for a + // split-band build, whose model band stays 16 kHz internally. + let rate_ok = sr == model_const::HOST_SAMPLE_RATE; if !rate_ok { eprintln!( "[{}] host sample rate is {sr} Hz, this plugin needs {}; \ passing audio through unprocessed. Set `audio.rate = {}` \ on the filter-chain node.", model_const::LADSPA_LABEL, - model_const::SAMPLE_RATE, - model_const::SAMPLE_RATE + model_const::HOST_SAMPLE_RATE, + model_const::HOST_SAMPLE_RATE ); } @@ -286,9 +332,20 @@ impl DpdfnetPlugin { // one, so an idle chain does not pay the OpenVINO Core and JIT cost; // when audio starts, the first frames pass through clean until the // build lands (170-460 ms depending on the model). + // The STFT runs at the host geometry (== the model geometry on a + // normal build, wider on a split-band one). The model band stays + // `FREQ_BINS` and is extracted from the low end of this spectrum. let mut planner = RealFftPlanner::::new(); - let fft_fwd = planner.plan_fft_forward(model_const::WIN_LEN); - let fft_inv = planner.plan_fft_inverse(model_const::WIN_LEN); + let fft_fwd = planner.plan_fft_forward(model_const::STFT_WIN_LEN); + let fft_inv = planner.plan_fft_inverse(model_const::STFT_WIN_LEN); + let hf_bins = model_const::STFT_FREQ_BINS - model_const::FREQ_BINS; + let highband = model_const::SPLIT_BAND.then(|| { + highband::HighBand::new( + hf_bins, + model_const::STFT_HOP, + model_const::HOST_SAMPLE_RATE, + ) + }); Self { inference: Inference::new(), @@ -318,40 +375,47 @@ impl DpdfnetPlugin { // periodic clicks / robotic timbre at the host quantum // rate). Latency cost: WIN_LEN/sr ≈ 20 ms. in_buf: { - // Primed with WIN_LEN zeros but reserved for a whole block on - // top: the first `run()` appends before it consumes, and that - // first callback is exactly the one this plugin must not + // Primed with STFT_WIN_LEN zeros but reserved for a whole block + // on top: the first `run()` appends before it consumes, and + // that first callback is exactly the one this plugin must not // stall in. - let mut buf = Vec::with_capacity(MAX_HOST_BLOCK + model_const::WIN_LEN); - buf.resize(model_const::WIN_LEN, 0.0); + let mut buf = Vec::with_capacity(MAX_HOST_BLOCK + model_const::STFT_WIN_LEN); + buf.resize(model_const::STFT_WIN_LEN, 0.0); buf }, - ola_buf: vec![0.0; model_const::WIN_LEN], - out_queue: VecDeque::with_capacity(MAX_HOST_BLOCK + model_const::WIN_LEN), - window: vorbis_window(model_const::WIN_LEN), - fft_real: vec![0.0; model_const::WIN_LEN], - fft_complex: vec![Complex::new(0.0, 0.0); model_const::FREQ_BINS], + ola_buf: vec![0.0; model_const::STFT_WIN_LEN], + out_queue: VecDeque::with_capacity(MAX_HOST_BLOCK + model_const::STFT_WIN_LEN), + window: vorbis_window(model_const::STFT_WIN_LEN), + fft_real: vec![0.0; model_const::STFT_WIN_LEN], + fft_complex: vec![Complex::new(0.0, 0.0); model_const::STFT_FREQ_BINS], + out_complex: vec![Complex::new(0.0, 0.0); model_const::STFT_FREQ_BINS], fft_fwd, fft_inv, spec_in: vec![0.0; model_const::FREQ_BINS * 2], spec_out: vec![0.0; model_const::FREQ_BINS * 2], + highband, + speech_gate: highband::SpeechGate::new(), } } /// Process exactly one analysis frame: window + FFT + ONNX + - /// spectral blend + iFFT + windowed OLA + flush HOP samples to - /// `out_queue`. Caller guarantees `in_buf.len() >= WIN_LEN`. + /// spectral blend + high-band reconstruction + iFFT + windowed OLA + flush + /// HOP samples to `out_queue`. Caller guarantees + /// `in_buf.len() >= STFT_WIN_LEN`. fn process_frame(&mut self, alpha: f32) { - for j in 0..model_const::WIN_LEN { + for j in 0..model_const::STFT_WIN_LEN { self.fft_real[j] = self.in_buf[j] * self.window[j]; } let _ = self .fft_fwd .process(&mut self.fft_real, &mut self.fft_complex); - for (k, c) in self.fft_complex.iter().enumerate() { - self.spec_in[k * 2] = c.re; - self.spec_in[k * 2 + 1] = c.im; + // The model only ever sees the low `FREQ_BINS` bins. On a split-band + // build the STFT is wider, so scale those bins to the magnitude the + // 16 kHz network was trained on (`SPECTRUM_SCALE` is 1.0 otherwise). + for k in 0..model_const::FREQ_BINS { + self.spec_in[k * 2] = self.fft_complex[k].re * model_const::SPECTRUM_SCALE; + self.spec_in[k * 2 + 1] = self.fft_complex[k].im * model_const::SPECTRUM_SCALE; } HOPS_TOTAL.fetch_add(1, Ordering::Relaxed); @@ -398,34 +462,59 @@ impl DpdfnetPlugin { } } + // Assemble the enhanced full-width spectrum for the inverse FFT. The + // model's low band is scaled back up to the host magnitude; on a + // split-band build the band above the model cutoff is reconstructed + // from the original high band and the clean low band. + let inv_scale = 1.0 / model_const::SPECTRUM_SCALE; for k in 0..model_const::FREQ_BINS { - self.fft_complex[k] = Complex::new(self.spec_out[k * 2], self.spec_out[k * 2 + 1]); + self.out_complex[k] = Complex::new( + self.spec_out[k * 2] * inv_scale, + self.spec_out[k * 2 + 1] * inv_scale, + ); } + if let Some(highband) = self.highband.as_mut() { + let input_energy: f32 = (0..model_const::FREQ_BINS) + .map(|k| self.spec_in[k * 2].powi(2) + self.spec_in[k * 2 + 1].powi(2)) + .sum(); + let enhanced_energy: f32 = (0..model_const::FREQ_BINS) + .map(|k| self.spec_out[k * 2].powi(2) + self.spec_out[k * 2 + 1].powi(2)) + .sum(); + let speech = self.speech_gate.update(input_energy, enhanced_energy); + + highband.process( + &self.fft_complex[model_const::FREQ_BINS..], + speech, + &mut self.out_complex[model_const::FREQ_BINS..], + ); + crossfade_boundary(&mut self.out_complex, model_const::FREQ_BINS); + } + let _ = self .fft_inv - .process(&mut self.fft_complex, &mut self.fft_real); + .process(&mut self.out_complex, &mut self.fft_real); // realfft inverse leaves a 1/N scaling — fold into the synthesis // window so OLA gets the correct amplitude. - let scale = 1.0 / model_const::WIN_LEN as f32; - for j in 0..model_const::WIN_LEN { + let scale = 1.0 / model_const::STFT_WIN_LEN as f32; + for j in 0..model_const::STFT_WIN_LEN { self.ola_buf[j] += self.fft_real[j] * scale * self.window[j]; } // First HOP samples of the OLA accumulator are now stable. - for j in 0..model_const::HOP_SIZE { + for j in 0..model_const::STFT_HOP { self.out_queue.push_back(self.ola_buf[j]); } self.ola_buf - .copy_within(model_const::HOP_SIZE..model_const::WIN_LEN, 0); - for j in (model_const::WIN_LEN - model_const::HOP_SIZE)..model_const::WIN_LEN { + .copy_within(model_const::STFT_HOP..model_const::STFT_WIN_LEN, 0); + for j in (model_const::STFT_WIN_LEN - model_const::STFT_HOP)..model_const::STFT_WIN_LEN { self.ola_buf[j] = 0.0; } // Slide analysis window forward by HOP samples. - self.in_buf.copy_within(model_const::HOP_SIZE.., 0); + self.in_buf.copy_within(model_const::STFT_HOP.., 0); self.in_buf - .truncate(self.in_buf.len() - model_const::HOP_SIZE); + .truncate(self.in_buf.len() - model_const::STFT_HOP); } /// Fill `spec_out` with the answer for hop `due`: the enhanced spectrum if @@ -581,7 +670,7 @@ impl Plugin for DpdfnetPlugin { // itself back off and sit at a few per cent enhanced. let now = Instant::now(); let started = *self.last_run.get_or_insert(now); - self.audio_seen += n as f64 / model_const::SAMPLE_RATE as f64; + self.audio_seen += n as f64 / model_const::HOST_SAMPLE_RATE as f64; if !self.offline && self.audio_seen > OFFLINE_WARMUP_S { let elapsed = now.duration_since(started).as_secs_f64(); self.offline = self.audio_seen > elapsed * 4.0; @@ -591,21 +680,23 @@ impl Plugin for DpdfnetPlugin { // the callback's own start: a hop that spends it all leaves the rest of // the block emitting noisy frames rather than pushing the whole graph // past its deadline. - self.wait_until = - now + Duration::from_secs_f64(n as f64 / model_const::SAMPLE_RATE as f64 * WAIT_BUDGET); + self.wait_until = now + + Duration::from_secs_f64( + n as f64 / model_const::HOST_SAMPLE_RATE as f64 * WAIT_BUDGET, + ); // One callback carries this many analysis hops, back to back with no // wall clock between them, so that is how many submissions the engine // has to accept before refusing. It used to be the emitted hop's lag // too, which is what made the latency the host's block. - let hops_per_callback = n.div_ceil(model_const::HOP_SIZE).max(1); + let hops_per_callback = n.div_ceil(model_const::STFT_HOP).max(1); if hops_per_callback != self.depth { self.depth = hops_per_callback; self.inference.set_depth(hops_per_callback); } self.in_buf.extend_from_slice(&input[..n]); - while self.in_buf.len() >= model_const::WIN_LEN { + while self.in_buf.len() >= model_const::STFT_WIN_LEN { self.process_frame(alpha); } @@ -678,7 +769,7 @@ mod tests { fn no_callback_grows_a_buffer() { let descriptor = get_ladspa_descriptor(0).expect("descriptor"); let ports = descriptor.ports.clone(); - let mut plugin = DpdfnetPlugin::new(MODEL_SAMPLE_RATE as u64); + let mut plugin = DpdfnetPlugin::new(model_const::HOST_SAMPLE_RATE as u64); plugin.activate(); // The largest block we reserve for, driven at once so the first