use std::collections::VecDeque;
use crate::fft::Complex;
use crate::gain::{compute_gain, multiband_specsub_gains, Algorithm, GainParams, SpecSubLaw};
use crate::noise::{NoiseConfig, NoiseEstimator};
use crate::perceptual::{apply_perceptual_weights, bin_to_bark_band, N_BARK_BANDS};
use crate::postfilter::{MusicalNoisePostFilter, PostFilterConfig};
use crate::stft::{Stft, StftConfig};
use crate::window::{WindowParams, WindowType};
#[derive(Clone, Debug)]
pub struct DenoiserConfig {
pub algorithm: Algorithm,
pub strength: f64,
pub frame_size: usize,
pub overlap: f64,
pub window: WindowType,
pub profile_ms: f64,
pub adapt: bool,
pub adaptive_noise: bool,
pub vad: bool,
pub vad_silence_gain: f64,
pub vad_speech_mix: f64,
pub smoothing: f64,
pub dc_block: bool,
pub makeup_gain_db: f64,
pub sample_rate: u32,
pub transient_protect: bool,
pub cepstral_smoothing: bool,
pub pre_emphasis: bool,
pub pre_emphasis_alpha: f64,
pub window_params: WindowParams,
pub multiband: bool,
pub perceptual_weighting: bool,
pub musical_noise_postfilter: bool,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Preset {
Speech,
Music,
Aggressive,
Gentle,
Restore,
HiFi,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ProcessingMode {
Speech,
Music,
Ambient,
}
impl ProcessingMode {
pub fn parse(value: &str) -> Option<Self> {
match value.to_ascii_lowercase().as_str() {
"speech" | "voice" => Some(Self::Speech),
"music" => Some(Self::Music),
"ambient" | "environment" => Some(Self::Ambient),
_ => None,
}
}
pub fn apply(self, config: &mut DenoiserConfig) {
match self {
Self::Speech => {
config.strength = config.strength.max(0.7);
config.vad = true;
config.adaptive_noise = true;
config.transient_protect = true;
config.cepstral_smoothing = true;
}
Self::Music => {
config.strength = config.strength.min(0.35);
config.vad = false;
config.adaptive_noise = false;
config.transient_protect = true;
config.perceptual_weighting = true;
config.musical_noise_postfilter = true;
config.smoothing = config.smoothing.max(0.75);
}
Self::Ambient => {
config.strength = config.strength.min(0.4);
config.vad = false;
config.adaptive_noise = true;
config.transient_protect = true;
config.perceptual_weighting = true;
config.smoothing = config.smoothing.max(0.7);
}
}
}
}
impl Preset {
pub fn parse(s: &str) -> Option<Self> {
Some(match s.to_ascii_lowercase().as_str() {
"speech" | "voice" => Preset::Speech,
"music" => Preset::Music,
"aggressive" => Preset::Aggressive,
"gentle" => Preset::Gentle,
"restore" => Preset::Restore,
"hifi" | "mastering" | "hi-fi" | "highfidelity" => Preset::HiFi,
_ => return None,
})
}
pub fn config(self, sample_rate: u32) -> DenoiserConfig {
let mut c = DenoiserConfig::default(sample_rate);
match self {
Preset::Speech => {
c.algorithm = Algorithm::Omlsa;
c.strength = 0.6;
c.frame_size = 2048;
c.smoothing = 0.6;
}
Preset::Music => {
c.algorithm = Algorithm::Omlsa;
c.strength = 0.4;
c.frame_size = 4096;
c.smoothing = 0.5;
c.overlap = 0.8;
c.transient_protect = true;
c.cepstral_smoothing = true;
c.perceptual_weighting = true;
c.musical_noise_postfilter = true;
}
Preset::Aggressive => {
c.algorithm = Algorithm::Omlsa;
c.strength = 0.85;
c.frame_size = 2048;
c.smoothing = 0.72;
}
Preset::Gentle => {
c.algorithm = Algorithm::LogMmse;
c.strength = 0.3;
c.frame_size = 2048;
c.smoothing = 0.45;
}
Preset::Restore => {
c.algorithm = Algorithm::LogMmse;
c.strength = 0.2;
c.frame_size = 2048;
c.smoothing = 0.4;
}
Preset::HiFi => {
c.algorithm = Algorithm::Omlsa;
c.strength = 0.28;
c.frame_size = 4096;
c.overlap = 0.875;
c.window = WindowType::Kaiser;
c.window_params.kaiser_beta = 10.0;
c.smoothing = 0.65;
c.transient_protect = true;
c.cepstral_smoothing = true;
c.perceptual_weighting = true;
c.musical_noise_postfilter = true;
c.pre_emphasis = false;
c.pre_emphasis_alpha = 0.72;
}
}
c
}
}
impl DenoiserConfig {
pub fn default(sample_rate: u32) -> Self {
DenoiserConfig {
algorithm: Algorithm::Omlsa,
strength: 0.6,
frame_size: 2048,
overlap: 0.75,
window: WindowType::Hann,
profile_ms: 0.0,
adapt: true,
adaptive_noise: false,
vad: false,
vad_silence_gain: 0.08,
vad_speech_mix: 0.85,
smoothing: 0.6,
dc_block: true,
makeup_gain_db: 0.0,
sample_rate,
transient_protect: true,
cepstral_smoothing: false, pre_emphasis: false,
pre_emphasis_alpha: 0.92,
window_params: WindowParams::default(),
multiband: false,
perceptual_weighting: false,
musical_noise_postfilter: false,
}
}
pub fn sanitized(mut self) -> Self {
self.strength = self.strength.clamp(0.0, 1.0);
self.smoothing = self.smoothing.clamp(0.0, 0.95);
self.overlap = self.overlap.clamp(0.5, 0.95);
self.vad_silence_gain = self.vad_silence_gain.clamp(0.0, 1.0);
self.vad_speech_mix = self.vad_speech_mix.clamp(0.0, 1.0);
if !self.frame_size.is_power_of_two() || self.frame_size < 256 {
self.frame_size = 2048;
}
self.pre_emphasis_alpha = self.pre_emphasis_alpha.clamp(0.0, 0.99);
self
}
}
pub struct Denoiser {
config: DenoiserConfig,
stft: Stft,
noise: NoiseEstimator,
noise_cfg: NoiseConfig,
gain_params: GainParams,
sample_rate: u32,
frame_size: usize,
hop: usize,
m: usize, alpha_dd: f64,
xi_min: f64,
makeup: f64,
prev_g: Vec<f64>,
prev_y2: Vec<f64>,
prev_lambda_d: Vec<f64>,
prev_gsmooth: Vec<f64>,
spec: Vec<Complex>,
frame: Vec<f64>,
y2: Vec<f64>,
g: Vec<f64>,
lambda_d_buf: Vec<f64>,
spp_buf: Vec<f64>,
y2_snapshot_buf: Vec<f64>,
prev_frame_energy: f64,
prev_mag: Vec<f64>, pre_emph_prev: f64, de_emph_prev: f64, dc_prev_x: f64,
dc_prev_y: f64,
cepstral_fft: crate::fft::Fft,
cepstral_spec: Vec<Complex>,
cepstral_orig: Vec<f64>,
bark_bands: Vec<usize>,
postfilter: MusicalNoisePostFilter,
}
impl Denoiser {
pub fn new(config: DenoiserConfig) -> Self {
let config = config.sanitized();
let strength = config.strength;
let musical_pf = config.musical_noise_postfilter;
let sample_rate = config.sample_rate;
let frame_size = config.frame_size;
let hop = (frame_size as f64 * (1.0 - config.overlap)).round() as usize;
let hop = hop.max(1);
let stft = Stft::new(StftConfig {
frame_size,
hop,
window: config.window,
window_params: config.window_params,
});
let m = stft.nbins();
let xi_min = 10f64.powf(-25.0 / 10.0); let g_min_db = -20.0 - 25.0 * config.strength;
let g_min = 10f64.powf(g_min_db / 20.0);
let alpha_os = 1.0 + 2.0 * config.strength; let beta_floor = 0.02;
let gain_params = GainParams {
xi_min,
g_min,
alpha_os,
beta_floor,
};
let noise_cfg = NoiseConfig {
adaptive_profile: config.adaptive_noise,
..NoiseConfig::default()
};
let noise = NoiseEstimator::new(noise_cfg, m, sample_rate, hop);
let makeup = 10f64.powf(config.makeup_gain_db / 20.0);
let cepstral_fft_size = (2 * m).next_power_of_two().max(32);
let cepstral_fft = crate::fft::Fft::new(cepstral_fft_size);
Denoiser {
config,
stft,
noise,
noise_cfg,
gain_params,
sample_rate,
frame_size,
hop,
m,
alpha_dd: 0.98,
xi_min,
makeup,
prev_g: vec![0.0; m],
prev_y2: vec![0.0; m],
prev_lambda_d: vec![1e-12; m],
prev_gsmooth: vec![1.0; m],
spec: vec![Complex::default(); frame_size],
frame: vec![0.0; frame_size],
y2: vec![0.0; m],
g: vec![0.0; m],
lambda_d_buf: vec![0.0; m],
spp_buf: vec![0.0; m],
y2_snapshot_buf: vec![0.0; m],
prev_frame_energy: 0.0,
prev_mag: vec![0.0; m],
pre_emph_prev: 0.0,
de_emph_prev: 0.0,
dc_prev_x: 0.0,
dc_prev_y: 0.0,
cepstral_fft,
cepstral_spec: vec![Complex::default(); cepstral_fft_size],
cepstral_orig: vec![0.0; m],
bark_bands: bin_to_bark_band(m, sample_rate),
postfilter: MusicalNoisePostFilter::new(
m,
PostFilterConfig {
enabled: musical_pf,
strength,
..PostFilterConfig::default()
},
),
}
}
pub fn config(&self) -> &DenoiserConfig {
&self.config
}
fn reset_for_channel(&mut self) {
self.noise = NoiseEstimator::new(self.noise_cfg, self.m, self.sample_rate, self.hop);
self.noise.adapt = self.config.adapt;
for v in &mut self.prev_g {
*v = 0.0;
}
for v in &mut self.prev_y2 {
*v = 0.0;
}
for v in &mut self.prev_lambda_d {
*v = 1e-12;
}
for v in &mut self.prev_gsmooth {
*v = 1.0;
}
self.prev_frame_energy = 0.0;
self.prev_mag.fill(0.0);
self.pre_emph_prev = 0.0;
self.de_emph_prev = 0.0;
self.dc_prev_x = 0.0;
self.dc_prev_y = 0.0;
self.postfilter.reset();
}
fn dc_block(input: &[f64]) -> Vec<f64> {
let r = 0.999;
let mut out = Vec::with_capacity(input.len());
let mut prev_x = 0.0;
let mut prev_y = 0.0;
for &x in input {
let y = x - prev_x + r * prev_y;
out.push(y);
prev_x = x;
prev_y = y;
}
out
}
fn dc_block_sample(&mut self, x: f64) -> f64 {
let r = 0.999;
let y = x - self.dc_prev_x + r * self.dc_prev_y;
self.dc_prev_x = x;
self.dc_prev_y = y;
y
}
fn pre_emphasize(&mut self, input: &[f64]) -> Vec<f64> {
let alpha = self.config.pre_emphasis_alpha;
let mut out = Vec::with_capacity(input.len());
let mut prev = self.pre_emph_prev;
for &x in input {
let y = x - alpha * prev;
out.push(y);
prev = x;
}
self.pre_emph_prev = prev;
out
}
fn pre_emphasize_sample(&mut self, x: f64) -> f64 {
let y = x - self.config.pre_emphasis_alpha * self.pre_emph_prev;
self.pre_emph_prev = x;
y
}
fn de_emphasize(&mut self, input: &[f64]) -> Vec<f64> {
let alpha = self.config.pre_emphasis_alpha;
let mut out = Vec::with_capacity(input.len());
let mut prev = self.de_emph_prev;
for &y in input {
let x = y + alpha * prev;
out.push(x);
prev = x;
}
self.de_emph_prev = prev;
out
}
fn de_emphasize_sample(&mut self, y: f64) -> f64 {
let x = y + self.config.pre_emphasis_alpha * self.de_emph_prev;
self.de_emph_prev = x;
x
}
fn compute_transient_score(&mut self) -> f64 {
let m = self.m;
let mut flux = 0.0;
let mut energy = 0.0;
for (k, &y2_k) in self.y2_snapshot_buf.iter().enumerate().take(m) {
let mag = y2_k.sqrt();
energy += y2_k;
let prev_mag = self.prev_mag[k];
flux += (mag - prev_mag).abs();
self.prev_mag[k] = mag * 0.7 + prev_mag * 0.3; }
let delta_e = (energy - self.prev_frame_energy).max(0.0);
self.prev_frame_energy = energy * 0.6 + self.prev_frame_energy * 0.4;
let norm_flux = if energy > 1e-12 {
(flux / (energy.sqrt() + 1e-9)).clamp(0.0, 8.0) / 8.0
} else {
0.0
};
let energy_rise = if energy > 1e-12 {
(delta_e / (energy + 1e-9)).clamp(0.0, 3.0) / 3.0
} else {
0.0
};
(0.75 * norm_flux + 0.25 * energy_rise).clamp(0.0, 1.0)
}
fn cepstral_smooth_gains(&mut self) {
let m = self.m;
if m < 8 {
return;
}
let mut min_g = 1.0f64;
let mut max_g = 0.0f64;
let mut sum = 0.0;
for &v in self.g.iter() {
min_g = min_g.min(v);
max_g = max_g.max(v);
sum += v;
}
let mean = sum / m as f64;
let variation = (max_g - min_g) / mean.max(1e-6);
if variation < 0.04 {
return;
}
self.cepstral_orig.copy_from_slice(&self.g);
let fft_size = self.cepstral_fft.size();
let keep = 6.min(fft_size / 10);
self.cepstral_spec.fill(Complex::default());
for (i, &gi) in self.g.iter().enumerate().take(m) {
self.cepstral_spec[i] = Complex::new(gi.max(1e-8).ln(), 0.0);
}
self.cepstral_fft.forward(&mut self.cepstral_spec);
for slot in self.cepstral_spec.iter_mut().take(fft_size - keep).skip(keep) {
*slot = Complex::default();
}
self.cepstral_fft.inverse(&mut self.cepstral_spec);
let blend = (0.35 + 0.45 * variation.min(1.0)).min(0.75);
for (i, gi) in self.g.iter_mut().enumerate().take(m) {
let liftered = self.cepstral_spec[i].re.exp().clamp(1e-6, 1.0);
*gi = (blend * liftered + (1.0 - blend) * self.cepstral_orig[i]).clamp(1e-6, 1.0);
}
}
fn detect_profile_frames(&mut self, input: &[f64]) -> usize {
let n = self.frame_size;
let m = self.m;
let hop = self.hop;
let frames_15s = (1.5 * self.sample_rate as f64 / hop as f64) as usize;
let max_check = frames_15s.max(8);
let mut spec = vec![crate::fft::Complex::default(); n];
let mut frame = vec![0.0; n];
let mut flatness = Vec::with_capacity(max_check);
let mut start = 0;
while start + n <= input.len() && flatness.len() < max_check {
frame[..n].copy_from_slice(&input[start..start + n]);
self.stft.analyze(&frame, &mut spec);
let mut sum_p = 0.0;
let mut sum_logp = 0.0;
let mut nz = 0usize;
for &c in spec.iter().take(m) {
let p = c.re * c.re + c.im * c.im;
if p > 1e-20 {
sum_p += p;
sum_logp += p.ln();
nz += 1;
}
}
let f = if nz > 0 {
let gm = (sum_logp / nz as f64).exp();
let am = sum_p / nz as f64;
(gm / am.max(1e-300)).clamp(0.0, 1.0)
} else {
0.0
};
flatness.push(f);
start += hop;
}
if flatness.is_empty() {
return 0;
}
let fmax = flatness.iter().cloned().fold(0.0f64, f64::max);
let fmin = flatness.iter().cloned().fold(1.0f64, f64::min);
if fmax - fmin < 0.08 {
return 0;
}
let flat_thr = fmin + 0.6 * (fmax - fmin);
let mut run = 0;
for &f in &flatness {
if f >= flat_thr {
run += 1;
} else {
break;
}
}
let min_frames = ((0.08 * self.sample_rate as f64 / hop as f64).round() as usize).max(1);
if run >= min_frames && run < flatness.len() {
run
} else {
0
}
}
fn collect_profile_y2(&mut self, input: &[f64], n_frames: usize) -> Vec<Vec<f64>> {
let n = self.frame_size;
let m = self.m;
let mut out = Vec::with_capacity(n_frames);
let mut start = 0;
let mut idx = 0;
while idx < n_frames && start + n <= input.len() {
self.frame[..n].copy_from_slice(&input[start..start + n]);
self.stft.analyze(&self.frame, &mut self.spec);
let y2: Vec<f64> = (0..m)
.map(|k| {
let c = self.spec[k];
c.re * c.re + c.im * c.im
})
.collect();
out.push(y2);
start += self.hop;
idx += 1;
}
out
}
fn apply_gain(&mut self) {
let n = self.frame_size;
let m = self.m;
self.spec[0] = self.spec[0].mul_real(self.g[0]);
for k in 1..m - 1 {
let gk = self.g[k];
self.spec[k] = self.spec[k].mul_real(gk);
let mir = n - k;
self.spec[mir] = self.spec[mir].mul_real(gk);
}
self.spec[n / 2] = self.spec[n / 2].mul_real(self.g[m - 1]);
}
fn process_frame(
&mut self,
input: &[f64],
start: usize,
frame_idx: usize,
out: &mut [f64],
norm: &mut [f64],
) {
let n = self.frame_size;
let m = self.m;
for i in 0..n {
self.frame[i] = if start + i < input.len() {
input[start + i]
} else {
0.0
};
}
self.stft.analyze(&self.frame, &mut self.spec);
for k in 0..m {
let c = self.spec[k];
self.y2[k] = c.re * c.re + c.im * c.im;
}
self.noise.update(&self.y2);
let frame_energy: f64 = self.y2.iter().sum();
let noise_energy: f64 = self.noise.noise_psd().iter().sum();
if frame_energy > noise_energy * 50.0 {
for k in 0..m {
self.g[k] = 1.0;
}
self.apply_gain();
self.stft.synthesize(&mut self.spec, out, norm, start);
for k in 0..m {
self.prev_g[k] = 1.0;
self.prev_y2[k] = self.y2[k];
self.prev_lambda_d[k] = self.noise.noise_psd()[k];
self.prev_gsmooth[k] = 1.0;
}
return;
}
self.lambda_d_buf.copy_from_slice(self.noise.noise_psd());
self.spp_buf.copy_from_slice(self.noise.speech_presence());
let g_min = self.gain_params.g_min;
let alpha_dd = self.alpha_dd;
let xi_min = self.xi_min;
let algo = self.config.algorithm;
let gp = self.gain_params;
let smoothing = self.config.smoothing;
let tscore = if self.config.transient_protect {
self.y2_snapshot_buf.copy_from_slice(&self.y2);
self.compute_transient_score()
} else {
0.0
};
let mut gamma_frame = vec![0.0f64; m];
let mut xi_frame = vec![0.0f64; m];
for (k, &y2_k) in self.y2.iter().enumerate().take(m) {
let lam = self.lambda_d_buf[k].max(1e-12);
let gamma = y2_k / lam;
let xi_hat = if frame_idx == 0 {
(gamma - 1.0).max(xi_min)
} else {
let prev_sig = self.prev_g[k] * self.prev_g[k] * self.prev_y2[k]
/ self.prev_lambda_d[k].max(1e-12);
alpha_dd * prev_sig + (1.0 - alpha_dd) * (gamma - 1.0).max(xi_min)
};
gamma_frame[k] = gamma;
xi_frame[k] = xi_hat.max(xi_min);
}
let use_mb_specsub = self.config.multiband
&& matches!(
algo,
Algorithm::SpectralSubtraction
| Algorithm::SpecSubNonlinear
| Algorithm::SpecSubGeometric
);
if use_mb_specsub {
let law = match algo {
Algorithm::SpecSubNonlinear => SpecSubLaw::PowerLaw(0.75),
Algorithm::SpecSubGeometric => SpecSubLaw::Geometric,
_ => SpecSubLaw::Linear,
};
let mb = multiband_specsub_gains(&gamma_frame, &self.bark_bands, N_BARK_BANDS, gp, law);
for (k, &mb_k) in mb.iter().enumerate().take(m) {
self.g[k] = mb_k.max(g_min);
}
} else {
for (k, &spp_k) in self.spp_buf.iter().enumerate().take(m) {
let mut gk = compute_gain(algo, xi_frame[k], gamma_frame[k], spp_k, gp);
if gk < g_min {
gk = g_min;
}
if tscore > 0.03 {
let protect = (tscore * 0.85).min(0.96);
gk = gk * (1.0 - protect) + 1.0 * protect;
gk = gk.clamp(g_min, 1.0);
}
let gs = if gk >= self.prev_gsmooth[k] {
gk
} else {
smoothing * self.prev_gsmooth[k] + (1.0 - smoothing) * gk
};
self.prev_gsmooth[k] = gs;
self.g[k] = gs;
}
}
if self.config.perceptual_weighting {
apply_perceptual_weights(&mut self.g, &self.bark_bands, self.config.strength, g_min);
}
if self.config.musical_noise_postfilter {
self.postfilter.apply(&self.y2, &self.lambda_d_buf, &mut self.g);
}
for (k, &lam_k) in self.lambda_d_buf.iter().enumerate().take(m) {
self.prev_g[k] = self.g[k];
self.prev_y2[k] = self.y2[k];
self.prev_lambda_d[k] = lam_k;
}
if self.config.cepstral_smoothing {
self.cepstral_smooth_gains();
let gmin = g_min;
for gi in &mut self.g {
if *gi < gmin {
*gi = gmin;
}
}
}
self.apply_gain();
self.stft.synthesize(&mut self.spec, out, norm, start);
}
pub fn process_channel(&mut self, input: &[f64]) -> Vec<f64> {
self.reset_for_channel();
let mut x: Vec<f64> = if self.config.dc_block {
Self::dc_block(input)
} else {
input.to_vec()
};
if self.config.pre_emphasis {
x = self.pre_emphasize(&x);
}
let total = x.len();
let profile_frames = if self.config.profile_ms > 0.0 {
((self.config.profile_ms / 1000.0 * self.sample_rate as f64 / self.hop as f64).round()
as usize)
.max(1)
} else if self.config.profile_ms == 0.0 {
self.detect_profile_frames(&x)
} else {
0
};
if profile_frames > 0 {
let prof = self.collect_profile_y2(&x, profile_frames);
if !prof.is_empty() {
self.noise.seed_from_profile(&prof);
}
}
let n = self.frame_size;
let hop = self.hop;
let plen = total + 2 * n;
let mut padded = vec![0.0; plen];
padded[n..n + total].copy_from_slice(&x);
let mut out = vec![0.0; plen];
let mut norm = vec![0.0; plen];
let mut start = 0usize;
let mut frame_idx = 0usize;
while start + n <= plen {
self.process_frame(&padded, start, frame_idx, &mut out, &mut norm);
start += hop;
frame_idx += 1;
}
let makeup = self.makeup;
let mut result = vec![0.0; total];
for i in 0..total {
let nv = norm[n + i];
if nv > 1e-9 {
result[i] = (out[n + i] / nv) * makeup;
} else {
result[i] = 0.0;
}
}
if self.config.pre_emphasis {
result = self.de_emphasize(&result);
}
result
}
pub fn process(&mut self, channels: &[Vec<f64>]) -> Vec<Vec<f64>> {
channels.iter().map(|ch| self.process_channel(ch)).collect()
}
}
pub struct StreamingDenoiser {
channels: Vec<ChannelStream>,
finished: bool,
}
struct ChannelStream {
denoiser: Denoiser,
input: VecDeque<f64>,
profile: Vec<f64>,
profile_target: usize,
profile_ready: bool,
frame: Vec<f64>,
frame_out: Vec<f64>,
frame_norm: Vec<f64>,
ola_out: Vec<f64>,
ola_norm: Vec<f64>,
pending: VecDeque<f64>,
frame_idx: usize,
input_frames: usize,
emitted_padded: usize,
discarded_left: usize,
returned_frames: usize,
finished: bool,
}
impl ChannelStream {
fn new(config: DenoiserConfig) -> Self {
let denoiser = Denoiser::new(config);
let n = denoiser.frame_size;
let profile_target = if denoiser.config.profile_ms < 0.0 {
0
} else if denoiser.config.profile_ms > 0.0 {
((denoiser.config.profile_ms / 1000.0 * denoiser.sample_rate as f64).round() as usize)
.saturating_add(n)
.max(n)
} else {
((1.5 * denoiser.sample_rate as f64).round() as usize)
.saturating_add(n)
.max(n)
};
let mut input = VecDeque::with_capacity(n * 2);
if profile_target == 0 {
input.extend(std::iter::repeat(0.0).take(n));
}
Self {
denoiser,
input,
profile: Vec::with_capacity(profile_target),
profile_ready: profile_target == 0,
profile_target,
frame: vec![0.0; n],
frame_out: vec![0.0; n],
frame_norm: vec![0.0; n],
ola_out: vec![0.0; n],
ola_norm: vec![0.0; n],
pending: VecDeque::with_capacity(n),
frame_idx: 0,
input_frames: 0,
emitted_padded: 0,
discarded_left: 0,
returned_frames: 0,
finished: false,
}
}
#[inline]
fn transform_sample(&mut self, sample: f64) -> f64 {
let mut value = sample;
if self.denoiser.config.dc_block {
value = self.denoiser.dc_block_sample(value);
}
if self.denoiser.config.pre_emphasis {
value = self.denoiser.pre_emphasize_sample(value);
}
value
}
fn initialize_profile(&mut self) {
if self.profile_ready {
return;
}
let profile = std::mem::take(&mut self.profile);
let profile_frames = if self.denoiser.config.profile_ms > 0.0 {
((self.denoiser.config.profile_ms / 1000.0 * self.denoiser.sample_rate as f64
/ self.denoiser.hop as f64)
.round() as usize)
.max(1)
} else if self.denoiser.config.profile_ms == 0.0 {
self.denoiser.detect_profile_frames(&profile)
} else {
0
};
if profile_frames > 0 {
let frames = self
.denoiser
.collect_profile_y2(&profile, profile_frames);
if !frames.is_empty() {
self.denoiser.noise.seed_from_profile(&frames);
}
}
self.profile_ready = true;
let n = self.denoiser.frame_size;
self.input.extend(std::iter::repeat(0.0).take(n));
self.input.extend(profile);
}
fn push_samples(&mut self, samples: &[f64]) {
for &sample in samples {
self.input_frames += 1;
let value = self.transform_sample(sample);
if self.profile_ready {
self.input.push_back(value);
} else {
self.profile.push(value);
if self.profile.len() >= self.profile_target {
self.initialize_profile();
}
}
}
if self.profile_ready {
self.process_available();
}
}
fn process_available(&mut self) {
let n = self.denoiser.frame_size;
let hop = self.denoiser.hop;
while self.input.len() >= n {
for i in 0..n {
self.frame[i] = self.input[i];
}
self.frame_out.fill(0.0);
self.frame_norm.fill(0.0);
self.denoiser.process_frame(
&self.frame,
0,
self.frame_idx,
&mut self.frame_out,
&mut self.frame_norm,
);
for i in 0..n {
self.ola_out[i] += self.frame_out[i];
self.ola_norm[i] += self.frame_norm[i];
}
let makeup = self.denoiser.makeup;
for i in 0..hop {
let norm = self.ola_norm[i];
let value = if norm > 1e-9 {
(self.ola_out[i] / norm) * makeup
} else {
0.0
};
self.pending.push_back(value);
}
self.ola_out.copy_within(hop..n, 0);
self.ola_norm.copy_within(hop..n, 0);
self.ola_out[n - hop..].fill(0.0);
self.ola_norm[n - hop..].fill(0.0);
for _ in 0..hop {
self.input.pop_front();
}
self.frame_idx += 1;
self.emitted_padded += hop;
}
}
fn drain_ready(&mut self) -> Vec<f64> {
let n = self.denoiser.frame_size;
while self.discarded_left < n {
if self.pending.pop_front().is_none() {
break;
}
self.discarded_left += 1;
}
let mut output = Vec::new();
while self.returned_frames < self.input_frames {
let Some(value) = self.pending.pop_front() else {
break;
};
let value = if self.denoiser.config.pre_emphasis {
self.denoiser.de_emphasize_sample(value)
} else {
value
};
output.push(value);
self.returned_frames += 1;
}
output
}
fn finish(&mut self) -> Vec<f64> {
if self.finished {
return Vec::new();
}
if !self.profile_ready {
self.initialize_profile();
}
let n = self.denoiser.frame_size;
self.input.extend(std::iter::repeat(0.0).take(n));
self.process_available();
let target = n.saturating_add(self.input_frames);
if self.emitted_padded < target {
let remaining = (target - self.emitted_padded).min(n);
let makeup = self.denoiser.makeup;
for i in 0..remaining {
let norm = self.ola_norm[i];
let value = if norm > 1e-9 {
(self.ola_out[i] / norm) * makeup
} else {
0.0
};
self.pending.push_back(value);
}
self.emitted_padded += remaining;
}
let output = self.drain_ready();
self.finished = true;
output
}
}
impl StreamingDenoiser {
pub fn new(config: DenoiserConfig, channels: usize) -> Result<Self, String> {
if channels == 0 {
return Err("streaming denoiser requires at least one channel".into());
}
Ok(Self {
channels: (0..channels)
.map(|_| ChannelStream::new(config.clone()))
.collect(),
finished: false,
})
}
pub fn process_block(&mut self, channels: &[Vec<f64>]) -> Result<Vec<Vec<f64>>, String> {
if self.finished {
return Err("streaming denoiser has already been finished".into());
}
if channels.len() != self.channels.len() {
return Err(format!(
"expected {} channels, got {}",
self.channels.len(),
channels.len()
));
}
let frames = channels.first().map(Vec::len).unwrap_or(0);
if channels.iter().any(|channel| channel.len() != frames) {
return Err("streaming blocks must have equal channel lengths".into());
}
for (stream, channel) in self.channels.iter_mut().zip(channels) {
stream.push_samples(channel);
}
Ok(self
.channels
.iter_mut()
.map(ChannelStream::drain_ready)
.collect())
}
pub fn finish(&mut self) -> Result<Vec<Vec<f64>>, String> {
if self.finished {
return Err("streaming denoiser has already been finished".into());
}
let output = self
.channels
.iter_mut()
.map(ChannelStream::finish)
.collect();
self.finished = true;
Ok(output)
}
}
#[cfg(test)]
mod tests {
use super::*;
struct Lcg(u64);
impl Lcg {
fn new(seed: u64) -> Self {
Lcg(seed.wrapping_add(0x9e3779b97f4a7c15))
}
fn uniform(&mut self) -> f64 {
self.0 = self
.0
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
let u = (self.0 >> 32) as f64 / (u32::MAX as f64 + 1.0);
u * 2.0 - 1.0 }
}
fn snr_db(clean: &[f64], test: &[f64]) -> f64 {
let mut sc = 0.0;
let mut sn = 0.0;
for i in 0..clean.len() {
sc += clean[i] * clean[i];
let e = test[i] - clean[i];
sn += e * e;
}
10.0 * (sc / sn.max(1e-300)).log10()
}
#[test]
fn denoising_improves_snr() {
let sr: u32 = 16000;
let dur = 2.0;
let n = (sr as f64 * dur) as usize;
let silence = (sr as f64 * 0.3) as usize;
let mut clean = vec![0.0; n];
for (i, c) in clean.iter_mut().enumerate().take(n).skip(silence) {
let t = i as f64 / sr as f64;
*c = 0.30 * (2.0 * std::f64::consts::PI * 440.0 * t).sin()
+ 0.15 * (2.0 * std::f64::consts::PI * 880.0 * t).sin();
}
let pc: f64 = clean[silence..].iter().map(|s| s * s).sum::<f64>() / (n - silence) as f64;
let pn = pc; let scale = (3.0 * pn).sqrt(); let mut rng = Lcg::new(12345);
let noise: Vec<f64> = (0..n).map(|_| scale * rng.uniform()).collect();
let noisy: Vec<f64> = (0..n).map(|i| clean[i] + noise[i]).collect();
let in_snr = snr_db(&clean[silence..], &noisy[silence..]);
let mut den = Denoiser::new(Preset::Speech.config(sr));
let out = den.process_channel(&noisy);
assert_eq!(out.len(), noisy.len());
let edge = 4096;
let lo = silence + edge;
let hi = n - edge;
let out_snr = snr_db(&clean[lo..hi], &out[lo..hi]);
assert!(
out_snr > in_snr + 3.0,
"expected SNR improvement > 3 dB, got in={in_snr:.2} out={out_snr:.2}"
);
}
#[test]
fn clean_signal_is_preserved() {
let sr: u32 = 16000;
let n = sr as usize * 2;
let silence = sr as usize / 3;
let mut clean = vec![0.0; n];
for (i, c) in clean.iter_mut().enumerate().take(n).skip(silence) {
let t = i as f64 / sr as f64;
*c = 0.25 * (2.0 * std::f64::consts::PI * 660.0 * t).sin();
}
let mut den = Denoiser::new(Preset::Restore.config(sr));
let out = den.process_channel(&clean);
let lo = silence + 4096;
let hi = n - 4096;
let in_rms = (clean[lo..hi].iter().map(|s| s * s).sum::<f64>() / (hi - lo) as f64).sqrt();
let out_rms = (out[lo..hi].iter().map(|s| s * s).sum::<f64>() / (hi - lo) as f64).sqrt();
let rel = (out_rms - in_rms).abs() / in_rms;
assert!(rel < 0.06, "tone amplitude changed by {rel:.3}");
}
#[test]
fn streaming_matches_batch_with_bounded_blocks() {
let sr = 16_000;
let mut config = Preset::Gentle.config(sr);
config.frame_size = 512;
config.overlap = 0.75;
config.profile_ms = -1.0;
config.dc_block = false;
config.pre_emphasis = true;
let signal: Vec<f64> = (0..sr as usize * 2)
.map(|i| {
let t = i as f64 / sr as f64;
0.25 * (2.0 * std::f64::consts::PI * 330.0 * t).sin()
+ 0.04 * (2.0 * std::f64::consts::PI * 2_700.0 * t).sin()
})
.collect();
let mut batch = Denoiser::new(config.clone());
let expected = batch.process_channel(&signal);
let mut streaming = StreamingDenoiser::new(config, 1).unwrap();
let mut actual = Vec::new();
let mut offset = 0;
for block_size in [37, 1_003, 257, 4_096, 89, 777] {
if offset >= signal.len() {
break;
}
let end = (offset + block_size).min(signal.len());
let block = vec![signal[offset..end].to_vec()];
actual.extend(streaming.process_block(&block).unwrap().remove(0));
offset = end;
}
if offset < signal.len() {
actual.extend(
streaming
.process_block(&[signal[offset..].to_vec()])
.unwrap()
.remove(0),
);
}
actual.extend(streaming.finish().unwrap().remove(0));
assert_eq!(actual.len(), expected.len());
let max_error = actual
.iter()
.zip(&expected)
.map(|(a, b)| (a - b).abs())
.fold(0.0, f64::max);
assert!(max_error < 1e-9, "streaming drifted from batch by {max_error}");
}
#[test]
fn hifi_preset_preserves_clean_and_enables_features() {
let sr: u32 = 48000;
let n = (sr as usize) * 2;
let silence = sr as usize / 4;
let mut clean = vec![0.0; n];
for (i, c) in clean.iter_mut().enumerate().take(n).skip(silence) {
let t = i as f64 / sr as f64;
*c = 0.18 * (2.0 * std::f64::consts::PI * 880.0 * t).sin()
+ 0.09 * (2.0 * std::f64::consts::PI * 1760.0 * t).sin();
}
let mut cfg = Preset::HiFi.config(sr);
cfg.cepstral_smoothing = true;
cfg.transient_protect = true;
cfg.pre_emphasis = false;
cfg.strength = 0.28;
let mut den = Denoiser::new(cfg);
let out = den.process_channel(&clean);
let edge = 4096;
let lo = silence + edge;
let hi = n - edge;
let in_rms = (clean[lo..hi].iter().map(|s| s * s).sum::<f64>() / (hi - lo) as f64).sqrt();
let out_rms = (out[lo..hi].iter().map(|s| s * s).sum::<f64>() / (hi - lo) as f64).sqrt();
let rel = (out_rms - in_rms).abs() / in_rms;
assert!(
rel < 0.12,
"hifi changed clean amplitude by {rel:.3} (too much for fidelity mode)"
);
let c = Preset::HiFi.config(sr);
assert!(c.transient_protect);
assert!(c.cepstral_smoothing);
assert!(c.perceptual_weighting);
assert!(c.musical_noise_postfilter);
assert_eq!(c.window, WindowType::Kaiser);
}
#[test]
fn content_modes_coordinate_processing_controls() {
let mut speech = DenoiserConfig::default(48_000);
ProcessingMode::Speech.apply(&mut speech);
assert!(speech.vad && speech.adaptive_noise);
assert!(speech.strength >= 0.7);
let mut music = DenoiserConfig::default(48_000);
ProcessingMode::Music.apply(&mut music);
assert!(!music.vad && music.transient_protect);
assert!(music.strength <= 0.35);
assert!(music.perceptual_weighting);
let mut ambient = DenoiserConfig::default(48_000);
ProcessingMode::Ambient.apply(&mut ambient);
assert!(ambient.adaptive_noise && !ambient.vad);
assert!(ambient.strength <= 0.4);
}
}