#[derive(Debug, Clone)]
pub struct KalmanFilter {
x: Vec<f64>,
p: Vec<Vec<f64>>,
f: Vec<Vec<f64>>,
h: Vec<Vec<f64>>,
q: Vec<Vec<f64>>,
r: Vec<Vec<f64>>,
n: usize,
m: usize,
}
impl KalmanFilter {
pub fn new(f: Vec<Vec<f64>>, h: Vec<Vec<f64>>, q: Vec<Vec<f64>>, r: Vec<Vec<f64>>) -> Self {
let n = f.len();
let m = h.len();
let x = vec![0.0; n];
let p = (0..n)
.map(|i| (0..n).map(|j| if i == j { 1.0 } else { 0.0 }).collect())
.collect();
Self {
x,
p,
f,
h,
q,
r,
n,
m,
}
}
pub fn set_state(&mut self, x: Vec<f64>) {
assert_eq!(x.len(), self.n, "State dimension mismatch");
self.x = x;
}
pub fn set_covariance(&mut self, p: Vec<Vec<f64>>) {
assert_eq!(p.len(), self.n, "Covariance dimension mismatch");
self.p = p;
}
pub fn state(&self) -> &[f64] {
&self.x
}
pub fn covariance(&self) -> &[Vec<f64>] {
&self.p
}
pub fn predict(&mut self) {
self.x = mat_vec_mul(&self.f, &self.x);
let fp = mat_mul(&self.f, &self.p);
let fpf = mat_mul(&fp, &transpose(&self.f));
self.p = mat_add(&fpf, &self.q);
}
pub fn update(&mut self, z: &[f64]) {
assert_eq!(z.len(), self.m, "Measurement dimension mismatch");
let hx = mat_vec_mul(&self.h, &self.x);
let y: Vec<f64> = z.iter().zip(hx.iter()).map(|(a, b)| a - b).collect();
let hp = mat_mul(&self.h, &self.p);
let hph = mat_mul(&hp, &transpose(&self.h));
let s = mat_add(&hph, &self.r);
let ht = transpose(&self.h);
let pht = mat_mul(&self.p, &ht);
let s_inv = mat_inv(&s);
let k = mat_mul(&pht, &s_inv);
let ky = mat_vec_mul(&k, &y);
self.x = vec_add(&self.x, &ky);
let kh = mat_mul(&k, &self.h);
let i_kh = mat_sub(&identity(self.n), &kh);
self.p = mat_mul(&i_kh, &self.p);
}
}
fn mat_vec_mul(a: &[Vec<f64>], x: &[f64]) -> Vec<f64> {
a.iter()
.map(|row| row.iter().zip(x.iter()).map(|(a, b)| a * b).sum())
.collect()
}
fn mat_mul(a: &[Vec<f64>], b: &[Vec<f64>]) -> Vec<Vec<f64>> {
let bt = transpose(b);
a.iter()
.map(|row| {
bt.iter()
.map(|col| row.iter().zip(col.iter()).map(|(a, b)| a * b).sum())
.collect()
})
.collect()
}
fn mat_add(a: &[Vec<f64>], b: &[Vec<f64>]) -> Vec<Vec<f64>> {
a.iter()
.zip(b.iter())
.map(|(row_a, row_b)| row_a.iter().zip(row_b.iter()).map(|(a, b)| a + b).collect())
.collect()
}
fn mat_sub(a: &[Vec<f64>], b: &[Vec<f64>]) -> Vec<Vec<f64>> {
a.iter()
.zip(b.iter())
.map(|(row_a, row_b)| row_a.iter().zip(row_b.iter()).map(|(a, b)| a - b).collect())
.collect()
}
fn vec_add(a: &[f64], b: &[f64]) -> Vec<f64> {
a.iter().zip(b.iter()).map(|(x, y)| x + y).collect()
}
fn transpose(m: &[Vec<f64>]) -> Vec<Vec<f64>> {
let rows = m.len();
let cols = m[0].len();
(0..cols)
.map(|j| (0..rows).map(|i| m[i][j]).collect())
.collect()
}
fn identity(n: usize) -> Vec<Vec<f64>> {
(0..n)
.map(|i| (0..n).map(|j| if i == j { 1.0 } else { 0.0 }).collect())
.collect()
}
fn mat_inv(m: &[Vec<f64>]) -> Vec<Vec<f64>> {
let n = m.len();
assert_eq!(n, m[0].len(), "Matrix must be square");
let mut aug = m
.iter()
.enumerate()
.map(|(i, row)| {
let mut r = row.clone();
r.extend((0..n).map(|j| if i == j { 1.0 } else { 0.0 }));
r
})
.collect::<Vec<_>>();
for i in 0..n {
let mut max_row = i;
for k in i + 1..n {
if aug[k][i].abs() > aug[max_row][i].abs() {
max_row = k;
}
}
aug.swap(i, max_row);
let pivot = aug[i][i];
for aug_val in aug[i].iter_mut().take(2 * n) {
*aug_val /= pivot;
}
for k in 0..n {
if k != i {
let factor = aug[k][i];
let aug_i_row = aug[i].clone();
for (j, aug_kj) in aug[k].iter_mut().enumerate().take(2 * n) {
*aug_kj -= factor * aug_i_row[j];
}
}
}
}
aug.iter().map(|row| row[n..].to_vec()).collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_kalman_1d() {
let f = vec![vec![1.0]];
let h = vec![vec![1.0]];
let q = vec![vec![0.01]];
let r = vec![vec![1.0]];
let mut kf = KalmanFilter::new(f, h, q, r);
kf.set_state(vec![0.0]);
kf.predict();
kf.update(&[1.0]);
assert!(kf.state()[0] > 0.0 && kf.state()[0] < 1.0);
}
#[test]
fn test_kalman_2d_velocity() {
let dt = 0.1;
let f = vec![vec![1.0, dt], vec![0.0, 1.0]];
let h = vec![vec![1.0, 0.0]];
let q = vec![vec![0.01, 0.0], vec![0.0, 0.01]];
let r = vec![vec![1.0]];
let mut kf = KalmanFilter::new(f, h, q, r);
kf.set_state(vec![0.0, 0.0]);
for i in 0..10 {
kf.predict();
let measurement = i as f64 * 0.5;
kf.update(&[measurement]);
}
assert!(kf.state()[0] > 3.0 && kf.state()[0] < 6.0);
}
}