#[derive(Clone, Copy, Debug)]
pub struct NoiseConfig {
pub alpha_s: f64,
pub alpha_d: f64,
pub window_frames: usize,
pub b_min: f64,
pub zeta0: f64,
pub sigma: f64,
pub anchor_cap_db: f64,
pub up_rate_dbps: f64,
}
impl Default for NoiseConfig {
fn default() -> Self {
NoiseConfig {
alpha_s: 0.9,
alpha_d: 0.95,
window_frames: 0, b_min: 2.0,
zeta0: 2.0,
sigma: 0.25,
anchor_cap_db: 15.0,
up_rate_dbps: 6.0,
}
}
}
pub struct NoiseEstimator {
cfg: NoiseConfig,
nbins: usize,
s: Vec<f64>, s_min: Vec<f64>, s_tmp: Vec<f64>, lambda_d: Vec<f64>, p: Vec<f64>, initialized: bool,
pub adapt: bool,
profile_psd: Vec<f64>,
has_profile: bool,
up_ratio: f64,
anchor_ratio: f64,
min_forget_fast: f64,
min_forget_slow: f64,
}
impl NoiseEstimator {
pub fn new(cfg: NoiseConfig, nbins: usize, sample_rate: u32, hop: usize) -> Self {
let window_frames = if cfg.window_frames == 0 {
((1.5 * sample_rate as f64 / hop as f64).round() as usize).max(8)
} else {
cfg.window_frames
};
let anchor_ratio = 10f64.powf(cfg.anchor_cap_db / 10.0);
let up_ratio = 10f64.powf(cfg.up_rate_dbps * hop as f64 / sample_rate as f64 / 10.0);
let dt = hop as f64 / sample_rate as f64; let min_forget_fast = 10f64.powf(9.0 * dt / 10.0);
let min_forget_slow = 10f64.powf(0.5 * dt / 10.0);
NoiseEstimator {
cfg: NoiseConfig {
window_frames,
..cfg
},
nbins,
s: vec![0.0; nbins],
s_min: vec![f64::MAX; nbins],
s_tmp: vec![f64::MAX; nbins],
lambda_d: vec![1e-10; nbins],
p: vec![0.0; nbins],
initialized: false,
adapt: true,
profile_psd: vec![0.0; nbins],
has_profile: false,
up_ratio,
anchor_ratio,
min_forget_fast,
min_forget_slow,
}
}
#[inline]
pub fn nbins(&self) -> usize {
self.nbins
}
pub fn seed_from_profile(&mut self, profile_frames: &[Vec<f64>]) {
if profile_frames.is_empty() {
return;
}
for k in 0..self.nbins {
let mut acc = 0.0;
let mut n = 0.0;
for fr in profile_frames {
if k < fr.len() {
acc += fr[k];
n += 1.0;
}
}
let val = (if n > 0.0 { acc / n } else { 1e-10 }).max(1e-12);
self.lambda_d[k] = val;
self.s[k] = val;
self.s_min[k] = val;
self.s_tmp[k] = val;
self.profile_psd[k] = val;
}
self.initialized = true;
self.has_profile = true;
}
#[inline]
pub fn noise_psd(&self) -> &[f64] {
&self.lambda_d
}
#[inline]
pub fn speech_presence(&self) -> &[f64] {
&self.p
}
pub fn update(&mut self, y2: &[f64]) {
debug_assert_eq!(y2.len(), self.nbins);
let cfg = self.cfg;
let frame_energy: f64 = y2.iter().sum();
if !self.initialized {
if frame_energy < 1e-9 {
self.p.fill(0.0);
return;
}
for k in 0..self.nbins {
let v = y2[k].max(1e-12);
self.s[k] = v;
self.s_min[k] = v;
self.s_tmp[k] = v;
self.lambda_d[k] = v;
}
self.initialized = true;
self.p.fill(0.0);
return;
}
for k in 0..self.nbins {
self.s[k] = cfg.alpha_s * self.s[k] + (1.0 - cfg.alpha_s) * y2[k];
}
let gamma_fast = if self.has_profile {
self.min_forget_slow } else {
self.min_forget_fast };
let gamma_slow = self.min_forget_slow;
for k in 0..self.nbins {
let f = self.s_min[k] * gamma_fast;
self.s_min[k] = if self.s[k] < f { self.s[k] } else { f };
let g = self.s_tmp[k] * gamma_slow;
self.s_tmp[k] = if self.s[k] < g { self.s[k] } else { g };
}
if self.has_profile {
let cap = self.anchor_ratio.sqrt();
for k in 0..self.nbins {
let m = self.profile_psd[k] * cap;
if self.s_min[k] > m {
self.s_min[k] = m;
}
if self.s_tmp[k] > m {
self.s_tmp[k] = m;
}
}
}
for k in 0..self.nbins {
let denom = (cfg.b_min * self.s_min[k]).max(1e-12);
let zeta = self.s[k] / denom;
let arg = (zeta - cfg.zeta0) / cfg.sigma;
let p = if arg >= 0.0 {
1.0 / (1.0 + (-arg).exp())
} else {
let e = arg.exp();
e / (1.0 + e)
};
self.p[k] = p;
if self.adapt {
let old = self.lambda_d[k];
let a_d_eff = cfg.alpha_d + (1.0 - cfg.alpha_d) * p;
let mut new_ld = a_d_eff * old + (1.0 - a_d_eff) * y2[k];
let cap_up = old * self.up_ratio;
if new_ld > cap_up {
new_ld = cap_up;
}
if self.has_profile {
let anchor = self.profile_psd[k] * self.anchor_ratio;
if new_ld > anchor {
new_ld = anchor;
}
}
self.lambda_d[k] = new_ld;
}
}
}
}