Skip to main content

embedded_dsp/
filtering.rs

1//! Digital filtering functions (FIR, Biquad IIR Direct Form I & II, LMS Adaptive Filter, Convolution, Correlation).
2
3use crate::types::*;
4
5// --- FIR Filter ---
6
7/// Instance structure for the floating-point FIR filter.
8pub struct FirInstanceF32<'a> {
9    pub num_taps: u16,
10    pub coeffs: &'a [f32],
11    pub state: &'a mut [f32],
12}
13
14impl<'a> FirInstanceF32<'a> {
15    pub fn init(num_taps: u16, coeffs: &'a [f32], state: &'a mut [f32]) -> Self {
16        state.fill(0.0);
17        Self {
18            num_taps,
19            coeffs,
20            state,
21        }
22    }
23}
24
25pub fn fir_f32(instance: &mut FirInstanceF32, src: &[f32], dst: &mut [f32]) {
26    let num_taps = instance.num_taps as usize;
27    let block_size = src.len().min(dst.len());
28
29    for i in 0..block_size {
30        // Shift state
31        for k in (1..num_taps).rev() {
32            instance.state[k] = instance.state[k - 1];
33        }
34        instance.state[0] = src[i];
35
36        // Compute dot product with coefficients
37        let mut acc = 0.0f32;
38        for k in 0..num_taps {
39            acc += instance.state[k] * instance.coeffs[k];
40        }
41        dst[i] = acc;
42    }
43}
44
45/// Instance structure for the Q31 FIR filter.
46pub struct FirInstanceQ31<'a> {
47    pub num_taps: u16,
48    pub coeffs: &'a [q31],
49    pub state: &'a mut [q31],
50}
51
52impl<'a> FirInstanceQ31<'a> {
53    pub fn init(num_taps: u16, coeffs: &'a [q31], state: &'a mut [q31]) -> Self {
54        state.fill(0);
55        Self {
56            num_taps,
57            coeffs,
58            state,
59        }
60    }
61}
62
63pub fn fir_q31(instance: &mut FirInstanceQ31, src: &[q31], dst: &mut [q31]) {
64    let num_taps = instance.num_taps as usize;
65    let block_size = src.len().min(dst.len());
66
67    for i in 0..block_size {
68        for k in (1..num_taps).rev() {
69            instance.state[k] = instance.state[k - 1];
70        }
71        instance.state[0] = src[i];
72
73        let mut acc: i64 = 0;
74        for k in 0..num_taps {
75            acc += (instance.state[k] as i64 * instance.coeffs[k] as i64) >> 31;
76        }
77        dst[i] = acc.clamp(i32::MIN as i64, i32::MAX as i64) as q31;
78    }
79}
80
81/// Instance structure for the Q15 FIR filter.
82pub struct FirInstanceQ15<'a> {
83    pub num_taps: u16,
84    pub coeffs: &'a [q15],
85    pub state: &'a mut [q15],
86}
87
88impl<'a> FirInstanceQ15<'a> {
89    pub fn init(num_taps: u16, coeffs: &'a [q15], state: &'a mut [q15]) -> Self {
90        state.fill(0);
91        Self {
92            num_taps,
93            coeffs,
94            state,
95        }
96    }
97}
98
99pub fn fir_q15(instance: &mut FirInstanceQ15, src: &[q15], dst: &mut [q15]) {
100    let num_taps = instance.num_taps as usize;
101    let block_size = src.len().min(dst.len());
102
103    for i in 0..block_size {
104        for k in (1..num_taps).rev() {
105            instance.state[k] = instance.state[k - 1];
106        }
107        instance.state[0] = src[i];
108
109        let mut acc: i32 = 0;
110        for k in 0..num_taps {
111            acc += (instance.state[k] as i32 * instance.coeffs[k] as i32) >> 15;
112        }
113        dst[i] = acc.clamp(i16::MIN as i32, i16::MAX as i32) as q15;
114    }
115}
116
117// --- Biquad Cascade Direct Form I Filter ---
118
119/// Instance structure for the floating-point Biquad Cascade Direct Form I filter.
120pub struct BiquadCascadeInstanceF32<'a> {
121    pub num_stages: u8,
122    pub coeffs: &'a [f32],    // 5 * num_stages: [b0, b1, b2, a1, a2]
123    pub state: &'a mut [f32], // 4 * num_stages: [x[n-1], x[n-2], y[n-1], y[n-2]]
124}
125
126impl<'a> BiquadCascadeInstanceF32<'a> {
127    pub fn init(num_stages: u8, coeffs: &'a [f32], state: &'a mut [f32]) -> Self {
128        state.fill(0.0);
129        Self {
130            num_stages,
131            coeffs,
132            state,
133        }
134    }
135}
136
137pub fn biquad_cascade_df1_f32(
138    instance: &mut BiquadCascadeInstanceF32,
139    src: &[f32],
140    dst: &mut [f32],
141) {
142    let num_stages = instance.num_stages as usize;
143    let block_size = src.len().min(dst.len());
144
145    let mut in_val;
146    let mut out_val;
147
148    for i in 0..block_size {
149        in_val = src[i];
150        for stage in 0..num_stages {
151            let b0 = instance.coeffs[stage * 5];
152            let b1 = instance.coeffs[stage * 5 + 1];
153            let b2 = instance.coeffs[stage * 5 + 2];
154            let a1 = instance.coeffs[stage * 5 + 3];
155            let a2 = instance.coeffs[stage * 5 + 4];
156
157            let x1 = instance.state[stage * 4];
158            let x2 = instance.state[stage * 4 + 1];
159            let y1 = instance.state[stage * 4 + 2];
160            let y2 = instance.state[stage * 4 + 3];
161
162            out_val = b0 * in_val + b1 * x1 + b2 * x2 + a1 * y1 + a2 * y2;
163
164            instance.state[stage * 4 + 1] = x1;
165            instance.state[stage * 4] = in_val;
166            instance.state[stage * 4 + 3] = y1;
167            instance.state[stage * 4 + 2] = out_val;
168
169            in_val = out_val;
170        }
171        dst[i] = in_val;
172    }
173}
174
175// --- LMS Adaptive Filter ---
176
177/// Instance structure for the floating-point LMS adaptive filter.
178pub struct LmsInstanceF32<'a> {
179    pub num_taps: u16,
180    pub coeffs: &'a mut [f32],
181    pub state: &'a mut [f32],
182    pub mu: f32,
183}
184
185impl<'a> LmsInstanceF32<'a> {
186    pub fn init(num_taps: u16, coeffs: &'a mut [f32], state: &'a mut [f32], mu: f32) -> Self {
187        state.fill(0.0);
188        coeffs.fill(0.0);
189        Self {
190            num_taps,
191            coeffs,
192            state,
193            mu,
194        }
195    }
196}
197
198pub fn lms_f32(
199    instance: &mut LmsInstanceF32,
200    src: &[f32],
201    ref_signal: &[f32],
202    out: &mut [f32],
203    err: &mut [f32],
204) {
205    let num_taps = instance.num_taps as usize;
206    let block_size = src
207        .len()
208        .min(ref_signal.len())
209        .min(out.len())
210        .min(err.len());
211
212    for i in 0..block_size {
213        for k in (1..num_taps).rev() {
214            instance.state[k] = instance.state[k - 1];
215        }
216        instance.state[0] = src[i];
217
218        let mut acc = 0.0f32;
219        for k in 0..num_taps {
220            acc += instance.state[k] * instance.coeffs[k];
221        }
222        out[i] = acc;
223        let e = ref_signal[i] - acc;
224        err[i] = e;
225
226        // Update coefficients: w[n+1] = w[n] + 2 * mu * e[n] * x[n]
227        let alpha = 2.0 * instance.mu * e;
228        for k in 0..num_taps {
229            instance.coeffs[k] += alpha * instance.state[k];
230        }
231    }
232}
233
234// --- Convolution ---
235
236pub fn conv_f32(src_a: &[f32], src_b: &[f32], dst: &mut [f32]) {
237    let len_a = src_a.len();
238    let len_b = src_b.len();
239    let out_len = (len_a + len_b - 1).min(dst.len());
240
241    dst[..out_len].fill(0.0);
242    for i in 0..len_a {
243        for j in 0..len_b {
244            if i + j < out_len {
245                dst[i + j] += src_a[i] * src_b[j];
246            }
247        }
248    }
249}
250
251pub fn conv_q31(src_a: &[q31], src_b: &[q31], dst: &mut [q31]) {
252    let len_a = src_a.len();
253    let len_b = src_b.len();
254    let out_len = (len_a + len_b - 1).min(dst.len());
255
256    for n in 0..out_len {
257        let mut acc: i64 = 0;
258        let k_min = if n >= len_b - 1 { n - (len_b - 1) } else { 0 };
259        let k_max = n.min(len_a - 1);
260        for k in k_min..=k_max {
261            acc += (src_a[k] as i64 * src_b[n - k] as i64) >> 31;
262        }
263        dst[n] = acc.clamp(i32::MIN as i64, i32::MAX as i64) as q31;
264    }
265}
266
267pub fn conv_q15(src_a: &[q15], src_b: &[q15], dst: &mut [q15]) {
268    let len_a = src_a.len();
269    let len_b = src_b.len();
270    let out_len = (len_a + len_b - 1).min(dst.len());
271
272    for n in 0..out_len {
273        let mut acc: i32 = 0;
274        let k_min = if n >= len_b - 1 { n - (len_b - 1) } else { 0 };
275        let k_max = n.min(len_a - 1);
276        for k in k_min..=k_max {
277            acc += (src_a[k] as i32 * src_b[n - k] as i32) >> 15;
278        }
279        dst[n] = acc.clamp(i16::MIN as i32, i16::MAX as i32) as q15;
280    }
281}
282
283pub fn conv_q7(src_a: &[q7], src_b: &[q7], dst: &mut [q7]) {
284    let len_a = src_a.len();
285    let len_b = src_b.len();
286    let out_len = (len_a + len_b - 1).min(dst.len());
287
288    for n in 0..out_len {
289        let mut acc: i32 = 0;
290        let k_min = if n >= len_b - 1 { n - (len_b - 1) } else { 0 };
291        let k_max = n.min(len_a - 1);
292        for k in k_min..=k_max {
293            acc += (src_a[k] as i32 * src_b[n - k] as i32) >> 7;
294        }
295        dst[n] = acc.clamp(i8::MIN as i32, i8::MAX as i32) as q7;
296    }
297}
298
299// --- Correlation ---
300
301pub fn correlate_f32(src_a: &[f32], src_b: &[f32], dst: &mut [f32]) {
302    let len_a = src_a.len();
303    let len_b = src_b.len();
304    let out_len = (len_a + len_b - 1).min(dst.len());
305
306    dst[..out_len].fill(0.0);
307    for n in 0..out_len {
308        let mut acc = 0.0f32;
309        for k in 0..len_a {
310            let idx_b = (k as isize) + (len_b as isize - 1) - (n as isize);
311            if idx_b >= 0 && (idx_b as usize) < len_b {
312                acc += src_a[k] * src_b[idx_b as usize];
313            }
314        }
315        dst[n] = acc;
316    }
317}
318
319pub fn correlate_q31(src_a: &[q31], src_b: &[q31], dst: &mut [q31]) {
320    let len_a = src_a.len();
321    let len_b = src_b.len();
322    let out_len = (len_a + len_b - 1).min(dst.len());
323
324    for n in 0..out_len {
325        let mut acc: i64 = 0;
326        for k in 0..len_a {
327            let idx_b = (k as isize) + (len_b as isize - 1) - (n as isize);
328            if idx_b >= 0 && (idx_b as usize) < len_b {
329                acc += (src_a[k] as i64 * src_b[idx_b as usize] as i64) >> 31;
330            }
331        }
332        dst[n] = acc.clamp(i32::MIN as i64, i32::MAX as i64) as q31;
333    }
334}
335
336pub fn correlate_q15(src_a: &[q15], src_b: &[q15], dst: &mut [q15]) {
337    let len_a = src_a.len();
338    let len_b = src_b.len();
339    let out_len = (len_a + len_b - 1).min(dst.len());
340
341    for n in 0..out_len {
342        let mut acc: i32 = 0;
343        for k in 0..len_a {
344            let idx_b = (k as isize) + (len_b as isize - 1) - (n as isize);
345            if idx_b >= 0 && (idx_b as usize) < len_b {
346                acc += (src_a[k] as i32 * src_b[idx_b as usize] as i32) >> 15;
347            }
348        }
349        dst[n] = acc.clamp(i16::MIN as i32, i16::MAX as i32) as q15;
350    }
351}