#![forbid(unsafe_code)]
use std::io::Write;
use std::process::Stdio;
use fsci_conformance::CompareLedger;
use fsci_linalg::{DecompOptions, eigvals, funm_with_error};
use serde::{Deserialize, Serialize};
const REQUIRE_SCIPY_ENV: &str = "FSCI_REQUIRE_SCIPY_ORACLE";
const FUNM_REL_TOL: f64 = 1e-12;
const FUNCS: [&str; 4] = ["exp", "sin", "cos", "poly"];
#[derive(Debug, Clone, Serialize)]
struct Case {
case_id: String,
a: Vec<Vec<f64>>,
funcs: Vec<String>,
}
#[derive(Debug, Clone, Deserialize)]
struct Row {
case_id: String,
func: String,
complex_pairs: usize,
funm: Option<Vec<Vec<f64>>>,
funm_err: f64,
reference: Vec<Vec<f64>>,
}
struct Rng(u64);
impl Rng {
fn next(&mut self) -> f64 {
self.0 = self.0.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = self.0;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^= z >> 31;
(z >> 11) as f64 / (1_u64 << 53) as f64 * 2.0 - 1.0
}
}
fn complex_pairs(a: &[Vec<f64>]) -> usize {
let (_, im) = eigvals(a, DecompOptions::default()).expect("eigvals");
im.iter().filter(|v| v.abs() > 1e-8).count() / 2
}
fn random_with_two_pairs() -> (u64, Vec<Vec<f64>>) {
(1..1000_u64)
.find_map(|seed| {
let mut rng = Rng(seed);
let a: Vec<Vec<f64>> = (0..6)
.map(|_| (0..6).map(|_| rng.next()).collect())
.collect();
(complex_pairs(&a) >= 2).then_some((seed, a))
})
.expect("a seed below 1000 gives a 6x6 with two complex pairs")
}
fn cases() -> Vec<Case> {
let all: Vec<String> = FUNCS.iter().map(|s| (*s).to_string()).collect();
let (seed, random) = random_with_two_pairs();
println!(
"random 6x6: seed {seed}, {} complex pairs",
complex_pairs(&random)
);
vec![
Case {
case_id: "rotation".into(),
a: vec![vec![0.0, -1.0], vec![1.0, 0.0]],
funcs: all.clone(),
},
Case {
case_id: "m3_complex_pair".into(),
a: vec![
vec![1.0, 2.0, 0.0],
vec![-3.0, 1.0, 1.0],
vec![0.0, 0.5, 2.0],
],
funcs: all.clone(),
},
Case {
case_id: format!("random6_seed{seed}"),
a: random,
funcs: all.clone(),
},
Case {
case_id: "symmetric4".into(),
a: vec![
vec![4.0, 1.0, 0.0, 0.0],
vec![1.0, 3.0, 1.0, 0.0],
vec![0.0, 1.0, 2.0, 1.0],
vec![0.0, 0.0, 1.0, 1.0],
],
funcs: all,
},
Case {
case_id: "jordan2".into(),
a: vec![vec![2.0, 1.0], vec![0.0, 2.0]],
funcs: vec!["exp".into()],
},
]
}
fn fsci_funm(a: &[Vec<f64>], func: &str) -> Option<(Vec<Vec<f64>>, f64)> {
let options = DecompOptions::default();
let out = match func {
"exp" => funm_with_error(a, |z| z.exp(), options),
"sin" => funm_with_error(a, |z| z.sin(), options),
"cos" => funm_with_error(a, |z| z.cos(), options),
"poly" => funm_with_error(a, |z| z * z * z - z * 2.0 + 1.0, options),
other => unreachable!("FUNCS has no {other}"),
};
out.ok()
}
fn rel_diff(got: &[Vec<f64>], want: &[Vec<f64>]) -> f64 {
assert_eq!(got.len(), want.len(), "shape");
let scale = want
.iter()
.flatten()
.fold(1.0_f64, |acc, v| acc.max(v.abs()));
got.iter()
.flatten()
.zip(want.iter().flatten())
.map(|(g, w)| (g - w).abs())
.fold(0.0_f64, f64::max)
/ scale
}
fn scipy_rows(cases: &[Case]) -> Option<Vec<Row>> {
let script = r#"
import json, sys
import numpy as np
from scipy.linalg import funm, expm, sinm, cosm
FUNCS = {"exp": np.exp, "sin": np.sin, "cos": np.cos, "poly": lambda z: z**3 - 2*z + 1}
REFS = {
"exp": expm, "sin": sinm, "cos": cosm,
"poly": lambda A: A @ A @ A - 2 * A + np.eye(A.shape[0]),
}
rows = []
for case in json.load(sys.stdin):
A = np.asarray(case["a"], dtype=np.float64)
pairs = int(np.sum(np.abs(np.linalg.eigvals(A).imag) > 1e-8) // 2)
for name in case["funcs"]:
F, err = funm(A, FUNCS[name], disp=False)
F = np.asarray(F)
real = None if np.max(np.abs(F.imag)) > 1e-10 else F.real.tolist()
rows.append({
"case_id": case["case_id"], "func": name, "complex_pairs": pairs,
"funm": real, "funm_err": float(err),
"reference": np.real(REFS[name](A)).tolist(),
})
print(json.dumps(rows, allow_nan=False))
"#;
let query = serde_json::to_string(cases).expect("serialize cases");
let mut child = match fsci_conformance::scipy_oracle_command()
.arg("-c")
.arg(script)
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()
{
Ok(c) => c,
Err(e) => {
assert!(
std::env::var(REQUIRE_SCIPY_ENV).is_err(),
"failed to spawn the funm oracle: {e}"
);
eprintln!("skipping funm oracle: python not available ({e})");
return None;
}
};
child
.stdin
.as_mut()
.expect("oracle stdin")
.write_all(query.as_bytes())
.expect("write oracle query");
let output = child.wait_with_output().expect("wait for the funm oracle");
if !output.status.success() {
let stderr = String::from_utf8_lossy(&output.stderr);
assert!(
std::env::var(REQUIRE_SCIPY_ENV).is_err(),
"funm oracle failed: {stderr}"
);
eprintln!("skipping funm oracle: scipy not available\n{stderr}");
return None;
}
Some(serde_json::from_slice(&output.stdout).expect("parse funm oracle JSON"))
}
#[test]
fn diff_linalg_funm_complex() {
let cases = cases();
let expected_rows: usize = cases.iter().map(|c| c.funcs.len()).sum();
let Some(rows) = scipy_rows(&cases) else {
return;
};
assert_eq!(
rows.len(),
expected_rows,
"the oracle must answer every row"
);
let flag = 1000.0 * f64::EPSILON;
let mut compared = 0;
let mut failures = Vec::new();
let mut ledger = CompareLedger::new("diff_linalg_funm_complex", &FUNCS);
for row in &rows {
let case = cases
.iter()
.find(|c| c.case_id == row.case_id)
.expect("case");
let Some((scipy, (fsci, err))) = ledger.both(
&row.func,
&row.case_id,
row.funm.as_ref(),
fsci_funm(&case.a, &row.func),
) else {
continue;
};
let (scipy_flat, fsci_flat) = (scipy.concat(), fsci.concat());
let Some(_) = ledger.slices(
&row.func,
&row.case_id,
Some(scipy_flat.as_slice()),
Some(fsci_flat.as_slice()),
) else {
continue;
};
let Some((_, err)) = ledger.pair(&row.func, &row.case_id, Some(row.funm_err), Some(err))
else {
continue;
};
let vs_funm = rel_diff(&fsci, scipy);
let vs_reference = rel_diff(&fsci, &row.reference);
println!(
"{}/{}: pairs {} | vs SciPy funm {vs_funm:.2e} | vs SciPy dedicated {vs_reference:.2e} | err fsci {err:.2e} SciPy {:.2e}",
row.case_id, row.func, row.complex_pairs, row.funm_err
);
let failed = if row.case_id == "jordan2" {
vs_funm > FUNM_REL_TOL || err < row.funm_err || err <= flag
} else {
vs_funm > FUNM_REL_TOL
|| vs_reference > FUNM_REL_TOL
|| (err > flag) != (row.funm_err > flag)
};
if failed {
failures.push(format!("{}/{}", row.case_id, row.func));
}
ledger.compared(&row.func, &row.case_id, !failed);
compared += 1;
}
let random = rows
.iter()
.find(|r| r.case_id.starts_with("random6"))
.expect("random row");
assert!(
random.complex_pairs >= 2,
"SciPy sees {} complex pairs in the random case",
random.complex_pairs
);
assert_eq!(
compared,
expected_rows,
"every row must be compared; {}",
ledger.verdict(1, false).err().unwrap_or_default()
);
assert!(
failures.is_empty(),
"funm disagrees with SciPy: {failures:?}"
);
let min_per_arm = FUNCS
.iter()
.map(|f| {
cases
.iter()
.filter(|c| c.funcs.iter().any(|g| g.as_str() == *f))
.count()
})
.min()
.unwrap_or(0);
ledger.finish(min_per_arm);
}