use num_complex::Complex64 as c64;
use super::SMatrix;
pub fn power(s: &[SMatrix], q: usize, p: usize) -> (f64, Vec<SMatrix>) {
let value = s.iter().map(|m| m.power(q, p)).sum();
let sensitivities = s
.iter()
.map(|m| only(m.size(), q, p, m[(q, p)].conj()))
.collect();
(value, sensitivities)
}
pub fn power_error(s: &[SMatrix], q: usize, p: usize, targets: &[f64]) -> (f64, Vec<SMatrix>) {
assert_eq!(s.len(), targets.len(), "one target per wavelength");
let mut value = 0.0;
let sensitivities = s
.iter()
.zip(targets)
.map(|(m, &t)| {
let miss = m.power(q, p) - t;
value += miss * miss;
only(m.size(), q, p, 2.0 * miss * m[(q, p)].conj())
})
.collect();
(value, sensitivities)
}
pub fn matrix_error(s: &[SMatrix], targets: &[SMatrix]) -> (f64, Vec<SMatrix>) {
assert_eq!(s.len(), targets.len(), "one target per wavelength");
let mut value = 0.0;
let sensitivities = s
.iter()
.zip(targets)
.map(|(m, t)| {
assert_eq!(m.size(), t.size(), "a target of the S-matrix's size");
let g = SMatrix::from_fn(m.size(), |q, p| (m[(q, p)] - t[(q, p)]).conj());
value += g.rows().iter().flatten().map(c64::norm_sqr).sum::<f64>();
g
})
.collect();
(value, sensitivities)
}
fn only(n: usize, q: usize, p: usize, v: c64) -> SMatrix {
let mut m = SMatrix::zeros(n);
m[(q, p)] = v;
m
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_module_example() {
let s = SMatrix::from_fn(2, |q, p| {
if q == p {
c64::new(0.0, 0.0)
} else {
c64::new(0.6, 0.3)
}
});
let (power, g) = power(std::slice::from_ref(&s), 1, 0);
assert!((power - 0.45).abs() < 1e-15);
assert_eq!(g[0][(1, 0)], c64::new(0.6, -0.3)); let (error, _) = power_error(&[s], 1, 0, &[0.5]);
assert!((error - 0.05f64.powi(2)).abs() < 1e-15);
}
#[test]
fn sensitivities_are_the_wirtinger_derivatives() {
let s = SMatrix::from_fn(3, |q, p| {
c64::new(0.1 * q as f64 - 0.2, 0.3 + 0.05 * p as f64)
});
let t = SMatrix::from_fn(3, |q, p| c64::new(0.2, -0.1 * (q + p) as f64));
type Objective<'a> = &'a dyn Fn(&[SMatrix]) -> (f64, Vec<SMatrix>);
let objectives: [Objective; 3] = [
&|m| power(m, 2, 1),
&|m| power_error(m, 2, 1, &[0.3]),
&|m| matrix_error(m, std::slice::from_ref(&t)),
];
for f in objectives {
let g = &f(std::slice::from_ref(&s)).1[0];
for (q, p) in [(2, 1), (0, 2)] {
for step in [c64::new(1e-6, 0.0), c64::new(0.0, 1e-6)] {
let mut plus = s.clone();
plus[(q, p)] += step;
let mut minus = s.clone();
minus[(q, p)] -= step;
let fd = (f(&[plus]).0 - f(&[minus]).0) / 2.0;
let predicted = 2.0 * (g[(q, p)] * step).re;
assert!((fd - predicted).abs() < 1e-14, "{fd} {predicted}");
}
}
}
}
}