use crate::butter::transfer_from_frequency;
use crate::error::{check_data_length, FilterError};
#[derive(Clone, Debug, PartialEq)]
pub struct Filter {
numerator: Vec<f64>,
denominator: Vec<f64>,
order: usize,
sample_rate: f64,
cutoff: Cutoff,
initial_state: Vec<f64>
}
#[derive(Clone, Debug, PartialEq)]
pub enum Cutoff {
LowPass(f64),
HighPass(f64),
BandPass(f64, f64),
BandStop(f64, f64),
}
impl Filter {
pub fn new(mut order: usize, sample_rate: f64, cutoff: Cutoff) -> Result<Self, FilterError> {
let (numerator, denominator) = transfer_from_frequency(order, sample_rate, cutoff.clone())?;
match cutoff {
Cutoff::BandPass(_, _) | Cutoff::BandStop(_, _) => {
order = order * 2;
}
_ => {}
}
let initial_state = create_initial_state(&numerator, &denominator, order);
Ok(Filter {
numerator,
denominator,
order,
sample_rate,
cutoff,
initial_state
})
}
pub fn sample_rate(&self) -> f64 {
self.sample_rate
}
pub fn order(&self) -> usize {
self.order
}
pub fn cutoff(&self) -> Cutoff {
self.cutoff.clone()
}
fn initial_state(&self) -> Vec<f64> {
self.initial_state.clone()
}
pub fn bidirectional(&self, data: &Vec<f64>) -> Result<Vec<f64>, FilterError> {
let padding_length = self.order * 3;
self.bidirectional_with_padding(data, padding_length)
}
pub fn bidirectional_with_padding(&self, data: &Vec<f64>, padding_length: usize)
-> Result<Vec<f64>, FilterError> {
check_data_length(data, padding_length)?;
let mut padded_data = vec![0.; padding_length * 2 + data.len()];
for i in 0..padding_length {
padded_data[i] = 2. * data[0] - data[padding_length - i];
padded_data[i + data.len() + padding_length] = 2. * data[data.len() - 1] - data[data.len() - 2 - i];
}
for i in 0..data.len() {
padded_data[i + padding_length] = data[i];
}
let forward = self.forward(&padded_data);
let reverse = self.forward(&forward.iter().rev().map(|x| *x).collect());
let mut output = vec![0.; data.len()];
for i in 0..data.len() {
output[i] = reverse[reverse.len() - padding_length - i - 1];
}
Ok(output)
}
pub fn forward(&self, data: &Vec<f64>) -> Vec<f64> {
let mut output = vec![0.; data.len()];
let mut state = self.initial_state();
for i in 0..self.order {
state[i] *= data[0];
}
let mut dot = vec![0.; self.order];
for i in 0..data.len() {
output[i] = self.numerator[0] * data[i] + state[0];
for j in 0..self.order - 1 {
dot[j] = -self.denominator[j + 1] * state[0] + state[j + 1];
}
dot[self.order - 1] = -self.denominator[self.order] * state[0];
for j in 0..self.order {
state[j] = dot[j] + (self.numerator[j + 1] - self.numerator[0] * self.denominator[j + 1]) * data[i];
}
}
output
}
}
fn create_initial_state(numerator: &Vec<f64>, denominator: &Vec<f64>, order: usize) -> Vec<f64> {
let mut initial_state = vec![0.; order];
let mut solution = vec![0.; order];
for i in 0..order {
solution[i] = numerator[i + 1] - numerator[0] * denominator[i + 1];
}
let solution_sum: f64 = solution.iter().sum();
let denominator_sum: f64 = denominator.iter().skip(1).sum();
initial_state[0] = solution_sum / (denominator_sum + 1.);
for i in 0..(order - 1) {
initial_state[i + 1] = denominator[i + 1] * initial_state[0]
+ initial_state[i] - solution[i];
}
initial_state
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_bidirectional_arbitrary_low() {
let filter = Filter::new(5, 81., Cutoff::LowPass(1.)).unwrap();
let data = vec![0., 1., 2., 3., 4., 5., 6., 7., 8., 9., 10., 11., 12., 13., 14., 15., 16., 17., 18., 19.];
let output = filter.bidirectional(&data).unwrap();
let expected = vec![-8.01042949, -7.92309844, -7.84282517, -7.76943633, -7.70272181, -7.64243774, -7.58830962, -7.54003594, -7.49729197, -7.45973391, -7.42700307, -7.39873029, -7.37454044, -7.35405680, -7.33690554, -7.32271990, -7.31114433, -7.30183816, -7.29447912, -7.28876625];
for i in 0..output.len() {
assert!((output[i] - expected[i]).abs() < 1e-8);
}
}
#[test]
fn test_bidirectional_arbitrary_low_even() {
let filter = Filter::new(4, 81., Cutoff::LowPass(1.)).unwrap();
let data = vec![0., 1., 2., 3., 4., 5., 6., 7., 8., 9., 10., 11., 12., 13., 14., 15., 16., 17., 18., 19.];
let output = filter.bidirectional(&data).unwrap();
let expected = vec![-4.40223009, -4.24702897, -4.10157216, -3.96598745, -3.84032942, -3.72457907, -3.61864436, -3.52236161, -3.43549778, -3.35775377, -3.28876856, -3.22812429, -3.17535224, -3.12993963, -3.09133713, -3.05896719, -3.03223281, -3.01052685, -2.99324160, -2.97977856];
for i in 0..output.len() {
assert!((output[i] - expected[i]).abs() < 1e-8);
}
}
#[test]
fn test_bidirectional_arbitrary_high() {
let filter = Filter::new(5, 81., Cutoff::HighPass(1.)).unwrap();
let data = vec![0., 1., 2., 3., 4., 5., 6., 7., 8., 9., 10., 11., 12., 13., 14., 15., 16., 17., 18., 19.];
let output = filter.bidirectional(&data).unwrap();
let expected = vec![0.97001270, 1.07085149, 1.14880480, 1.20477865, 1.23974996, 1.25476586, 1.25094293, 1.22946634, 1.19158893, 1.13863025, 1.07197548, 0.99307439, 0.90344019, 0.80464836, 0.69833547, 0.58619796, 0.46999086, 0.35152658, 0.23267361, 0.11535522];
for i in 0..output.len() {
assert!((output[i] - expected[i]).abs() < 1e-6);
}
}
#[test]
fn test_bidirectional_arbitrary_band_pass() {
let filter = Filter::new(3, 81., Cutoff::BandPass(1., 10.)).unwrap();
let data = vec![0., 1., 2., 3., 4., 5., 6., 7., 8., 9., 10., 11., 12., 13., 14., 15., 16., 17., 18., 19.];
let output = filter.bidirectional(&data).unwrap();
let expected = vec![0.97054031, 1.03467778, 1.08183994, 1.11279947, 1.12823833, 1.12892733, 1.11579913, 1.08993227, 1.05249333, 1.00468001, 0.94768740, 0.88269768, 0.81088195, 0.73340281, 0.65141423, 0.56606446, 0.47851213, 0.38996018, 0.30169734, 0.21511972];
for i in 0..output.len() {
assert!((output[i] - expected[i]).abs() < 1e-6);
}
}
#[test]
fn test_bidirectional_arbitrary_band_stop() {
let filter = Filter::new(3, 81., Cutoff::BandStop(1., 10.)).unwrap();
let data = vec![0., 1., 2., 3., 4., 5., 6., 7., 8., 9., 10., 11., 12., 13., 14., 15., 16., 17., 18., 19.];
let output = filter.bidirectional(&data).unwrap();
let expected = vec![0.53645593, 1.19729383, 1.84359749, 2.47402623, 3.08745172, 3.68277859, 4.25887845, 4.81462623, 5.34899261, 5.86113966, 6.35047539, 6.81663938, 7.25941371, 7.67858615, 8.07383114, 8.44470296, 8.79082373, 9.11227133, 9.41002743, 9.68618722];
for i in 0..output.len() {
assert!((output[i] - expected[i]).abs() < 1e-6);
}
}
#[test]
fn test_recover_sin_low() {
let filter = Filter::new(4, 100., Cutoff::LowPass(8.)).unwrap();
let data = (0..=100).map(|x| x as f64).map(|x| (x * 0.1).sin() + (x * 0.75).sin()).collect();
let output = filter.bidirectional(&data).unwrap();
let expected = (0..=100).map(|x| x as f64).map(|x| (x * 0.1).sin()).collect::<Vec<f64>>();
for i in 0..output.len() {
assert!((output[i] - expected[i]).abs() < 5e-1);
}
for i in 8..output.len() - 8 {
assert!((output[i] - expected[i]).abs() < 5e-2);
}
}
#[test]
fn test_get_properties() {
let filter = Filter::new(4, 100., Cutoff::LowPass(8.)).unwrap();
assert_eq!(filter.sample_rate(), 100.);
assert_eq!(filter.order(), 4);
assert_eq!(filter.cutoff(), Cutoff::LowPass(8.));
}
}