use anyhow::{Result, bail};
use super::FIR;
use crate::{
Flt, StrictlyPositive, Vd,
filter::{Filter, FilterMethods},
twopi, *,
};
enum State {
FIR0Active,
FIR1Active,
Switching(bool, usize),
}
#[cfg_attr(feature = "python-bindings", gen_stub_pyclass, pyclass)]
pub struct AdaptableFIR {
firs: [FIR; 2],
state: State,
transition_samples: usize,
block_size: usize,
}
impl AdaptableFIR {
pub fn new(
fs: StrictlyPositive,
block_size: usize,
transition_time: StrictlyPositive,
init_coeffs: Option<&[Flt]>,
) -> Self {
let coefs = init_coeffs.unwrap_or(&[0.]);
let transition_samples = (*fs * *transition_time) as usize;
assert!(transition_samples > 0);
Self {
block_size,
transition_samples,
firs: [
FIR::new(coefs, block_size).unwrap(),
FIR::new(coefs, block_size).unwrap(),
],
state: State::FIR0Active,
}
}
pub fn updateCoefficients(&mut self, coefs: &[Flt]) -> Result<()> {
match self.state {
State::FIR0Active => {
self.firs[1] = FIR::new(coefs, self.block_size).unwrap();
self.state = State::Switching(true, 0);
Ok(())
}
State::FIR1Active => {
self.firs[0] = FIR::new(coefs, self.block_size).unwrap();
self.state = State::Switching(false, 0);
Ok(())
}
State::Switching(..) => {
bail!("Cannot update coefficients while an update is in progress");
}
}
}
}
impl FilterMethods for AdaptableFIR {
fn filter(&mut self, input: &[Flt], output: &mut [Flt]) {
let Self {
firs,
state,
transition_samples,
..
} = self;
match state {
State::FIR0Active => firs[0].filter(input, output),
State::FIR1Active => firs[1].filter(input, output),
State::Switching(from_0to1, n) => {
firs[0].filter(input, output);
let mut o1 = vec![0.0; output.len()];
firs[1].filter(input, &mut o1);
for (o0i, o1i) in output.iter_mut().zip(o1.iter()) {
let (gain0, gain1) = calc_gains(*from_0to1, *n, *transition_samples);
*o0i = gain0 * *o0i + gain1 * o1i;
*n += 1;
}
if n >= transition_samples {
if *from_0to1 {
*state = State::FIR1Active;
} else {
*state = State::FIR0Active;
}
}
}
}
}
fn reset(&mut self) {
match self.state {
State::FIR0Active => self.firs[0].reset(),
State::FIR1Active => self.firs[1].reset(),
State::Switching(_, _) => {
self.firs[0].reset();
self.firs[1].reset();
}
}
}
}
#[cfg(feature = "python-bindings")]
#[cfg_attr(feature = "python-bindings", gen_stub_pymethods, pymethods)]
impl AdaptableFIR {
#[new]
fn new_py(fs: StrictlyPositive, block_size: usize, transition_time: StrictlyPositive) -> Self {
Self::new(fs, block_size, transition_time, None)
}
#[pyo3(name = "updateCoefficients")]
fn updateCoefficients_py(&mut self, coefs: PyReadonlyArray1<Flt>) -> Result<()> {
let coefs = coefs.as_slice().unwrap();
self.updateCoefficients(coefs)
}
#[pyo3(name = "filter")]
fn filter_py<'py>(
&mut self,
py: Python<'py>,
input: PyReadonlyArray1<Flt>,
) -> PyResult<Bound<'py, PyArray1<Flt>>> {
let mut output = vec![0.0; input.len()?];
match input.as_slice().ok() {
Some(i) => {
self.filter(i, &mut output);
Ok(output.into_pyarray(py))
}
None => {
let input = &input.as_array().iter().copied().collect::<Vec<_>>();
self.filter(input, &mut output);
Ok(output.into_pyarray(py))
}
}
}
}
#[inline]
fn calc_gains(from_0to1: bool, n: usize, transition_samples: usize) -> (Flt, Flt) {
let gainA = skewsine((n as Flt) / (transition_samples - 1) as Flt);
let gainB = 1. - gainA;
if from_0to1 {
(gainB, gainA)
} else {
(gainA, gainB)
}
}
#[inline]
fn skewsine(val: Flt) -> Flt {
if val < 0. {
return 0.;
}
if val > 1. {
return 1.;
}
val - 1. / twopi * Flt::sin(twopi * val)
}