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 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,
}
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,
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);
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>,
prev_frame_energy: f64,
prev_mag: Vec<f64>, pre_emph_prev: f64, de_emph_prev: 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::default();
let noise = NoiseEstimator::new(noise_cfg, m, sample_rate, hop);
let makeup = 10f64.powf(config.makeup_gain_db / 20.0);
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],
prev_frame_energy: 0.0,
prev_mag: vec![0.0; m],
pre_emph_prev: 0.0,
de_emph_prev: 0.0,
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.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 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 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 compute_transient_score(&mut self, y2: &[f64]) -> f64 {
let m = self.m;
let mut flux = 0.0;
let mut energy = 0.0;
for k in 0..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
};
let score = (0.75 * norm_flux + 0.25 * energy_rise).clamp(0.0, 1.0);
score
}
fn cepstral_smooth_gains(g: &mut [f64]) {
let m = g.len();
if m < 8 {
return;
}
let mut min_g = 1.0f64;
let mut max_g = 0.0f64;
let mut sum = 0.0;
for &v in 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;
}
let original: Vec<f64> = g.to_vec();
let fft_size = (2 * m).next_power_of_two().max(32);
let keep = 6.min(fft_size / 10);
let mut spec = vec![Complex::default(); fft_size];
let fft = crate::fft::Fft::new(fft_size);
for i in 0..m {
spec[i] = Complex::new(g[i].max(1e-8).ln(), 0.0);
}
fft.forward(&mut spec);
for k in keep..fft_size - keep {
spec[k] = Complex::default();
}
fft.inverse(&mut spec);
let blend = (0.35 + 0.45 * variation.min(1.0)).min(0.75);
for i in 0..m {
let liftered = spec[i].re.exp().clamp(1e-6, 1.0);
g[i] = (blend * liftered + (1.0 - blend) * original[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 k in 0..m {
let p = spec[k].re * spec[k].re + spec[k].im * spec[k].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;
}
let lambda_d: Vec<f64> = self.noise.noise_psd().to_vec();
let spp: Vec<f64> = self.noise.speech_presence().to_vec();
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 {
let y2_snapshot: Vec<f64> = self.y2.clone();
self.compute_transient_score(&y2_snapshot)
} else {
0.0
};
let mut gamma_frame = vec![0.0f64; m];
let mut xi_frame = vec![0.0f64; m];
for k in 0..m {
let gamma = self.y2[k] / lambda_d[k].max(1e-12);
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 in 0..m {
self.g[k] = mb[k].max(g_min);
}
} else {
for k in 0..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, &lambda_d, &mut self.g);
}
for k in 0..m {
self.prev_g[k] = self.g[k];
self.prev_y2[k] = self.y2[k];
self.prev_lambda_d[k] = lambda_d[k];
}
if self.config.cepstral_smoothing {
Self::cepstral_smooth_gains(&mut self.g);
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()
}
}
#[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 in silence..n {
let t = i as f64 / sr as f64;
clean[i] = 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 in silence..n {
let t = i as f64 / sr as f64;
clean[i] = 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 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 in silence..n {
let t = i as f64 / sr as f64;
clean[i] = 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);
}
}