Files
foxhunt/crates/ml-alpha/tests
jgrusewski 5f10bcde3a feat(ml-alpha): Phase 2A-C+ — device-aggregated gate diag
Adds the specialization-signal diagnostics that Phase 2A-C deferred.
4 new stats are emitted under policy_diagnostic.multi_head_policy.*
when FOXHUNT_USE_MULTI_HEAD_POLICY=1: gate_probs_mean[K],
gate_argmax_mass[K], gate_entropy_mean, per_head_entropy_mean[K].

These are the load-bearing signals for Phase 2A-D's verdict —
without them we'd see a pnl delta but couldn't distinguish "the
mixture genuinely specializes by regime" from "the mixture
accidentally acts as a single-head regularizer".

## ISV slots (25 new)

* 765-772 RL_POLICY_GATE_PROBS_MEAN_BASE (8 slots, MAX_K_HEADS=8 stride)
* 773-780 RL_POLICY_GATE_ARGMAX_MASS_BASE
* 781     RL_POLICY_GATE_ENTROPY_MEAN
* 782-789 RL_POLICY_PER_HEAD_ENTROPY_MEAN_BASE
* RL_SLOTS_END = 790

## Aggregator kernel (multi_head_policy_aggregate_diag.cu)

* Grid = (MAX_K_HEADS + 1, 1, 1) = 9 blocks. Block = (128, 1, 1).
* Per-head blocks (k_block 0..7): if k_block ≥ runtime K, thread 0
  writes 0.0 to its two ISV destinations. Else parallel-sum over B
  in shared memory, tree-reduce by halving stride, thread 0 writes
  gate_probs_mean[k] + per_head_entropy_mean[k].
* Global block (k_block = MAX_K_HEADS): single block computes gate
  entropy + argmax-mass in one pass. Per-thread argmax counter
  array in registers scatters to s_am[MAX_K_HEADS][BLOCK_THREADS]
  shared mem; tree-reduce; thread 0 writes 9 ISV scalars.
* No-atomicAdd: every ISV destination has a single writer thread.
  Tree-reduce via shared memory + __syncthreads(). Deterministic
  fixed-order pairwise sum.

## Trainer wiring

* MultiHeadPolicy::emit_diag_stats(isv_dev_ptr) launches the
  aggregator on the struct stream. Called at all 3 train-path
  forward sites in integrated.rs immediately after mhp.forward()
  (gate_probs and pi_probs_k are forward outputs — must aggregate
  before next step's forward overwrites them).
* build_diag_value emits the 4 fields inside the existing
  multi_head_policy.* block under if self.use_multi_head_policy.
  Flag-off schema unchanged (multi_head_policy key absent).

## Verification (all gates pass)

* 13/13 invariants: 11 pre-existing + 2 new
  - aggregate_diag_gate_probs_mean_correct: uniform 1/K + asymmetric
    ramp case, max abs err < 1e-5
  - aggregate_diag_entropy_pins_top_and_bottom: top (uniform gate
    → entropy = log(K) = 1.0986; uniform pi → per-head = log(11) =
    2.398); bottom (one-hot gate → 0; one-hot pi → 0). Tie-broken-
    low argmax verified.
* FOXHUNT_USE_MULTI_HEAD_POLICY=0 determinism-check.sh --quick:
  exit 0. Flag-off diag.jsonl preserved (multi_head_policy key
  ABSENT, EXPECTED_LEAVES=712 unchanged).
* FOXHUNT_USE_MULTI_HEAD_POLICY=1 determinism-check.sh --quick:
  exit 0. Aggregator is deterministic.
* 200-step flag-on smoke shows the diag working as designed:
  - gate_probs_mean ≈ [0.328, 0.339, 0.333, 0,0,0,0,0] sum ≈ 1.0
  - gate_argmax_mass ≈ [0, 1.0, 0, ...] (Head 1 Long-bias dominant
    at init — matches init bias asymmetry)
  - gate_entropy_mean ≈ 1.0985 (at max log(3) ≈ 1.0986 — gate has
    NOT collapsed)
  - per_head_entropy_mean ≈ [2.36, 2.28, 2.29] vs max 2.398 —
    heads slightly less than uniform, expected at random init
* Pre-commit: 0 atomicAdd, 0 raw memcpy_htod/dtoh, 0 TODO.

## Surprises handled

None — all STOP-on-surprise predictions held:
1. Slot bootstrap zero-init contributes 0² to isv_state checksum,
   so unconditional bootstrap preserves flag-off bit-equality
   without needing a flag-gated block (unlike the 2A-C ISV
   surprise).
2. EXPECTED_LEAVES (712) preserved — flag-off schema unchanged.
3. K consistency (Rust MAX_K_HEADS=8 + CUDA #define) maintained.
4. emit_diag_stats placed after mhp.forward() and before any
   downstream consumer that would overwrite forward outputs.
5. CUDA graph capture: aggregator launch is deterministic (fixed
   grid/block, fixed reductions). Replay single-path. No conflict.

## Linked

* Phase 2A-A: 0b3e40150 (MultiHeadPolicy foundation, inert)
* Phase 2A-B: e22da61cf (backward + aux KL prior)
* Phase 2A-C: 3e36f4a0e (trainer integration behind flag)
* Plan: docs/superpowers/plans/2026-06-03-multi-head-policy-
  implementation.md
* Spec ADDENDUM §R.5 (falsification gates for specialization)

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-06-03 12:15:39 +02:00
..