use crate::error::LinalgError;
use crate::linear_algebra::Matrix;
use crate::scalar::Numeric;
const MAXIMUM_PASSES: usize = 64;
pub fn solve_discrete_riccati<const N: usize, const M: usize, T: Numeric>(
a: Matrix<N, N, T>,
b: Matrix<N, M, T>,
q: Matrix<N, N, T>,
r: Matrix<M, M, T>,
) -> Result<Matrix<N, N, T>, LinalgError> {
if !a.is_finite() || !b.is_finite() || !q.is_finite() || !r.is_finite() {
return Err(LinalgError::NonFinite);
}
if !q.is_symmetric() || !r.is_symmetric() {
return Err(LinalgError::NotSymmetric);
}
let input_cost = r.cholesky()?;
let scaled_input = input_cost.solve_matrix::<N>(b.transpose());
let mut reach = b * scaled_input;
let mut state = a;
let mut cost = q;
let mut passes_taken = MAXIMUM_PASSES;
for pass in 0..MAXIMUM_PASSES {
let coupling = (Matrix::<N, N, T>::identity() + reach * cost).inverse()?;
let folded = coupling * state;
let cost_increment = state.transpose() * cost * folded;
let reach_increment = state * (coupling * reach) * state.transpose();
let next_state = state * folded;
let next_cost = (cost + cost_increment).symmetrized();
let next_reach = (reach + reach_increment).symmetrized();
state = next_state;
cost = next_cost;
reach = next_reach;
let increment_size = cost_increment.frobenius_norm();
let cost_size = cost.frobenius_norm();
if !cost.is_finite()
|| !reach.is_finite()
|| !state.is_finite()
|| !increment_size.is_finite()
|| !cost_size.is_finite()
{
return Err(LinalgError::DidNotConverge { iters: pass + 1 });
}
if increment_size <= T::EPSILON_X30 * cost_size.max(T::ONE) {
passes_taken = pass + 1;
break;
}
}
let input_weight = r + b.transpose() * cost * b;
let input_weight_factor = input_weight.cholesky()?;
let coupling_term = b.transpose() * cost * a;
let correction =
coupling_term.transpose() * input_weight_factor.solve_matrix::<N>(coupling_term);
let residual = a.transpose() * cost * a - cost - correction + q;
let allowed = T::EPSILON.sqrt() * cost.frobenius_norm().max(T::ONE);
if !residual.is_finite() || residual.frobenius_norm() > allowed {
return Err(LinalgError::DidNotConverge {
iters: passes_taken,
});
}
Ok(cost)
}