1use crate::audio::SAMPLE_RATE;
19use crate::delay::{ReferenceRing, estimate_delay_samples};
20use crate::fdaf::{FdafConfig, FdafNlms};
21use crate::residual::{ResidualWeights, embedded_residual_weights};
22use anyhow::{Result, ensure};
23use std::path::Path;
24
25#[derive(Debug, Clone)]
26pub struct AecConfig {
27 pub n_fft: usize,
28 pub frame_samples: usize,
29 pub step_size: f32,
30 pub max_delay_ms: u32,
31 pub residual: bool,
32 pub delay_reestimate_frames: usize,
33}
34
35impl Default for AecConfig {
36 fn default() -> Self {
37 Self {
38 n_fft: 1024,
39 frame_samples: 160,
40 step_size: 0.05,
41 max_delay_ms: 300,
42 residual: true,
43 delay_reestimate_frames: 50,
44 }
45 }
46}
47
48pub struct AecSession {
49 cfg: AecConfig,
50 delay_samples: usize,
51 max_delay_samples: usize,
52 reference_ring: ReferenceRing,
53 fdaf: FdafNlms,
54 frame_count: usize,
55 pending_far: Vec<f32>,
56}
57
58impl AecSession {
59 pub fn new(cfg: AecConfig) -> Result<Self> {
60 let max_delay_samples = (cfg.max_delay_ms as usize * SAMPLE_RATE) / 1000;
61 let residual = if cfg.residual {
62 Some(embedded_residual_weights()?)
63 } else {
64 None
65 };
66 let fdaf_cfg = FdafConfig {
67 n_fft: cfg.n_fft,
68 frame_samples: cfg.frame_samples,
69 step_size: cfg.step_size,
70 adapt: true,
71 use_residual: cfg.residual,
72 };
73 let fdaf = FdafNlms::new(fdaf_cfg, residual)?;
74 Ok(Self {
75 cfg,
76 delay_samples: 0,
77 max_delay_samples,
78 reference_ring: ReferenceRing::new(max_delay_samples),
79 fdaf,
80 frame_count: 0,
81 pending_far: Vec::new(),
82 })
83 }
84
85 pub fn with_residual_weights(cfg: AecConfig, weights: ResidualWeights) -> Result<Self> {
86 let max_delay_samples = (cfg.max_delay_ms as usize * SAMPLE_RATE) / 1000;
87 let fdaf_cfg = FdafConfig {
88 n_fft: cfg.n_fft,
89 frame_samples: cfg.frame_samples,
90 step_size: cfg.step_size,
91 adapt: true,
92 use_residual: cfg.residual,
93 };
94 let residual = if cfg.residual { Some(weights) } else { None };
95 let fdaf = FdafNlms::new(fdaf_cfg, residual)?;
96 Ok(Self {
97 cfg,
98 delay_samples: 0,
99 max_delay_samples,
100 reference_ring: ReferenceRing::new(max_delay_samples),
101 fdaf,
102 frame_count: 0,
103 pending_far: Vec::new(),
104 })
105 }
106
107 pub fn reset(&mut self) {
108 self.delay_samples = 0;
109 self.frame_count = 0;
110 self.reference_ring.clear();
111 self.pending_far.clear();
112 self.fdaf.reset();
113 }
114
115 pub fn delay_samples(&self) -> usize {
116 self.delay_samples
117 }
118
119 pub fn push_reference(&mut self, far_end: &[f32]) {
121 self.reference_ring.push(far_end);
122 self.pending_far.extend_from_slice(far_end);
123 }
124
125 pub fn process(&mut self, mic: &[f32], far_end: &[f32]) -> Result<Vec<f32>> {
127 ensure!(mic.len() == far_end.len());
128 self.push_reference(far_end);
129 self.process_mic(mic)?
130 .ok_or_else(|| anyhow::anyhow!("empty mic buffer"))
131 }
132
133 pub fn process_mic(&mut self, mic: &[f32]) -> Result<Option<Vec<f32>>> {
135 if mic.is_empty() {
136 return Ok(None);
137 }
138 let hop = self.cfg.frame_samples;
139 let mut out = vec![0.0f32; mic.len()];
140 let mut pos = 0;
141 while pos < mic.len() {
142 let end = (pos + hop).min(mic.len());
143 let chunk = end - pos;
144 let mut aligned = vec![0.0f32; hop];
145 self.reference_ring
146 .read_delayed(self.delay_samples, hop, &mut aligned);
147 if chunk < hop {
148 aligned.fill(0.0);
149 self.reference_ring
150 .read_delayed(self.delay_samples, hop, &mut aligned);
151 }
152 let mut frame_out = vec![0.0f32; hop];
153 if chunk == hop {
154 self.fdaf
155 .process_frame(&mic[pos..end], &aligned, &mut frame_out)?;
156 out[pos..end].copy_from_slice(&frame_out);
157 } else {
158 let mut mp = vec![0.0f32; hop];
159 mp[..chunk].copy_from_slice(&mic[pos..end]);
160 self.fdaf.process_frame(&mp, &aligned, &mut frame_out)?;
161 out[pos..end].copy_from_slice(&frame_out[..chunk]);
162 }
163 pos = end;
164 self.frame_count += 1;
165 if self.frame_count % self.cfg.delay_reestimate_frames == 1 {
166 self.maybe_reestimate_delay(mic);
167 }
168 }
169 Ok(Some(out))
170 }
171
172 fn maybe_reestimate_delay(&mut self, mic: &[f32]) {
173 let n = self
174 .pending_far
175 .len()
176 .min(mic.len())
177 .min(self.cfg.n_fft * 4);
178 if n < 256 {
179 return;
180 }
181 let est = estimate_delay_samples(
182 &self.pending_far[..n],
183 &mic[..n],
184 self.cfg.n_fft,
185 self.max_delay_samples,
186 );
187 self.delay_samples = est;
188 }
189
190 pub fn process_aligned_buffers(&mut self, mic: &[f32], far: &[f32]) -> Result<Vec<f32>> {
192 ensure!(mic.len() == far.len());
193 let n = mic.len().min(far.len()).min(self.cfg.n_fft * 8);
194 if n >= 256 {
195 self.delay_samples = estimate_delay_samples(
196 &far[..n],
197 &mic[..n],
198 self.cfg.n_fft,
199 self.max_delay_samples,
200 );
201 }
202 let mut aligned_far = vec![0.0f32; far.len()];
203 if self.delay_samples < far.len() {
204 aligned_far[self.delay_samples..]
205 .copy_from_slice(&far[..far.len() - self.delay_samples]);
206 }
207 let mut out = vec![0.0f32; mic.len()];
208 self.fdaf.process_buffer(mic, &aligned_far, &mut out)?;
209 Ok(out)
210 }
211
212 pub fn process_wav_files(
214 mic_path: &Path,
215 ref_path: &Path,
216 cfg: &AecConfig,
217 ) -> Result<(Vec<f32>, usize)> {
218 let mic = crate::audio::load_wav_16k(mic_path)?;
219 let far = crate::audio::load_wav_16k(ref_path)?;
220 let len = mic.len().min(far.len());
221 let mut session = AecSession::new(cfg.clone())?;
222 let out = session.process_aligned_buffers(&mic[..len], &far[..len])?;
223 Ok((out, session.delay_samples()))
224 }
225}