#![cfg(feature = "bytecode")]
use echidna::record;
use echidna::BReverse;
use num_traits::Float as _;
#[test]
fn bytecode_hypot_origin_absorbs_infinite_adjoint_and_replays() {
let (mut tape, _) = record(|v: &[BReverse<f64>]| v[0].hypot(v[1]).sqrt(), &[3.0, 4.0]);
let g0 = tape.gradient(&[0.0, 0.0]);
assert_eq!(
g0[0], 0.0,
"zero partial absorbs the Inf adjoint, got {}",
g0[0]
);
assert_eq!(g0[1], 0.0);
let g = tape.gradient(&[3.0, 4.0]);
let expect_x = (3.0 / 5.0) * 0.5 / 5.0_f64.sqrt();
let expect_y = (4.0 / 5.0) * 0.5 / 5.0_f64.sqrt();
assert!(
(g[0] - expect_x).abs() < 1e-12,
"replay gradient x: {}",
g[0]
);
assert!(
(g[1] - expect_y).abs() < 1e-12,
"replay gradient y: {}",
g[1]
);
}
#[test]
fn bytecode_atan2_origin_absorbs_infinite_adjoint() {
let (mut tape, _) = record(|v: &[BReverse<f64>]| v[0].atan2(v[1]).sqrt(), &[3.0, 4.0]);
let g0 = tape.gradient(&[0.0, 0.0]);
assert_eq!(g0[0], 0.0);
assert_eq!(g0[1], 0.0);
}
#[test]
fn regular_zero_partials_absorb_infinite_adjoints() {
let (mut tape, _) = record(
|v: &[BReverse<f64>]| (v[0].max(v[1]) - 4.0).sqrt(),
&[5.0, 1.0],
);
let g = tape.gradient(&[4.0, 1.0]);
assert!(
g[0].is_infinite(),
"winner keeps the live Inf, got {}",
g[0]
);
assert_eq!(g[1], 0.0, "loser's zero partial absorbs the Inf adjoint");
let (mut tape2, _) = record(|v: &[BReverse<f64>]| (v[0] * v[1]).sqrt(), &[2.0, 3.0]);
let g2 = tape2.gradient(&[5.0, 0.0]);
assert_eq!(
g2[0], 0.0,
"zero partial × Inf adjoint must be 0, got {}",
g2[0]
);
assert!(g2[1].is_infinite(), "live partial keeps Inf, got {}", g2[1]);
}
#[test]
fn nan_partials_still_propagate() {
let (mut tape, _) = record(|v: &[BReverse<f64>]| v[0].ln().sqrt(), &[2.0]);
let g = tape.gradient(&[-1.0]);
assert!(
g[0].is_nan(),
"out-of-domain partial must stay NaN, got {}",
g[0]
);
}
#[test]
fn singular_point_convention_table() {
use echidna::Dual;
type Elemental = fn(Dual<f64>) -> Dual<f64>;
let cases: &[(&str, f64, Elemental, f64)] = &[
("sqrt@0", 0.0, |x| x.sqrt(), f64::INFINITY),
("cbrt@0", 0.0, |x| x.cbrt(), f64::INFINITY),
("ln@+0", 0.0, |x| x.ln(), f64::INFINITY),
("ln@-0", -0.0, |x| x.ln(), f64::NEG_INFINITY),
("recip@+0", 0.0, |x| x.recip(), f64::NEG_INFINITY),
("asin@1", 1.0, |x| x.asin(), f64::INFINITY),
("acos@1", 1.0, |x| x.acos(), f64::NEG_INFINITY),
("acosh@1", 1.0, |x| x.acosh(), f64::INFINITY),
("atanh@1", 1.0, |x| x.atanh(), f64::INFINITY),
];
for &(name, x0, f, expect) in cases {
let live = f(Dual::new(x0, 1.0));
assert_eq!(
live.eps, expect,
"{name}: live tangent must be {expect}, got {}",
live.eps
);
let dead = f(Dual::new(x0, 0.0));
assert_eq!(
dead.eps, 0.0,
"{name}: structurally zero tangent must stay 0, got {}",
dead.eps
);
}
}