use crate::context::{max_variables, set_truncation_order, truncation_order};
use crate::da::Da;
use crate::error::{codes, dace_panic};
use crate::eval::CompiledDa;
pub trait DaVector {
fn cons(&self) -> Vec<f64>;
fn linear(&self) -> Vec<Vec<f64>>;
fn deriv(&self, var: u32) -> Vec<Da>;
fn integ(&self, var: u32) -> Vec<Da>;
fn eval(&self, args: &[f64]) -> Vec<f64>;
fn plug(&self, var: u32, val: f64) -> Vec<Da>;
fn trim(&self, min_order: u32, max_order: u32) -> Vec<Da>;
fn invert(&self) -> Vec<Da>;
}
impl DaVector for [Da] {
fn cons(&self) -> Vec<f64> {
self.iter().map(|d| d.cons()).collect()
}
fn linear(&self) -> Vec<Vec<f64>> {
self.iter().map(|d| d.linear()).collect()
}
fn deriv(&self, var: u32) -> Vec<Da> {
self.iter().map(|d| d.deriv(var)).collect()
}
fn integ(&self, var: u32) -> Vec<Da> {
self.iter().map(|d| d.integ(var)).collect()
}
fn eval(&self, args: &[f64]) -> Vec<f64> {
self.iter().map(|d| d.eval(args)).collect()
}
fn plug(&self, var: u32, val: f64) -> Vec<Da> {
self.iter().map(|d| d.plug(var, val)).collect()
}
fn trim(&self, min_order: u32, max_order: u32) -> Vec<Da> {
self.iter().map(|d| d.trim(min_order, max_order)).collect()
}
fn invert(&self) -> Vec<Da> {
let ord = truncation_order();
let nvar = self.len();
if nvar > max_variables() as usize {
dace_panic(
codes::TOO_MANY_VARIABLES,
"dimension of vector exceeds maximum number of DA variables",
);
}
let dda: Vec<Da> = (1..=nvar as u32).map(Da::variable).collect();
let ac: Vec<f64> = self.cons();
let m: Vec<Da> = self.iter().map(|d| d.trim(1, u32::MAX)).collect();
let an: Vec<Da> = m.iter().map(|d| d.trim(2, u32::MAX)).collect();
let mut ai = m.linear();
matrix_inverse(&mut ai);
let aloan: Vec<Da> = (0..nvar)
.map(|i| (0..nvar).fold(Da::constant(0.0), |acc, j| acc + ai[i][j] * an[j].clone()))
.collect();
let aioan = CompiledDa::from_das(&aloan);
let linv: Vec<Da> = (0..nvar)
.map(|i| (0..nvar).fold(Da::constant(0.0), |acc, j| acc + ai[i][j] * dda[j].clone()))
.collect();
let mut mi = linv.clone();
for i in 1..ord {
set_truncation_order(i + 1);
let correction = aioan.eval_da(&mi);
mi = linv
.iter()
.zip(correction)
.map(|(l, c)| l.clone() - c)
.collect();
}
set_truncation_order(ord);
let args: Vec<Da> = dda.iter().zip(&ac).map(|(d, &c)| d.clone() - c).collect();
mi.iter().map(|m| m.eval_da(&args)).collect()
}
}
impl DaVector for Vec<Da> {
fn cons(&self) -> Vec<f64> {
self.as_slice().cons()
}
fn linear(&self) -> Vec<Vec<f64>> {
self.as_slice().linear()
}
fn deriv(&self, var: u32) -> Vec<Da> {
self.as_slice().deriv(var)
}
fn integ(&self, var: u32) -> Vec<Da> {
self.as_slice().integ(var)
}
fn eval(&self, args: &[f64]) -> Vec<f64> {
self.as_slice().eval(args)
}
fn plug(&self, var: u32, val: f64) -> Vec<Da> {
self.as_slice().plug(var, val)
}
fn trim(&self, min_order: u32, max_order: u32) -> Vec<Da> {
self.as_slice().trim(min_order, max_order)
}
fn invert(&self) -> Vec<Da> {
self.as_slice().invert()
}
}
fn matrix_inverse(a: &mut [Vec<f64>]) {
let n = a.len();
let mut indexc = vec![0usize; n];
let mut indexr = vec![0usize; n];
let mut ipiv = vec![0usize; n];
for i in 0..n {
let mut icol = 0usize;
let mut irow = 0usize;
let mut big = 0.0f64;
for (j, jp) in ipiv.iter().enumerate() {
if *jp != 0 {
continue;
}
for (k, kp) in ipiv.iter().enumerate() {
if *kp == 0 && a[j][k].abs() >= big {
big = a[j][k].abs();
irow = j;
icol = k;
}
}
}
ipiv[icol] = 1;
if irow != icol {
a.swap(irow, icol);
}
indexr[i] = irow;
indexc[i] = icol;
if a[icol][icol] == 0.0 {
dace_panic(
codes::INVERSE_DOES_NOT_EXIST,
"linear matrix inverse does not exist",
);
}
let pivinv = 1.0 / a[icol][icol];
a[icol][icol] = 1.0;
for v in a[icol].iter_mut() {
*v *= pivinv;
}
for ll in 0..n {
if ll != icol {
let temp = a[ll][icol];
a[ll][icol] = 0.0;
let src = a[icol].clone();
for (v, s) in a[ll].iter_mut().zip(&src) {
*v -= s * temp;
}
}
}
}
for i in (0..n).rev() {
if indexr[i] != indexc[i] {
for row in a.iter_mut() {
row.swap(indexr[i], indexc[i]);
}
}
}
}
pub fn dot(a: &[f64], b: &[f64]) -> f64 {
assert_eq!(a.len(), b.len(), "dot: length mismatch");
a.iter().zip(b).map(|(x, y)| x * y).sum()
}
pub fn dot_da(a: &[Da]) -> Da {
a.iter()
.skip(1)
.fold(a[0].clone(), |acc, d| acc + d.clone())
}
pub fn cross(a: &[f64], b: &[f64]) -> Vec<f64> {
assert_eq!(a.len(), 3, "cross: not a 3-vector");
assert_eq!(b.len(), 3, "cross: not a 3-vector");
vec![
a[1] * b[2] - a[2] * b[1],
a[2] * b[0] - a[0] * b[2],
a[0] * b[1] - a[1] * b[0],
]
}
pub fn vnorm(a: &[f64]) -> f64 {
a.iter().map(|x| x * x).sum::<f64>().sqrt()
}
pub fn normalize(a: &[f64]) -> Vec<f64> {
let n = vnorm(a);
assert!(n > 0.0, "normalize: zero vector");
a.iter().map(|x| x / n).collect()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_support::CONTEXT_LOCK;
#[test]
fn invert_roundtrip() {
let _g = CONTEXT_LOCK.lock();
crate::context::init(8, 3).unwrap();
let x = Da::variable(1);
let y = Da::variable(2);
let z = Da::variable(3);
let map = vec![
x.clone() + 0.3 * (x.clone() * y.clone()),
y.clone() - 0.2 * (y.clone() * z.clone()),
z.clone() + 0.1 * x.clone() * z.clone(),
];
let inv = map.invert();
assert_eq!(inv.len(), 3);
for &(px, py, pz) in &[(0.05, -0.03, 0.04), (-0.08, 0.06, 0.02), (0.0, 0.0, 0.0)] {
let img = map.eval(&[px, py, pz]);
let back = inv.eval(&img);
assert!((back[0] - px).abs() < 1e-10, "x: {} vs {px}", back[0]);
assert!((back[1] - py).abs() < 1e-10, "y: {} vs {py}", back[1]);
assert!((back[2] - pz).abs() < 1e-10, "z: {} vs {pz}", back[2]);
}
assert_eq!(inv.cons(), vec![0.0, 0.0, 0.0]);
}
#[test]
fn scalar_helpers() {
assert_eq!(dot(&[1.0, 2.0, 3.0], &[4.0, -5.0, 6.0]), 12.0);
assert_eq!(
cross(&[1.0, 0.0, 0.0], &[0.0, 1.0, 0.0]),
vec![0.0, 0.0, 1.0]
);
assert!((vnorm(&[3.0, 4.0]) - 5.0).abs() < 1e-15);
assert_eq!(normalize(&[3.0, 4.0]), vec![0.6, 0.8]);
}
}