butterworth 0.1.0

A library for simple Butterworth filters.
Documentation
use crate::butter::transfer_from_frequency;
use crate::error::{check_data_length, FilterError};

/// Represents a digital Butterworth filter in transfer function form. The filter can be applied to
/// time series data using the `bidirectional`, `bidirectional_with_padding`, and `forward`
/// functions.
#[derive(Clone, Debug, PartialEq)]
pub struct Filter {
    numerator: Vec<f64>,
    denominator: Vec<f64>,
    order: usize,
    sample_rate: f64,
    cutoff: Cutoff,
    initial_state: Vec<f64>
}

/// Represents filter type and cutoff frequencies for a Butterworth filter.
#[derive(Clone, Debug, PartialEq)]
pub enum Cutoff {
    LowPass(f64),
    HighPass(f64),
    BandPass(f64, f64),
    BandStop(f64, f64),
}

impl Filter {
    /// Create a new Butterworth filter with the specified order, sample rate, and `Cutoff` variant.
    /// Note that the filter order is doubled for bandpass and bandstop filters for correct use of
    /// coefficients as these filter types are effectively two filters in series.
    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
        })
    }

    /// Get the sample rate of the filter.
    pub fn sample_rate(&self) -> f64 {
        self.sample_rate
    }

    /// Get the order of the filter.
    pub fn order(&self) -> usize {
        self.order
    }

    /// Get a clone of the Cutoff variant of the filter describing the filter type and cutoff
    /// frequencies.
    pub fn cutoff(&self) -> Cutoff {
        self.cutoff.clone()
    }

    fn initial_state(&self) -> Vec<f64> {
        self.initial_state.clone()
    }

    /// Apply a bidirectional filter to the input data. This function is designed to match MATLAB's
    /// `filtfilt` function.
    pub fn bidirectional(&self, data: &Vec<f64>) -> Result<Vec<f64>, FilterError> {
        let padding_length = self.order * 3;
        self.bidirectional_with_padding(data, padding_length)
    }

    /// Apply a bidirectional filter to the input data with a specified 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)
    }

    /// Apply a forward filter to the input data. Note that unlike the bidirectional filter, using
    /// `filter` alone introduces phase lag.
    pub fn forward(&self, data: &Vec<f64>) -> Vec<f64> {
        let mut output = vec![0.; data.len()];
        let mut state = self.initial_state();
        // Scale state with first data value
        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
    }
}

/// Create initial state for filter based on numerator and denominator coefficients and filter
/// order. This is performed once on filter creation.
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];
    }

    // Solve sparse system of equations for initial state
    let solution_sum: f64 = solution.iter().sum();
    let denominator_sum: f64 = denominator.iter().skip(1).sum();
    initial_state[0] = solution_sum / (denominator_sum + 1.);
    // Solve remaining states
    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);
    }
    // More strict test outside edges
    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.));
    }

}