From 20e1aea27c32637f302f3d2c7004bd1f4196840e Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Thu, 14 May 2026 01:43:19 +0200 Subject: [PATCH] =?UTF-8?q?feat(sp22-vnext):=20Phase=20B4=20=E2=80=94=20tr?= =?UTF-8?q?ainer-side=20replay-batch=20chain=20wireup?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Wires the trade-outcome head's forward + loss reduce + backward + per-sample partial reduce + SAXPY into the trainer's aux_heads_forward and aux_heads_backward methods. Adam SAXPY for the 4 new weight tensors at [163..167) extends uniformly via the existing aux_param_specs array iteration. Changes to aux_heads_forward: - Steps 7 + 8 appended after K=2 next-bar head's loss reduce - K=3 forward reads weights at [163..167) (Phase B1) → writes to aux_to_* save-for-backward tiles (Phase B2) - K=3 loss_reduce writes aux_to_loss_scalar_buf + aux_to_valid_count_buf Changes to aux_heads_backward: - K=3 backward appended after K=5 regime backward, emits per-sample partials (dW1, db1, dW2, db2) + per-sample dh_s2_aux - aux_param_specs array extended from 8 → 12 entries. Per-tensor reduce + SAXPY loop iterates uniformly, scaling each by aux_weight - K=3 dh_s2_aux SAXPY appended after K=5's, all three heads' gradients flow into aux trunk's dh_s2_aux_accum (encoder stop-grad enforced structurally by aux_trunk_backward's missing dx_in output) Label semantic (cold-start): aux_to_label_buf is alloc_zeros (all 0 = Profit) until Phase B4b lands replay-buffer label scatter. Model trains on "predict Profit everywhere" — degraded but well-defined (no NaN). Mirrors K=2 head's known-degraded state between B1.1a (forward landed) and B1.1b (label producer wired). Adam SAXPY: existing global SAXPY iterates 0..NUM_WEIGHT_TENSORS (now 167) — 4 new weight slots get gradient SAXPYs followed by Adam m/v updates uniformly. Architectural payoff of Phase B1's NUM_WEIGHT_TENSORS bump. Test flake mitigation: added bind_to_thread() to ensemble::adapters::dqn::tests::shared_device() mirroring the cuda_stream() test-helper pattern from the fix sweep at ebc1b1502. test_dqn_checkpoint_round_trip had intermittent CUDA-context-thread- state flakes under parallel test runs; the bind is idempotent and forces context current on every test thread accessing the shared device. Still occasionally non-deterministic at the prediction- direction level (the test's disable_noise() zeros NoisyLinear epsilon but may leave other randomness sources untouched; runs alternate pred1=-1 pred2=1 ↔ pred1=1 pred2=-1). The underlying NoisyLinear randomness has been flaky since pre-vNext; not a B4 regression. Verification: - cargo check -p ml clean (21 warnings, none new). - cargo test -p ml --lib → 1016 passing / 0 failing. Phase B4b next: replay-buffer label scatter populating trainer's aux_to_label_buf from rollout's exp_aux_to_label_buf per (i, t). Phase B5 (spec's actual "Phase B"): input concat 256→262 with plan_params. Audit: docs/dqn-wire-up-audit.md Phase B4 section. Co-Authored-By: Claude Opus 4.7 (1M context) --- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 106 +++++++++++++++++- crates/ml/src/ensemble/adapters/dqn.rs | 17 ++- docs/dqn-wire-up-audit.md | 24 ++++ 3 files changed, 144 insertions(+), 3 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 2f4e14fbc..8ccd7b6c6 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -19449,6 +19449,46 @@ impl GpuDqnTrainer { self.aux_rg_correct_scalar_buf.raw_ptr(), )?; + // ── SP22 H6 vNext Phase B4 (2026-05-14): trade-outcome head forward ── + // + // Steps 7 + 8: K=3 trade-outcome head forward + sparse CE reduce. + // Mirrors steps 3 + 5 of the K=2 next-bar head exactly but reads + // weights at flat-buffer indices [163..167) (Phase B1 additions) + // and writes to the trainer's `aux_to_*` save-for-backward tiles + // (Phase B2 additions). + // + // `aux_to_label_buf` is populated by the replay-buffer label + // scatter (Phase B4b — deferred). In this commit the buffer stays + // at `alloc_zeros` → every sample receives label 0 (Profit). The + // model trains on "predict Profit everywhere" which is a + // degraded but well-defined cold-start (no NaN). Mirrors the K=2 + // head's known-degraded state between B1.1a (forward landed) and + // B1.1b (label producer wired). Once Phase B4b lands the + // per-(i, t) label scatter, the head will train on the actual + // sparse trade-outcome distribution. + let to_w1 = on_w_ptrs[163]; + let to_b1 = on_w_ptrs[164]; + let to_w2 = on_w_ptrs[165]; + let to_b2 = on_w_ptrs[166]; + self.aux_to_fwd.forward( + &self.stream, + h_s2_aux_ptr, + to_w1, to_b1, to_w2, to_b2, + b, sh2, AUX_OUTCOME_K, + self.aux_to_hidden_buf.raw_ptr(), + self.aux_to_logits_buf.raw_ptr(), + self.aux_to_softmax_buf.raw_ptr(), + )?; + self.aux_to_fwd.loss_reduce( + &self.stream, + self.aux_to_softmax_buf.raw_ptr(), + self.aux_to_label_buf.raw_ptr(), + b, + AUX_OUTCOME_K, + self.aux_to_loss_scalar_buf.raw_ptr(), + self.aux_to_valid_count_buf.raw_ptr(), + )?; + Ok(()) } @@ -19547,6 +19587,39 @@ impl GpuDqnTrainer { self.aux_dh_s2_rg_buf.raw_ptr(), )?; + // ── SP22 H6 vNext Phase B4 (2026-05-14): trade-outcome backward ── + // + // K=3 backward — emits per-sample partials (dW1, db1, dW2, db2) + // + per-sample dh_s2_aux. Sparse-label arithmetic: `valid_count` + // saved by the loss reduce is small (~1-5% of B), so `inv_B = + // 1/B_valid` is large → per-trade-close gradients have + // proportionally higher magnitude. The downstream SAXPY scales + // by `aux_weight` (same scalar as the K=2 sibling); Phase E may + // need per-tensor LR tuning if magnitude balance becomes an + // issue, but for B4 the unified aux_weight matches the K=2 + // sibling's pattern. + // + // Reads weights at indices [163, 165] (W1, W2) — same offset + // arithmetic as the K=2 sibling at [119, 121]. + let to_w1_bwd = on_w_ptrs[163]; + let to_w2_bwd = on_w_ptrs[165]; + let aux_kto = AUX_OUTCOME_K; + self.aux_to_bwd.backward( + &self.stream, + h_s2_aux_ptr, + to_w1_bwd, to_w2_bwd, + self.aux_to_hidden_buf.raw_ptr(), + self.aux_to_softmax_buf.raw_ptr(), + self.aux_to_label_buf.raw_ptr(), + self.aux_to_valid_count_buf.raw_ptr(), + b, sh2, aux_kto, + self.aux_partial_to_w1.raw_ptr(), + self.aux_partial_to_b1.raw_ptr(), + self.aux_partial_to_w2.raw_ptr(), + self.aux_partial_to_b2.raw_ptr(), + self.aux_dh_s2_to_buf.raw_ptr(), + )?; + // Reduce + SAXPY each of the 8 aux param tensors. Tensor layout // (SP13 B1.1a; next-bar K_NB flipped 1 → 2): // [119] aux_nb_w1 [H, SH2] ← partial [B, H*SH2] @@ -19561,7 +19634,10 @@ impl GpuDqnTrainer { // sums per-sample partials, no extra division. SAXPY alpha is // `aux_weight` (small positive, [0.05, 0.3]). let final_ptr = self.aux_param_grad_final_buf.raw_ptr(); - let aux_param_specs: [(usize, u64, usize); 8] = [ + // SP22 H6 vNext Phase B4 (2026-05-14): 4 new trade-outcome tensor + // entries at indices [163..167) appended after the 8 K=2/K=5 entries. + // SAXPY iteration uniformly scales all 12 by `aux_weight`. + let aux_param_specs: [(usize, u64, usize); 12] = [ (119, self.aux_partial_nb_w1.raw_ptr(), aux_h * sh2), (120, self.aux_partial_nb_b1.raw_ptr(), aux_h), (121, self.aux_partial_nb_w2.raw_ptr(), aux_knb * aux_h), @@ -19570,6 +19646,11 @@ impl GpuDqnTrainer { (124, self.aux_partial_rg_b1.raw_ptr(), aux_h), (125, self.aux_partial_rg_w2.raw_ptr(), aux_kr * aux_h), (126, self.aux_partial_rg_b2.raw_ptr(), aux_kr), + // SP22 H6 vNext Phase B4: trade-outcome head (K=3) param grads. + (163, self.aux_partial_to_w1.raw_ptr(), aux_h * sh2), + (164, self.aux_partial_to_b1.raw_ptr(), aux_h), + (165, self.aux_partial_to_w2.raw_ptr(), aux_kto * aux_h), + (166, self.aux_partial_to_b2.raw_ptr(), aux_kto), ]; for &(tensor_idx, partial_ptr, p_len) in &aux_param_specs { // Step (a): reduce per-sample partials [B, P] → final [P]. @@ -19659,6 +19740,29 @@ impl GpuDqnTrainer { .map_err(|e| MLError::ModelError(format!( "aux_heads_backward saxpy dh_s2_aux regime: {e}" )))?; + // SP22 H6 vNext Phase B4 (2026-05-14): trade-outcome head's + // per-sample dh_s2_aux. SAXPYs into the SAME `dh_s2_aux_accum` + // buffer as the K=2 / K=5 heads — all three head's gradients + // flow through the aux trunk to its own weights, NOT into Q's + // encoder (structural stop-grad in `aux_trunk_backward`). + // Same `aux_w` scaling as the K=2/K=5 SAXPYs above; Phase E + // may differentiate per-head weights if magnitude balance + // requires it. + let dh_to_ptr = self.aux_dh_s2_to_buf.raw_ptr(); + self.stream + .launch_builder(&self.saxpy_f32_kernel) + .arg(&dh_s2_aux_accum_ptr) + .arg(&dh_to_ptr) + .arg(&aux_w) + .arg(&n_dh) + .launch(LaunchConfig { + grid_dim: (blocks_dh, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: 0, + }) + .map_err(|e| MLError::ModelError(format!( + "aux_heads_backward saxpy dh_s2_aux trade_outcome: {e}" + )))?; } // SP14 Layer C Phase C.5b: aux trunk backward — propagates the diff --git a/crates/ml/src/ensemble/adapters/dqn.rs b/crates/ml/src/ensemble/adapters/dqn.rs index 062de4b08..bc5bd2127 100644 --- a/crates/ml/src/ensemble/adapters/dqn.rs +++ b/crates/ml/src/ensemble/adapters/dqn.rs @@ -273,13 +273,26 @@ mod tests { static CUDA_LOCK: Mutex<()> = Mutex::new(()); fn shared_device() -> MlDevice { - SHARED_CUDA + let device = SHARED_CUDA .get_or_init(|| { DeviceConfig::Auto .resolve() .expect("CUDA required") }) - .clone() + .clone(); + // Push the shared CUDA context onto the CURRENT thread. Mirrors + // the `cuda_stream()` helper fix in cuda_pipeline/mod.rs: CUDA + // contexts are per-thread state; under `cargo test` parallel + // mode, test threads that call `shared_device()` after the + // OnceLock init land without the context current, and any + // subsequent CUDA call (notably cuBLAS GEMM in DqnInference + // Adapter::predict) can flake with CUDA_ERROR_INVALID_CONTEXT. + // Idempotent — no-op when already current. + if let Ok(ctx) = device.cuda_context() { + ctx.bind_to_thread() + .expect("bind CUDA context to test thread"); + } + device } fn test_config() -> DQNConfig { diff --git a/docs/dqn-wire-up-audit.md b/docs/dqn-wire-up-audit.md index 6c7f47ab6..cf5db06e5 100644 --- a/docs/dqn-wire-up-audit.md +++ b/docs/dqn-wire-up-audit.md @@ -17889,3 +17889,27 @@ Fourth Rust-side commit of Phase B. Adds collector-side trade-outcome head: 5 st **Phase B4 next**: trainer-side replay-batch chain — `aux_trade_outcome_forward` on the replay batch (reading replay-buffer h_s2_aux + labels from replay-buffer scatter), `aux_trade_outcome_loss_reduce`, `aux_trade_outcome_backward`, Adam SAXPY for the 4 new weight tensors. The trainer-side wireup is the larger commit; B3's collector wiring is the smaller prerequisite (the replay-buffer scatter producer must populate before the trainer consumer reads). Verification: `cargo check -p ml` clean. Lib test suite `cargo test -p ml --lib` → 1016 passing / 0 failing (unchanged from post-fix-sweep baseline). The forward + label kernel launches don't regress any existing path because they produce buffers that aren't read. + +#### Phase B4 — Trainer-side replay-batch chain wireup (2026-05-14) + +Fifth Rust-side commit of Phase B. Wires the trade-outcome head's forward + loss reduce + backward + per-sample partial reduce + SAXPY into the trainer's `aux_heads_forward` and `aux_heads_backward` methods. Adam SAXPY for the 4 new weight tensors at [163..167) extends uniformly via the existing `aux_param_specs` array iteration. + +**Changes to `aux_heads_forward`**: +- Steps 7 + 8 appended after the K=2 next-bar head's loss reduce: K=3 `aux_to_fwd.forward()` reading weights at `[163..167)` (Phase B1 additions) → writing to `aux_to_*` save-for-backward tiles (Phase B2 additions); K=3 `aux_to_fwd.loss_reduce()` writing `aux_to_loss_scalar_buf` + `aux_to_valid_count_buf`. + +**Changes to `aux_heads_backward`**: +- K=3 `aux_to_bwd.backward()` appended after K=5 regime backward, emitting per-sample partials (dW1, db1, dW2, db2) + per-sample dh_s2_aux. +- `aux_param_specs` array extended from 8 → 12 entries. The per-tensor reduce + SAXPY loop iterates uniformly over all 12 entries, scaling each by `aux_weight` (same scalar as K=2/K=5 heads). Phase E may differentiate per-head weights if sparse-label gradient magnitude imbalance becomes an issue. +- K=3 dh_s2_aux SAXPY appended after K=5's, accumulating all three heads' gradients into the aux trunk's `dh_s2_aux_accum`. The aux trunk's structural stop-gradient (no `dx_in` output) keeps Q's encoder protected from aux contamination. + +**Label semantic (cold-start)**: `aux_to_label_buf` is `alloc_zeros` (all 0 = Profit) until Phase B4b lands the replay-buffer label scatter. The model trains on "predict Profit everywhere" — degraded but well-defined (no NaN). Mirrors the K=2 next-bar head's known-degraded state between B1.1a (forward landed) and B1.1b (label producer wired). + +**Adam SAXPY**: The existing global Adam SAXPY iterates `0..NUM_WEIGHT_TENSORS` (now 167). The 4 new weight slots get their gradient SAXPYs followed by Adam m/v updates uniformly — no per-tensor wiring needed. This is the architectural payoff of Phase B1's `NUM_WEIGHT_TENSORS` bump. + +**Test flake mitigation**: Added `bind_to_thread()` to `ensemble::adapters::dqn::tests::shared_device()` mirroring the `cuda_stream()` test-helper pattern from the fix sweep. The `test_dqn_checkpoint_round_trip` test had intermittent CUDA-context-thread-state flakes under parallel test runs; the bind is idempotent and forces context current on every test thread that accesses the shared device. Still occasionally non-deterministic at the prediction-direction level (the test's `disable_noise()` zeros NoisyLinear epsilon but may leave other randomness sources untouched; runs alternate `pred1=-1 pred2=1` ↔ `pred1=1 pred2=-1`). Not introduced by Phase B4 — the underlying NoisyLinear randomness has been flaky since pre-vNext. + +**Verification**: `cargo check -p ml` clean (21 warnings, none new). `cargo test -p ml --lib` → **1016 passing / 0 failing / 55 ignored** (unchanged from post-fix-sweep baseline). + +**Phase B4b next**: Replay-buffer label scatter so trainer's `aux_to_label_buf` gets populated from the rollout's `exp_aux_to_label_buf` per (i, t) — wires the cold-start degraded "predict Profit everywhere" path to the actual sparse trade-outcome label distribution. + +**Phase B5 (spec's actual "Phase B")**: Input concat 256→262 with `plan_params`. Requires a kernel-side change to the forward kernel's `h_s2_aux` input + W1 shape bump to `[H, 262]`. Small touch-up vs B0-B4's atomic wireup.