#[inline(always)]
pub fn faa_top2(m: [f64; 2], q: &[f64; 4]) -> f64 {
m[1] * q[1] * q[2] + m[0] * q[3]
}
#[inline(always)]
pub fn faa_top3(m: [f64; 3], q: &[f64; 8]) -> f64 {
let (a, b, u) = (1usize, 2, 4);
let p = q[a] * q[b];
let p_u = q[a | u] * q[b] + q[a] * q[b | u];
m[2] * (q[u] * p) + m[1] * (p_u + q[u] * q[a | b]) + m[0] * q[a | b | u]
}
#[inline(always)]
pub fn faa_top4(m: [f64; 4], q: &[f64; 16]) -> f64 {
let (a, b, u, v) = (1usize, 2, 4, 8);
let p = q[a] * q[b];
let p_u = q[a | u] * q[b] + q[a] * q[b | u];
let p_v = q[a | v] * q[b] + q[a] * q[b | v];
let p_uv = q[a | u | v] * q[b] + q[a | u] * q[b | v] + q[a | v] * q[b | u] + q[a] * q[b | u | v];
let uv = q[u] * q[v];
m[3] * (uv * p)
+ m[2] * (q[u | v] * p + q[u] * p_v + q[v] * p_u + uv * q[a | b])
+ m[1] * (p_uv + q[u | v] * q[a | b] + q[u] * q[a | b | v] + q[v] * q[a | b | u])
+ m[0] * q[a | b | u | v]
}
#[inline(always)]
pub fn faa_bundle3<const OUTPUTS: usize>(m: [f64; 3], q: &[[f64; 8]; OUTPUTS]) -> [f64; OUTPUTS] {
std::array::from_fn(|output| faa_top3(m, &q[output]))
}
#[inline(always)]
pub fn curve_wiggle_bundle3(
m: [f64; 3],
q_u: f64,
a: f64,
b: f64,
a_u: f64,
b_u: f64,
xi: f64,
) -> [f64; 6] {
faa_bundle3(
m,
&[
[0.0, a, a, b, q_u, a_u, a_u, b_u],
[0.0, a, 1.0, 0.0, q_u, a_u, 0.0, 0.0],
[0.0, a, 0.0, 1.0, q_u, a_u, xi, 0.0],
[0.0, a, 0.0, 0.0, q_u, a_u, 0.0, xi],
[0.0, 1.0, 1.0, 0.0, q_u, 0.0, 0.0, 0.0],
[0.0, 0.0, 1.0, 0.0, q_u, xi, 0.0, 0.0],
],
)
}
#[inline(always)]
pub fn curve_wiggle_bundle4(
m: [f64; 4],
q: [f64; 3],
a_jet: [f64; 4],
b_jet: [f64; 4],
xi: [f64; 2],
) -> [f64; 9] {
let [m1, m2, m3, m4] = m;
let [q_u, q_v, q_uv] = q;
let [a, a_u, a_v, a_uv] = a_jet;
let [b, b_u, b_v, b_uv] = b_jet;
let [xi_u, xi_v] = xi;
let xi_uv = xi_u * xi_v;
let ww_bb = m4 * q_u * q_v + m3 * q_uv;
let ww_db = m3 * (xi_u * q_v + xi_v * q_u);
let ww_ddb = m2 * xi_uv;
let ww_dd = 2.0 * ww_ddb;
let eta_w_b = ww_bb * a + m3 * (a_u * q_v + a_v * q_u) + m2 * a_uv;
let eta_w_d1 = ww_db * a + m3 * q_u * q_v + m2 * (q_uv + a_u * xi_v + a_v * xi_u);
let eta_w_d2 = ww_ddb * a + m2 * (q_u * xi_v + q_v * xi_u);
let eta_w_d3 = m1 * xi_uv;
let eta_eta = eta_w_b * a
+ m3 * (a * (q_u * a_v + q_v * a_u) + q_u * q_v * b)
+ m2 * (a * a_uv + 2.0 * a_u * a_v + q_uv * b + q_u * b_v + q_v * b_u)
+ m1 * b_uv;
[
eta_eta, eta_w_b, eta_w_d1, eta_w_d2, eta_w_d3, ww_bb, ww_db, ww_ddb, ww_dd,
]
}
#[cfg(test)]
mod oracle_tests {
use super::*;
use crate::jet_algebra::faa_di_bruno;
use crate::paired_timing::{SpeedGate, batched, paired_interleaved};
use std::hint::black_box;
fn stream(seed: u64) -> impl FnMut() -> f64 {
let mut s = seed;
move || {
s = s
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((s >> 11) as f64 / (1u64 << 53) as f64) * 2.0 - 1.0
}
}
fn walker_top(n: usize, derivs: &[f64], q: &[f64]) -> f64 {
let positions: Vec<usize> = (0..n).collect();
faa_di_bruno(&positions, derivs, |block| {
let mask: usize = block.iter().fold(0usize, |acc, &p| acc | (1 << p));
q[mask]
})
}
#[test]
fn faa_top2_matches_runtime_walker() {
let mut next = stream(0x2);
for _ in 0..500 {
let m = [next(), next()];
let mut q = [0.0; 4];
for (mask, qm) in q.iter_mut().enumerate() {
if mask != 0 {
*qm = next();
}
}
let derivs = [0.0, m[0], m[1]];
let got = faa_top2(m, &q);
let want = walker_top(2, &derivs, &q);
assert!(
(got - want).abs() <= 1e-12 * want.abs().max(1.0),
"faa_top2 {got:+.17e} vs walker {want:+.17e}"
);
}
}
#[test]
fn faa_top3_matches_runtime_walker() {
let mut next = stream(0x3);
for _ in 0..500 {
let m = [next(), next(), next()];
let mut q = [0.0; 8];
for (mask, qm) in q.iter_mut().enumerate() {
if mask != 0 {
*qm = next();
}
}
let derivs = [0.0, m[0], m[1], m[2]];
let got = faa_top3(m, &q);
let want = walker_top(3, &derivs, &q);
assert!(
(got - want).abs() <= 1e-12 * want.abs().max(1.0),
"faa_top3 {got:+.17e} vs walker {want:+.17e}"
);
}
}
#[test]
fn faa_top4_matches_runtime_walker() {
let mut next = stream(0x4);
for _ in 0..500 {
let m = [next(), next(), next(), next()];
let mut q = [0.0; 16];
for (mask, qm) in q.iter_mut().enumerate() {
if mask != 0 {
*qm = next();
}
}
let derivs = [0.0, m[0], m[1], m[2], m[3]];
let got = faa_top4(m, &q);
let want = walker_top(4, &derivs, &q);
assert!(
(got - want).abs() <= 1e-12 * want.abs().max(1.0),
"faa_top4 {got:+.17e} vs walker {want:+.17e}"
);
}
}
#[inline(never)]
fn compiled_bundle3(x: [f64; 9]) -> [f64; 6] {
curve_wiggle_bundle3([x[0], x[1], x[2]], x[3], x[4], x[5], x[6], x[7], x[8])
}
#[inline(never)]
fn strongest_hand_bundle3(x: [f64; 9]) -> [f64; 6] {
let [m1, m2, m3, q_u, a, b, a_u, b_u, xi] = x;
let ww_bb = m3 * q_u;
let ww_db = m2 * xi;
let eta_w_b = ww_bb * a + m2 * a_u;
[
eta_w_b * a + m2 * (a * a_u + q_u * b) + m1 * b_u,
eta_w_b,
m2 * (a * xi + q_u),
m1 * xi,
ww_bb,
ww_db,
]
}
#[inline(never)]
fn compiled_bundle4(x: [f64; 17]) -> [f64; 9] {
curve_wiggle_bundle4(
[x[0], x[1], x[2], x[3]],
[x[4], x[5], x[6]],
[x[7], x[9], x[10], x[11]],
[x[8], x[12], x[13], x[14]],
[x[15], x[16]],
)
}
fn canonical_bundle4(x: [f64; 17]) -> [f64; 9] {
let [
m1,
m2,
m3,
m4,
q_u,
q_v,
q_uv,
a,
b,
a_u,
a_v,
a_uv,
b_u,
b_v,
b_uv,
xi_u,
xi_v,
] = x;
let xi_uv = xi_u * xi_v;
let q = [
[
0.0, a, a, b, q_u, a_u, a_u, b_u, q_v, a_v, a_v, b_v, q_uv, a_uv, a_uv, b_uv,
],
[
0.0, a, 1.0, 0.0, q_u, a_u, 0.0, 0.0, q_v, a_v, 0.0, 0.0, q_uv, a_uv, 0.0, 0.0,
],
[
0.0, a, 0.0, 1.0, q_u, a_u, xi_u, 0.0, q_v, a_v, xi_v, 0.0, q_uv, 0.0, 0.0, 0.0,
],
[
0.0, a, 0.0, 0.0, q_u, 0.0, 0.0, xi_u, q_v, 0.0, 0.0, xi_v, 0.0, 0.0, xi_uv, 0.0,
],
[
0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, xi_uv,
],
[
0.0, 1.0, 1.0, 0.0, q_u, 0.0, 0.0, 0.0, q_v, 0.0, 0.0, 0.0, q_uv, 0.0, 0.0, 0.0,
],
[
0.0, 0.0, 1.0, 0.0, q_u, xi_u, 0.0, 0.0, q_v, xi_v, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0,
],
[
0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, xi_uv, 0.0, 0.0,
],
[
0.0, 0.0, 0.0, 0.0, 0.0, xi_u, xi_u, 0.0, 0.0, xi_v, xi_v, 0.0, 0.0, 0.0, 0.0, 0.0,
],
];
std::array::from_fn(|output| faa_top4([m1, m2, m3, m4], &q[output]))
}
#[inline(never)]
fn strongest_hand_bundle4(x: [f64; 17]) -> [f64; 9] {
let [
m1,
m2,
m3,
m4,
q_u,
q_v,
q_uv,
a,
b,
a_u,
a_v,
a_uv,
b_u,
b_v,
b_uv,
xi_u,
xi_v,
] = x;
let eta_eta = m4 * q_u * q_v * a * a
+ m3 * (q_uv * a * a
+ q_u * (a_v * a + a * a_v)
+ q_v * (a_u * a + a * a_u)
+ q_u * q_v * b)
+ m2 * (a_uv * a + a_u * a_v + a_v * a_u + a * a_uv + q_uv * b + q_u * b_v + q_v * b_u)
+ m1 * b_uv;
let eta_w_b = m4 * q_u * q_v * a + m3 * (q_uv * a + q_u * a_v + q_v * a_u) + m2 * a_uv;
let eta_w_d1 = (m3 * q_u * a + m2 * a_u) * xi_v
+ (m3 * q_v * a + m2 * a_v) * xi_u
+ m3 * q_u * q_v
+ m2 * q_uv;
let xi_uv = xi_u * xi_v;
let eta_w_d2 = m2 * a * xi_uv + m2 * q_u * xi_v + m2 * q_v * xi_u;
let eta_w_d3 = m1 * xi_uv;
let ww_bb = m4 * q_u * q_v + m3 * q_uv;
let ww_db = m3 * q_v * xi_u + m3 * q_u * xi_v;
let ww_ddb = m2 * xi_uv;
let ww_dd = 2.0 * m2 * xi_uv;
[
eta_eta, eta_w_b, eta_w_d1, eta_w_d2, eta_w_d3, ww_bb, ww_db, ww_ddb, ww_dd,
]
}
fn assert_close<const N: usize>(actual: [f64; N], expected: [f64; N]) {
for channel in 0..N {
let band = 2e-12 * actual[channel].abs().max(expected[channel].abs()).max(1.0);
assert!(
(actual[channel] - expected[channel]).abs() <= band,
"bundle channel {channel}: compiled={} hand={} band={band}",
actual[channel],
expected[channel],
);
}
}
#[test]
fn higher_order_curve_wiggle_bundles_beat_strongest_hand_932() {
let order3 = [0.31, 0.72, -0.41, 0.83, 1.13, -0.27, 0.19, -0.08, 0.47];
let order4 = [
0.31, 0.72, -0.41, 0.26, 0.83, -0.54, 0.17, 1.13, -0.27, 0.19, -0.12, 0.07, -0.08,
0.05, -0.03, 0.47, -0.38,
];
assert_close(compiled_bundle3(order3), strongest_hand_bundle3(order3));
assert_close(compiled_bundle4(order4), strongest_hand_bundle4(order4));
assert_close(compiled_bundle4(order4), canonical_bundle4(order4));
let mut next = stream(0x9324_BA7C);
for _ in 0..500 {
let random3 = std::array::from_fn(|_| next());
let random4 = std::array::from_fn(|_| next());
assert_close(compiled_bundle3(random3), strongest_hand_bundle3(random3));
assert_close(compiled_bundle4(random4), canonical_bundle4(random4));
assert_close(compiled_bundle4(random4), strongest_hand_bundle4(random4));
}
if cfg!(debug_assertions) {
return;
}
let mut gate = SpeedGate::open("CURVE-WIGGLE-BUNDLE-932");
for (order, timing) in [
(
3,
paired_interleaved(
15,
3_000,
0x5153_9320_0C03,
batched(64, |nudge| {
let mut x = order3;
x[0] += nudge;
compiled_bundle3(black_box(x)).into_iter().sum::<f64>()
}),
batched(64, |nudge| {
let mut x = order3;
x[0] += nudge;
strongest_hand_bundle3(black_box(x)).into_iter().sum::<f64>()
}),
),
),
(
4,
paired_interleaved(
15,
1_500,
0x5153_9320_0C04,
batched(64, |nudge| {
let mut x = order4;
x[0] += nudge;
compiled_bundle4(black_box(x)).into_iter().sum::<f64>()
}),
batched(64, |nudge| {
let mut x = order4;
x[0] += nudge;
strongest_hand_bundle4(black_box(x)).into_iter().sum::<f64>()
}),
),
),
] {
gate.faster(&format!("order={order}"), &timing, "compiled", "strongest_hand");
}
gate.finish();
}
}