1use crate::context::{max_variables, set_truncation_order, truncation_order};
12use crate::da::Da;
13use crate::error::{codes, dace_panic};
14use crate::eval::CompiledDa;
15
16pub trait DaVector {
18 fn cons(&self) -> Vec<f64>;
20
21 fn linear(&self) -> Vec<Vec<f64>>;
23
24 fn deriv(&self, var: u32) -> Vec<Da>;
26
27 fn integ(&self, var: u32) -> Vec<Da>;
29
30 fn eval(&self, args: &[f64]) -> Vec<f64>;
32
33 fn plug(&self, var: u32, val: f64) -> Vec<Da>;
35
36 fn trim(&self, min_order: u32, max_order: u32) -> Vec<Da>;
38
39 fn invert(&self) -> Vec<Da>;
46}
47
48impl DaVector for [Da] {
49 fn cons(&self) -> Vec<f64> {
50 self.iter().map(|d| d.cons()).collect()
51 }
52
53 fn linear(&self) -> Vec<Vec<f64>> {
54 self.iter().map(|d| d.linear()).collect()
55 }
56
57 fn deriv(&self, var: u32) -> Vec<Da> {
58 self.iter().map(|d| d.deriv(var)).collect()
59 }
60
61 fn integ(&self, var: u32) -> Vec<Da> {
62 self.iter().map(|d| d.integ(var)).collect()
63 }
64
65 fn eval(&self, args: &[f64]) -> Vec<f64> {
66 self.iter().map(|d| d.eval(args)).collect()
67 }
68
69 fn plug(&self, var: u32, val: f64) -> Vec<Da> {
70 self.iter().map(|d| d.plug(var, val)).collect()
71 }
72
73 fn trim(&self, min_order: u32, max_order: u32) -> Vec<Da> {
74 self.iter().map(|d| d.trim(min_order, max_order)).collect()
75 }
76
77 fn invert(&self) -> Vec<Da> {
78 let ord = truncation_order();
79 let nvar = self.len();
80 if nvar > max_variables() as usize {
81 dace_panic(
82 codes::TOO_MANY_VARIABLES,
83 "dimension of vector exceeds maximum number of DA variables",
84 );
85 }
86
87 let dda: Vec<Da> = (1..=nvar as u32).map(Da::variable).collect();
89
90 let ac: Vec<f64> = self.cons();
93 let m: Vec<Da> = self.iter().map(|d| d.trim(1, u32::MAX)).collect();
94 let an: Vec<Da> = m.iter().map(|d| d.trim(2, u32::MAX)).collect();
95
96 let mut ai = m.linear();
99 matrix_inverse(&mut ai);
100
101 let aloan: Vec<Da> = (0..nvar)
103 .map(|i| (0..nvar).fold(Da::constant(0.0), |acc, j| acc + ai[i][j] * an[j].clone()))
104 .collect();
105 let aioan = CompiledDa::from_das(&aloan);
106 let linv: Vec<Da> = (0..nvar)
107 .map(|i| (0..nvar).fold(Da::constant(0.0), |acc, j| acc + ai[i][j] * dda[j].clone()))
108 .collect();
109
110 let mut mi = linv.clone();
112 for i in 1..ord {
113 set_truncation_order(i + 1);
114 let correction = aioan.eval_da(&mi);
115 mi = linv
116 .iter()
117 .zip(correction)
118 .map(|(l, c)| l.clone() - c)
119 .collect();
120 }
121 set_truncation_order(ord);
122
123 let args: Vec<Da> = dda.iter().zip(&ac).map(|(d, &c)| d.clone() - c).collect();
125 mi.iter().map(|m| m.eval_da(&args)).collect()
126 }
127}
128
129impl DaVector for Vec<Da> {
130 fn cons(&self) -> Vec<f64> {
131 self.as_slice().cons()
132 }
133
134 fn linear(&self) -> Vec<Vec<f64>> {
135 self.as_slice().linear()
136 }
137
138 fn deriv(&self, var: u32) -> Vec<Da> {
139 self.as_slice().deriv(var)
140 }
141
142 fn integ(&self, var: u32) -> Vec<Da> {
143 self.as_slice().integ(var)
144 }
145
146 fn eval(&self, args: &[f64]) -> Vec<f64> {
147 self.as_slice().eval(args)
148 }
149
150 fn plug(&self, var: u32, val: f64) -> Vec<Da> {
151 self.as_slice().plug(var, val)
152 }
153
154 fn trim(&self, min_order: u32, max_order: u32) -> Vec<Da> {
155 self.as_slice().trim(min_order, max_order)
156 }
157
158 fn invert(&self) -> Vec<Da> {
159 self.as_slice().invert()
160 }
161}
162
163fn matrix_inverse(a: &mut [Vec<f64>]) {
170 let n = a.len();
171 let mut indexc = vec![0usize; n];
172 let mut indexr = vec![0usize; n];
173 let mut ipiv = vec![0usize; n];
174
175 for i in 0..n {
176 let mut icol = 0usize;
177 let mut irow = 0usize;
178 let mut big = 0.0f64;
179 for (j, jp) in ipiv.iter().enumerate() {
180 if *jp != 0 {
181 continue;
182 }
183 for (k, kp) in ipiv.iter().enumerate() {
184 if *kp == 0 && a[j][k].abs() >= big {
185 big = a[j][k].abs();
186 irow = j;
187 icol = k;
188 }
189 }
190 }
191 ipiv[icol] = 1;
192 if irow != icol {
193 a.swap(irow, icol);
194 }
195 indexr[i] = irow;
196 indexc[i] = icol;
197 if a[icol][icol] == 0.0 {
198 dace_panic(
199 codes::INVERSE_DOES_NOT_EXIST,
200 "linear matrix inverse does not exist",
201 );
202 }
203 let pivinv = 1.0 / a[icol][icol];
204 a[icol][icol] = 1.0;
205 for v in a[icol].iter_mut() {
206 *v *= pivinv;
207 }
208 for ll in 0..n {
209 if ll != icol {
210 let temp = a[ll][icol];
211 a[ll][icol] = 0.0;
212 let src = a[icol].clone();
213 for (v, s) in a[ll].iter_mut().zip(&src) {
214 *v -= s * temp;
215 }
216 }
217 }
218 }
219
220 for i in (0..n).rev() {
222 if indexr[i] != indexc[i] {
223 for row in a.iter_mut() {
224 row.swap(indexr[i], indexc[i]);
225 }
226 }
227 }
228}
229
230pub fn dot(a: &[f64], b: &[f64]) -> f64 {
236 assert_eq!(a.len(), b.len(), "dot: length mismatch");
237 a.iter().zip(b).map(|(x, y)| x * y).sum()
238}
239
240pub fn dot_da(a: &[Da]) -> Da {
242 a.iter()
243 .skip(1)
244 .fold(a[0].clone(), |acc, d| acc + d.clone())
245}
246
247pub fn cross(a: &[f64], b: &[f64]) -> Vec<f64> {
253 assert_eq!(a.len(), 3, "cross: not a 3-vector");
254 assert_eq!(b.len(), 3, "cross: not a 3-vector");
255 vec![
256 a[1] * b[2] - a[2] * b[1],
257 a[2] * b[0] - a[0] * b[2],
258 a[0] * b[1] - a[1] * b[0],
259 ]
260}
261
262pub fn vnorm(a: &[f64]) -> f64 {
264 a.iter().map(|x| x * x).sum::<f64>().sqrt()
265}
266
267pub fn normalize(a: &[f64]) -> Vec<f64> {
269 let n = vnorm(a);
270 assert!(n > 0.0, "normalize: zero vector");
271 a.iter().map(|x| x / n).collect()
272}
273
274#[cfg(test)]
275mod tests {
276 use super::*;
277 use crate::test_support::CONTEXT_LOCK;
278
279 #[test]
280 fn invert_roundtrip() {
281 let _g = CONTEXT_LOCK.lock();
282 crate::context::init(8, 3).unwrap();
283 let x = Da::variable(1);
284 let y = Da::variable(2);
285 let z = Da::variable(3);
286
287 let map = vec![
289 x.clone() + 0.3 * (x.clone() * y.clone()),
290 y.clone() - 0.2 * (y.clone() * z.clone()),
291 z.clone() + 0.1 * x.clone() * z.clone(),
292 ];
293
294 let inv = map.invert();
295 assert_eq!(inv.len(), 3);
296
297 for &(px, py, pz) in &[(0.05, -0.03, 0.04), (-0.08, 0.06, 0.02), (0.0, 0.0, 0.0)] {
299 let img = map.eval(&[px, py, pz]);
300 let back = inv.eval(&img);
301 assert!((back[0] - px).abs() < 1e-10, "x: {} vs {px}", back[0]);
302 assert!((back[1] - py).abs() < 1e-10, "y: {} vs {py}", back[1]);
303 assert!((back[2] - pz).abs() < 1e-10, "z: {} vs {pz}", back[2]);
304 }
305
306 assert_eq!(inv.cons(), vec![0.0, 0.0, 0.0]);
308 }
309
310 #[test]
311 fn scalar_helpers() {
312 assert_eq!(dot(&[1.0, 2.0, 3.0], &[4.0, -5.0, 6.0]), 12.0);
313 assert_eq!(
314 cross(&[1.0, 0.0, 0.0], &[0.0, 1.0, 0.0]),
315 vec![0.0, 0.0, 1.0]
316 );
317 assert!((vnorm(&[3.0, 4.0]) - 5.0).abs() < 1e-15);
318 assert_eq!(normalize(&[3.0, 4.0]), vec![0.6, 0.8]);
319 }
320}