use crate::error::LinalgError;
use crate::linear_algebra::Matrix;
use crate::scalar::Numeric;
pub fn zoh<const N: usize, const M: usize, const NM: usize, T: Numeric>(
a: Matrix<N, N, T>,
b: Matrix<N, M, T>,
dt: T,
) -> Result<(Matrix<N, N, T>, Matrix<N, M, T>), LinalgError> {
const { assert!(NM == N + M, "zoh: NM must equal N + M") };
let aug = Matrix::<NM, NM, T>::from_fn(|i, j| {
if i < N && j < N {
a[(i, j)] * dt
} else if i < N {
b[(i, j - N)] * dt
} else {
T::ZERO
}
});
let e = aug.expm()?;
let f = Matrix::<N, N, T>::from_fn(|i, j| e[(i, j)]);
let g = Matrix::<N, M, T>::from_fn(|i, j| e[(i, N + j)]);
Ok((f, g))
}
pub fn van_loan<const N: usize, const N2: usize, T: Numeric>(
a: Matrix<N, N, T>,
qc: Matrix<N, N, T>,
dt: T,
) -> Result<(Matrix<N, N, T>, Matrix<N, N, T>), LinalgError> {
const { assert!(N2 == 2 * N, "van_loan: N2 must equal 2*N") };
let xi = Matrix::<N2, N2, T>::from_fn(|i, j| {
if i < N && j < N {
-a[(i, j)] * dt
} else if i < N {
qc[(i, j - N)] * dt
} else if j >= N {
a[(j - N, i - N)] * dt } else {
T::ZERO
}
});
let e = xi.expm()?;
let g12 = Matrix::<N, N, T>::from_fn(|i, j| e[(i, N + j)]);
let g22 = Matrix::<N, N, T>::from_fn(|i, j| e[(N + i, N + j)]);
let f = g22.transpose();
let qd = f * g12;
Ok((f, qd))
}
pub fn q_discrete_white_noise<const DIM: usize, T: Numeric>(
dt: T,
variance: T,
) -> Matrix<DIM, DIM, T> {
const {
assert!(
DIM >= 2 && DIM <= 4,
"q_discrete_white_noise: DIM must be 2, 3, or 4"
)
};
let dt2 = dt * dt;
let dt3 = dt2 * dt;
let dt4 = dt3 * dt;
let dt5 = dt4 * dt;
let dt6 = dt5 * dt;
Matrix::from_fn(|i, j| {
let c = match (DIM, i, j) {
(2, 0, 0) => dt4 * T::from_f64(0.25),
(2, 0, 1) | (2, 1, 0) => dt3 * T::HALF,
(2, 1, 1) => dt2,
(3, 0, 0) => dt4 * T::from_f64(0.25),
(3, 0, 1) | (3, 1, 0) => dt3 * T::HALF,
(3, 0, 2) | (3, 2, 0) => dt2 * T::HALF,
(3, 1, 1) => dt2,
(3, 1, 2) | (3, 2, 1) => dt,
(3, 2, 2) => T::ONE,
(4, 0, 0) => dt6 / T::from_f64(36.0),
(4, 0, 1) | (4, 1, 0) => dt5 / T::from_f64(12.0),
(4, 0, 2) | (4, 2, 0) => dt4 / T::from_f64(6.0),
(4, 0, 3) | (4, 3, 0) => dt3 / T::from_f64(6.0),
(4, 1, 1) => dt4 * T::from_f64(0.25),
(4, 1, 2) | (4, 2, 1) => dt3 * T::HALF,
(4, 1, 3) | (4, 3, 1) => dt2 * T::HALF,
(4, 2, 2) => dt2,
(4, 2, 3) | (4, 3, 2) => dt,
(4, 3, 3) => T::ONE,
_ => T::ZERO,
};
c * variance
})
}