Skip to main content

rlx_aec/
session.rs

1// RLX — versatile ML compiler + runtime.
2// Copyright (C) 2026 Eugene Hauptmann, Nataliya Kosmyna.
3//
4// This program is free software: you can redistribute it and/or modify
5// it under the terms of the GNU General Public License as published by
6// the Free Software Foundation, version 3.
7//
8// This program is distributed in the hope that it will be useful,
9// but WITHOUT ANY WARRANTY; without even the implied warranty of
10// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
11// GNU General Public License for more details.
12//
13// You should have received a copy of the GNU General Public License
14// along with this program. If not, see <https://www.gnu.org/licenses/>.
15
16//! Top-level AEC session: delay alignment + FDAF + optional residual.
17
18use 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    /// Push far-end reference (e.g. TTS playback tap) without mic input.
120    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    /// Process mic + synchronous far-end chunk.
126    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    /// Process mic using buffered reference ring.
134    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    /// Offline: estimate delay once then process full buffers.
191    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    /// Process paired WAV files at 16 kHz.
213    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}