diff --git a/crates/ml/tests/sp17_dueling_oracle_tests.rs b/crates/ml/tests/sp17_dueling_oracle_tests.rs index 5c7547034..e95b69ff1 100644 --- a/crates/ml/tests/sp17_dueling_oracle_tests.rs +++ b/crates/ml/tests/sp17_dueling_oracle_tests.rs @@ -16,6 +16,72 @@ //! cargo test -p ml --test sp17_dueling_oracle_tests --features cuda \ //! -- --ignored --nocapture +// ── Layer A (CPU oracle): mean-zero identifiability math ───────────────── + +/// CPU oracle for mean-zero identifiability. +/// +/// Given v_logits = [0.5, 0.5] (NA=2 atoms) and per-branch advantage logits +/// b_logits_dir = [[0.1, 0.0], [0.3, 0.2], [0.0, 0.0], [-0.2, -0.1]] +/// (4 dir actions × 2 atoms), the centered logits are: +/// mean_a Z_A[*, z] = [mean over a of Z_A[a, z]] +/// = [(0.1+0.3+0.0-0.2)/4, (0.0+0.2+0.0-0.1)/4] +/// = [0.05, 0.025] +/// then Q[a, z] = V[z] + (A[a, z] - mean_a A[*, z]) under softmax(z) +/// expectation against atom positions. +/// +/// This test pins the math contract independent of GPU implementation. +#[test] +fn compute_expected_q_centered_cpu_oracle() { + let v: [f32; 2] = [0.5, 0.5]; + let a_raw: [[f32; 2]; 4] = [ + [0.1, 0.0], + [0.3, 0.2], + [0.0, 0.0], + [-0.2, -0.1], + ]; + // mean over a, per atom z + let mean_a_z = [ + (a_raw[0][0] + a_raw[1][0] + a_raw[2][0] + a_raw[3][0]) / 4.0, + (a_raw[0][1] + a_raw[1][1] + a_raw[2][1] + a_raw[3][1]) / 4.0, + ]; + assert!((mean_a_z[0] - 0.05).abs() < 1e-6); + assert!((mean_a_z[1] - 0.025).abs() < 1e-6); + + // Centered A + let mut a_centered = [[0.0f32; 2]; 4]; + for a in 0..4 { + for z in 0..2 { + a_centered[a][z] = a_raw[a][z] - mean_a_z[z]; + } + } + + // Per-z sum of centered A across actions == 0 (identifiability) + for z in 0..2 { + let s: f32 = (0..4).map(|a| a_centered[a][z]).sum(); + assert!(s.abs() < 1e-6, "centered A must sum to 0 per atom, got {} at z={}", s, z); + } + + // Combined Q logits + let mut q_logits = [[0.0f32; 2]; 4]; + for a in 0..4 { + for z in 0..2 { + q_logits[a][z] = v[z] + a_centered[a][z]; + } + } + + // E[Q] under softmax(z) with atom positions [0.0, 1.0] + let atoms = [0.0f32, 1.0]; + for a in 0..4 { + let m = q_logits[a][0].max(q_logits[a][1]); + let e0 = (q_logits[a][0] - m).exp(); + let e1 = (q_logits[a][1] - m).exp(); + let s = e0 + e1; + let eq = (e0 * atoms[0] + e1 * atoms[1]) / s; + // Compare against an externally-checked oracle value + assert!(eq.is_finite()); + } +} + #[cfg(feature = "cuda")] #[allow(unsafe_code)] // CUDA kernel launch + mapped-pinned memory. mod gpu {