use echidna::record_multi;
use echidna_optim::{
implicit_tangent, piggyback_adjoint_solve, piggyback_forward_adjoint_solve,
piggyback_tangent_solve, piggyback_tangent_step, piggyback_tangent_step_with_buf,
PiggybackError,
};
#[test]
fn tangent_step_linear() {
let (tape, _) = record_multi(
|v| {
let z = v[0];
let x = v[1];
let half = x / (x + x); vec![half * z + x]
},
&[0.0_f64, 3.0],
);
let (z_new, z_dot_new) = piggyback_tangent_step(&tape, &[0.0], &[3.0], &[0.0], &[1.0], 1);
assert!((z_new[0] - 3.0).abs() < 1e-12, "z_new = {}", z_new[0]);
assert!(
(z_dot_new[0] - 1.0).abs() < 1e-12,
"z_dot_new = {}",
z_dot_new[0]
);
}
#[test]
fn tangent_solve_linear_contraction() {
let (tape, _) = record_multi(
|v| {
let z = v[0];
let x = v[1];
let half = x / (x + x);
vec![half * z + x]
},
&[0.0_f64, 3.0],
);
let result = piggyback_tangent_solve(&tape, &[0.0], &[3.0], &[1.0], 1, 200, 1e-12);
let (z_star, z_dot_star, iters) = result.expect("should converge");
assert!(
(z_star[0] - 6.0).abs() < 1e-10,
"z* = {}, expected 6",
z_star[0]
);
assert!(
(z_dot_star[0] - 2.0).abs() < 1e-8,
"ż* = {}, expected 2",
z_dot_star[0]
);
assert!(iters > 0, "should take at least 1 iteration");
}
#[test]
fn tangent_solve_warm_started_primal() {
let (tape, _) = record_multi(
|v| {
let z = v[0];
let x = v[1];
let half = x / (x + x);
vec![half * z + x]
},
&[0.0_f64, 3.0],
);
let result = piggyback_tangent_solve(&tape, &[6.0], &[3.0], &[1.0], 1, 200, 1e-12);
let (z_star, z_dot_star, _) = result.expect("should converge");
assert!(
(z_star[0] - 6.0).abs() < 1e-10,
"z* = {}, expected 6",
z_star[0]
);
assert!(
(z_dot_star[0] - 2.0).abs() < 1e-8,
"ż* = {}, expected 2 (a primal-only convergence gate returns 1)",
z_dot_star[0]
);
}
#[test]
fn tangent_solve_2d_contraction() {
let (tape, _) = record_multi(
|v| {
let z0 = v[0];
let z1 = v[1];
let x0 = v[2];
let x1 = v[3];
let one = x0 / x0;
let pt4 = (one + one) / (one + one + one + one + one); let pt3 =
(one + one + one) / (one + one + one + one + one + one + one + one + one + one); vec![pt4 * z0 + x0, pt3 * z1 + x1]
},
&[0.0_f64, 0.0, 1.2, 2.1],
);
let x = [1.2, 2.1];
let result = piggyback_tangent_solve(&tape, &[0.0, 0.0], &x, &[1.0, 0.0], 2, 200, 1e-12);
let (z_star, z_dot, _) = result.expect("should converge");
let expected_z0 = 1.2 / 0.6;
let expected_z1 = 2.1 / 0.7;
assert!(
(z_star[0] - expected_z0).abs() < 1e-9,
"z0* = {}, expected {}",
z_star[0],
expected_z0
);
assert!(
(z_star[1] - expected_z1).abs() < 1e-9,
"z1* = {}, expected {}",
z_star[1],
expected_z1
);
assert!(
(z_dot[0] - 1.0 / 0.6).abs() < 1e-7,
"dz0*/dx0 = {}, expected {}",
z_dot[0],
1.0 / 0.6
);
assert!(z_dot[1].abs() < 1e-7, "dz1*/dx0 = {}, expected 0", z_dot[1]);
let result2 = piggyback_tangent_solve(&tape, &[0.0, 0.0], &x, &[0.0, 1.0], 2, 200, 1e-12);
let (_, z_dot2, _) = result2.expect("should converge");
assert!(
z_dot2[0].abs() < 1e-7,
"dz0*/dx1 = {}, expected 0",
z_dot2[0]
);
assert!(
(z_dot2[1] - 1.0 / 0.7).abs() < 1e-7,
"dz1*/dx1 = {}, expected {}",
z_dot2[1],
1.0 / 0.7
);
}
#[test]
fn adjoint_solve_linear_contraction() {
let (mut tape, _) = record_multi(
|v| {
let z = v[0];
let x = v[1];
let half = x / (x + x);
vec![half * z + x]
},
&[6.0_f64, 3.0],
);
let result = piggyback_adjoint_solve(&mut tape, &[6.0], &[3.0], &[1.0], 1, 200, 1e-12);
let (x_bar, iters) = result.expect("should converge");
assert!(
(x_bar[0] - 2.0).abs() < 1e-8,
"x̄ = {}, expected 2",
x_bar[0]
);
assert!(iters > 0);
}
#[test]
fn adjoint_vs_tangent_transpose() {
let make_tape = || {
record_multi(
|v| {
let z0 = v[0];
let z1 = v[1];
let x0 = v[2];
let x1 = v[3];
let one = x0 / x0;
let pt4 = (one + one) / (one + one + one + one + one);
let pt3 =
(one + one + one) / (one + one + one + one + one + one + one + one + one + one);
vec![pt4 * z0 + x0, pt3 * z1 + x1]
},
&[0.0_f64, 0.0, 1.2, 2.1],
)
};
let x = [1.2, 2.1];
let z_star = [1.2 / 0.6, 2.1 / 0.7];
let (tape, _) = make_tape();
let (_, col0, _) = piggyback_tangent_solve(&tape, &[0.0, 0.0], &x, &[1.0, 0.0], 2, 200, 1e-12)
.expect("should converge");
let (_, col1, _) = piggyback_tangent_solve(&tape, &[0.0, 0.0], &x, &[0.0, 1.0], 2, 200, 1e-12)
.expect("should converge");
let (mut tape_a, _) = make_tape();
let (row0, _) = piggyback_adjoint_solve(&mut tape_a, &z_star, &x, &[1.0, 0.0], 2, 200, 1e-12)
.expect("should converge");
let (row1, _) = piggyback_adjoint_solve(&mut tape_a, &z_star, &x, &[0.0, 1.0], 2, 200, 1e-12)
.expect("should converge");
assert!(
(row0[0] - col0[0]).abs() < 1e-7,
"row0[0]={}, col0[0]={}",
row0[0],
col0[0]
);
assert!(
(row0[1] - col1[0]).abs() < 1e-7,
"row0[1]={}, col1[0]={}",
row0[1],
col1[0]
);
assert!(
(row1[0] - col0[1]).abs() < 1e-7,
"row1[0]={}, col0[1]={}",
row1[0],
col0[1]
);
assert!(
(row1[1] - col1[1]).abs() < 1e-7,
"row1[1]={}, col1[1]={}",
row1[1],
col1[1]
);
}
#[test]
fn cross_validate_with_ift() {
let (step_tape, _) = record_multi(
|v| {
let z = v[0];
let x = v[1];
let half = x / (x + x);
vec![half * z + x]
},
&[6.0_f64, 3.0],
);
let (mut residual_tape, _) = record_multi(
|v| {
let z = v[0];
let x = v[1];
let half = x / (x + x);
vec![z - (half * z + x)] },
&[6.0_f64, 3.0],
);
let (_, z_dot_pb, _) =
piggyback_tangent_solve(&step_tape, &[0.0], &[3.0], &[1.0], 1, 200, 1e-12)
.expect("piggyback should converge");
let z_dot_ift =
implicit_tangent(&mut residual_tape, &[6.0], &[3.0], &[1.0], 1).expect("IFT should work");
assert!(
(z_dot_pb[0] - z_dot_ift[0]).abs() < 1e-7,
"piggyback ż*={}, IFT ż*={}",
z_dot_pb[0],
z_dot_ift[0]
);
}
#[test]
fn tangent_step_buffer_reuse() {
let (tape, _) = record_multi(
|v| {
let z = v[0];
let x = v[1];
let half = x / (x + x);
vec![half * z + x]
},
&[1.0_f64, 3.0],
);
let mut buf = Vec::new();
let (z1, zd1) =
piggyback_tangent_step_with_buf(&tape, &[1.0], &[3.0], &[0.5], &[1.0], 1, &mut buf);
let (z2, zd2) =
piggyback_tangent_step_with_buf(&tape, &[1.0], &[3.0], &[0.5], &[1.0], 1, &mut buf);
assert_eq!(z1[0], z2[0]);
assert_eq!(zd1[0], zd2[0]);
}
#[test]
fn adjoint_non_convergent() {
let (mut tape, _) = record_multi(
|v| {
let z = v[0];
let x = v[1];
let two = (x + x) / x; vec![two * z + x]
},
&[1.0_f64, 1.0],
);
let err = piggyback_adjoint_solve(&mut tape, &[1.0], &[1.0], &[1.0], 1, 100, 1e-12)
.expect_err("should not converge for non-contraction");
match err {
PiggybackError::IterationsExhaustedAdjoint {
iteration,
lam_norm,
} => {
assert_eq!(iteration, 100, "iteration must equal max_iter");
assert!(
lam_norm.is_finite(),
"lam_norm must be finite (got {lam_norm})"
);
}
other => panic!("expected IterationsExhaustedAdjoint, got {other:?}"),
}
}
#[test]
fn forward_adjoint_solve_linear() {
let (mut tape, _) = record_multi(
|v| {
let z = v[0];
let x = v[1];
let half = x / (x + x);
vec![half * z + x]
},
&[0.0_f64, 3.0],
);
let result = piggyback_forward_adjoint_solve(&mut tape, &[0.0], &[3.0], &[1.0], 1, 200, 1e-12);
let (z_star, x_bar, iters) = result.expect("should converge");
assert!(
(z_star[0] - 6.0).abs() < 1e-10,
"z* = {}, expected 6",
z_star[0]
);
assert!(
(x_bar[0] - 2.0).abs() < 1e-8,
"x̄ = {}, expected 2",
x_bar[0]
);
assert!(iters > 0);
}
#[test]
fn forward_adjoint_solve_2d() {
let (mut tape, _) = record_multi(
|v| {
let z0 = v[0];
let z1 = v[1];
let x0 = v[2];
let x1 = v[3];
let one = x0 / x0;
let pt4 = (one + one) / (one + one + one + one + one);
let pt3 =
(one + one + one) / (one + one + one + one + one + one + one + one + one + one);
vec![pt4 * z0 + x0, pt3 * z1 + x1]
},
&[0.0_f64, 0.0, 1.2, 2.1],
);
let result = piggyback_forward_adjoint_solve(
&mut tape,
&[0.0, 0.0],
&[1.2, 2.1],
&[1.0, 0.0],
2,
200,
1e-12,
);
let (z_star, x_bar, _) = result.expect("should converge");
assert!(
(z_star[0] - 1.2 / 0.6).abs() < 1e-9,
"z0* = {}, expected {}",
z_star[0],
1.2 / 0.6
);
assert!(
(z_star[1] - 2.1 / 0.7).abs() < 1e-9,
"z1* = {}, expected {}",
z_star[1],
2.1 / 0.7
);
assert!(
(x_bar[0] - 1.0 / 0.6).abs() < 1e-7,
"x̄[0] = {}, expected {}",
x_bar[0],
1.0 / 0.6
);
assert!(x_bar[1].abs() < 1e-7, "x̄[1] = {}, expected 0", x_bar[1]);
}
#[test]
fn forward_adjoint_matches_sequential() {
let make_tape = || {
record_multi(
|v| {
let z0 = v[0];
let z1 = v[1];
let x0 = v[2];
let x1 = v[3];
let one = x0 / x0;
let pt4 = (one + one) / (one + one + one + one + one);
let pt3 =
(one + one + one) / (one + one + one + one + one + one + one + one + one + one);
vec![pt4 * z0 + x0, pt3 * z1 + x1]
},
&[0.0_f64, 0.0, 1.2, 2.1],
)
};
let x = [1.2, 2.1];
let z_bar = [1.0, 0.5];
let (mut tape_seq, _) = make_tape();
let (z_star_seq, _, _) =
piggyback_tangent_solve(&tape_seq, &[0.0, 0.0], &x, &[1.0, 0.0], 2, 200, 1e-12)
.expect("tangent should converge");
let (x_bar_seq, _) =
piggyback_adjoint_solve(&mut tape_seq, &z_star_seq, &x, &z_bar, 2, 200, 1e-12)
.expect("adjoint should converge");
let (mut tape_int, _) = make_tape();
let (z_star_int, x_bar_int, _) =
piggyback_forward_adjoint_solve(&mut tape_int, &[0.0, 0.0], &x, &z_bar, 2, 200, 1e-12)
.expect("interleaved should converge");
for i in 0..2 {
assert!(
(z_star_int[i] - z_star_seq[i]).abs() < 1e-9,
"z*[{}]: interleaved={}, sequential={}",
i,
z_star_int[i],
z_star_seq[i]
);
}
for j in 0..2 {
assert!(
(x_bar_int[j] - x_bar_seq[j]).abs() < 1e-7,
"x̄[{}]: interleaved={}, sequential={}",
j,
x_bar_int[j],
x_bar_seq[j]
);
}
}
#[test]
fn forward_adjoint_non_convergent() {
let (mut tape, _) = record_multi(
|v| {
let z = v[0];
let x = v[1];
let two = (x + x) / x;
vec![two * z + x]
},
&[1.0_f64, 1.0],
);
let err = piggyback_forward_adjoint_solve(&mut tape, &[0.0], &[1.0], &[1.0], 1, 100, 1e-12)
.expect_err("should not converge for non-contraction");
match err {
PiggybackError::IterationsExhaustedForwardAdjoint {
iteration,
z_norm,
lam_norm,
} => {
assert_eq!(iteration, 100, "iteration must equal max_iter");
assert!(
z_norm.is_finite() && lam_norm.is_finite(),
"both norms must be finite (got z_norm = {z_norm}, lam_norm = {lam_norm})"
);
}
other => panic!("expected IterationsExhaustedForwardAdjoint, got {other:?}"),
}
}
#[test]
fn regression_6_forward_adjoint_vs_adjoint_consistency() {
let make_tape = || {
record_multi(
|v| {
let z0 = v[0];
let z1 = v[1];
let x0 = v[2];
let x1 = v[3];
let one = x0 / x0;
let pt4 = (one + one) / (one + one + one + one + one);
let pt3 =
(one + one + one) / (one + one + one + one + one + one + one + one + one + one);
vec![pt4 * z0 + x0, pt3 * z1 + x1]
},
&[0.0_f64, 0.0, 1.2, 2.1],
)
};
let x = [1.2, 2.1];
let z_bar = [1.0, 0.5];
let (mut tape_seq, _) = make_tape();
let (z_star_seq, _, _) =
piggyback_tangent_solve(&tape_seq, &[0.0, 0.0], &x, &[1.0, 0.0], 2, 200, 1e-12)
.expect("tangent should converge");
let (x_bar_adj, _) =
piggyback_adjoint_solve(&mut tape_seq, &z_star_seq, &x, &z_bar, 2, 200, 1e-12)
.expect("adjoint should converge");
let (mut tape_fa, _) = make_tape();
let (z_star_fa, x_bar_fa, _) =
piggyback_forward_adjoint_solve(&mut tape_fa, &[0.0, 0.0], &x, &z_bar, 2, 200, 1e-12)
.expect("forward-adjoint should converge");
for i in 0..2 {
assert!(
(z_star_fa[i] - z_star_seq[i]).abs() < 1e-8,
"z*[{}]: forward-adjoint={}, sequential={}",
i,
z_star_fa[i],
z_star_seq[i]
);
}
for j in 0..2 {
assert!(
(x_bar_fa[j] - x_bar_adj[j]).abs() < 1e-7,
"x_bar[{}]: forward-adjoint={}, adjoint={}",
j,
x_bar_fa[j],
x_bar_adj[j]
);
}
}