diff --git a/crates/ml/build.rs b/crates/ml/build.rs index 739d71b51..7591f16b3 100644 --- a/crates/ml/build.rs +++ b/crates/ml/build.rs @@ -812,6 +812,15 @@ fn main() { // `docs/plans/2026-05-14-sp22-h6-vNext-trade-outcome-aux.md` for // the full vNext architecture spec. "aux_trade_outcome_forward_kernel.cu", + // SP22 H6 vNext Phase A4 (2026-05-13): sparse cross-entropy reduce + // over the K=3 softmax tile from `aux_trade_outcome_forward`. + // Mirrors `aux_next_bar_loss_reduce`'s single-block shmem-tree + // pattern: mean over B_valid (skip label=-1 rows), fmaxf(p, 1e-30) + // numerical floor, saved valid_count for backward. Dead code until + // Phase A5 (backward) + Phase B (Rust launcher wireup) land. Kept + // separate from the K=2 sibling for diagnostic isolation, sparse- + // label semantic clarity, and future per-class weighting headroom. + "aux_trade_outcome_loss_reduce_kernel.cu", // SP22 H6 (2026-05-12): per-env p_up extractor that copies // `aux_softmax[env, 1]` → `prev_aux_dir_prob[env]` after each aux // forward. The cache buffer is read by `experience_state_gather` diff --git a/crates/ml/src/cuda_pipeline/aux_trade_outcome_loss_reduce_kernel.cu b/crates/ml/src/cuda_pipeline/aux_trade_outcome_loss_reduce_kernel.cu new file mode 100644 index 000000000..a1ce1e0f0 --- /dev/null +++ b/crates/ml/src/cuda_pipeline/aux_trade_outcome_loss_reduce_kernel.cu @@ -0,0 +1,150 @@ +// crates/ml/src/cuda_pipeline/aux_trade_outcome_loss_reduce_kernel.cu +// +// SP22 H6 vNext Phase A4 (2026-05-13) — sparse cross-entropy reduce over +// the K=3 trade-outcome softmax tile produced by +// `aux_trade_outcome_forward`. +// +// Mirrors `aux_next_bar_loss_reduce`'s single-block shmem-tree reduce +// pattern (mean over valid-row count, fmaxf(p_tgt, 1e-30) numerical floor, +// saved valid_count for backward). Kept as a SEPARATE kernel from the K=2 +// sibling because: +// +// 1. Diagnostic isolation: the trade-outcome head's loss feeds a +// DIFFERENT HEALTH_DIAG slot (`aux_outcome_ce_ema`, allocated in +// Phase F) than the next-bar head's `aux_next_bar_mse_ema_idx`. +// Two distinct kernels = two distinct cubins = clean +// per-loss-source attribution in profiles and reward decomposition. +// +// 2. Sparse-label semantics: the trade-outcome producer emits a `-1` +// mask on EVERY bar without a trade-close event, which is the +// common case (~95-99% of bars). The next-bar head sees ~50-100% +// valid labels (only the trailing H bars near series-end get +// masked). The aggressive sparsity changes the gradient-scale +// arithmetic: an active-batch fraction of ~1-5% means each valid +// row's gradient gets divided by a much smaller B_valid → higher +// per-trade-close gradient magnitude. The kernel logic handles +// this correctly via `fmaxf(valid, 1.0f)` divisor (matches the +// backward kernel's `inv_B = 1 / B_valid` semantics) but the +// independent kernel lets Phase F's Adam-LR controller tune the +// trade-outcome head separately from the next-bar head if +// magnitude calibration becomes necessary. +// +// 3. Future per-class weighting: Profit/Stop/Timeout class imbalance +// can be ~3:1 to ~10:1 depending on the regime. Phase D+ may need +// class-weighted CE; an isolated kernel makes the modification +// surgical without touching the K=2 head's contract. +// +// Inputs / outputs (mirror the K=2 sibling contract) +// ────────────────────────────────────────────────── +// softmax [B, K=3] f32 row-major — saved by +// `aux_trade_outcome_forward`. +// labels [B] i32 in {-1, 0, 1, 2} — from +// `trade_outcome_label_kernel` at the +// same step (sparse: -1 = no trade +// close at this bar). +// B batch size (= N × L in the replay-batch use case; +// = n_envs in the rollout-per-step use +// case). K-runtime arg even though this +// kernel is K=3-semantic — keeps the +// launcher signature symmetric with +// the K=2 sibling for testing. +// loss_out [1] mean CE over the B_valid valid rows. +// valid_count_out [1] B_valid (saved for backward — matches the K=2 +// sibling's contract). +// +// Discipline +// ────────── +// - No atomicAdd per `feedback_no_atomicadd.md` — single-block shmem- +// tree reduction. +// - GPU-only per `feedback_cpu_is_read_only.md`. +// - Numerical floor at 1e-30 prevents -log(0) propagation through the +// EMA when an untrained softmax emits near-zero probabilities. +// +// Block: AUX_BLOCK threads. Grid: (1, 1, 1). +// Shared memory: 2 * AUX_BLOCK floats (loss numerator + valid-count +// partials, lockstep-reduced). + +#include +#include + +#ifndef AUX_BLOCK +#define AUX_BLOCK 256 +#endif + +extern "C" __global__ void aux_trade_outcome_loss_reduce( + /* [B, K=3] f32 row-major softmax tile from + * `aux_trade_outcome_forward`. */ + const float* __restrict__ softmax, + /* [B] i32 labels in `{-1, 0=Profit, 1=Stop, 2=Timeout}`. Producer is + * `trade_outcome_label_kernel.cu` — emits -1 (mask) on bars with no + * trade-close event (the common case at ~95-99% of bars). */ + const int* __restrict__ labels, + int B, + int K, + /* [1] mean CE over the B_valid rows whose label != -1. An all-skip + * batch produces 0 (no NaN; fmaxf(valid, 1) divisor). */ + float* __restrict__ loss_out, + /* [1] B_valid saved for backward. Matches the K=2 sibling's contract + * — the backward kernel uses this as `inv_B = 1 / B_valid` for the + * d_logits scaling. An all-skip batch saves 0; backward then emits + * zero gradients across the board (no NaN). */ + float* __restrict__ valid_count_out +) { + if (blockIdx.x != 0) return; + extern __shared__ float smem[]; + float* sh_loss = smem; + float* sh_valid = smem + blockDim.x; + const int tid = threadIdx.x; + const int block = blockDim.x; + + /* Per-thread strided accumulation: each thread folds its slice of + * `[0, B)` into private (loss_numer, valid_count) accumulators. + * Mask-encoded rows (label == -1) contribute 0 to BOTH accumulators — + * the common case for trade-outcome labels where most bars have no + * trade-close event. The valid-row branch reads `softmax[i, label]` + * directly (no max-shift needed; the forward kernel already + * materialised stable softmax probs). */ + float local_loss = 0.0f; + float local_valid = 0.0f; + for (int i = tid; i < B; i += block) { + const int tgt = labels[i]; + if (tgt == -1) continue; + /* Defensive bounds — out-of-range labels degrade to skip + * rather than read past the end of the [B, K] softmax row. The + * label producer (Phase A2 `trade_outcome_label_kernel`) emits + * only -1 / {0, 1, 2} so this branch is dead under correct + * producer behavior; kept for kernel robustness against + * producer bugs. */ + if (tgt < 0 || tgt >= K) continue; + const float p_tgt = softmax[(size_t)i * K + tgt]; + /* Numerical floor at 1e-30 — matches the K=2 sibling's clip. + * An untrained `aux_trade_outcome_forward` with extreme W2 can + * emit p ≈ 0 for some class; -log(0) = +inf would propagate + * to the EMA. log(1e-30) ≈ -69 is well outside the trained + * regime — the floor is robustness, not a substitute for + * training stability. */ + const float p_clipped = fmaxf(p_tgt, 1e-30f); + local_loss += -logf(p_clipped); + local_valid += 1.0f; + } + sh_loss[tid] = local_loss; + sh_valid[tid] = local_valid; + __syncthreads(); + + /* Standard log2(BLOCK) tree reduction over both arrays in lockstep + * — identical pattern to the K=2 sibling so the proven correctness + * carries over directly. */ + for (int s = block / 2; s > 0; s >>= 1) { + if (tid < s) { + sh_loss[tid] += sh_loss[tid + s]; + sh_valid[tid] += sh_valid[tid + s]; + } + __syncthreads(); + } + + if (tid == 0) { + const float valid = sh_valid[0]; + loss_out[0] = sh_loss[0] / fmaxf(valid, 1.0f); + valid_count_out[0] = valid; + } +} diff --git a/docs/dqn-wire-up-audit.md b/docs/dqn-wire-up-audit.md index 92e9364fc..96f0fd7f6 100644 --- a/docs/dqn-wire-up-audit.md +++ b/docs/dqn-wire-up-audit.md @@ -17693,3 +17693,32 @@ Second foundation commit for the SP22 H6 vNext trade-outcome aux head. Lands two - Phase F: validation smoke Cargo check clean. Registry tests clean. Ready for Phase A4. + +#### Phase A4 — aux_trade_outcome loss reduce kernel (2026-05-13) + +Third foundation commit for the SP22 H6 vNext trade-outcome aux head. Lands the sparse cross-entropy reduce. + +**NEW kernel `aux_trade_outcome_loss_reduce_kernel.cu`** — K=3 mean CE over the trade-outcome softmax tile with `-1` label masking. +- Pattern: single-block shmem-tree reduce, two parallel partial strips (loss_numer + valid_count) reduced in lockstep. Identical structural shape to `aux_next_bar_loss_reduce` (K=2 sibling); same `fmaxf(p_tgt, 1e-30)` numerical floor, same `fmaxf(valid, 1.0)` all-skip-batch guard, same `valid_count_out[1]` save-for-backward contract. +- Sparse-label arithmetic: most bars (~95-99%) have `label=-1` → contribute 0 to both numerator and valid_count. The few valid-trade-close rows produce CE numerator divided by a small B_valid → higher per-trade-close gradient magnitude. The kernel correctly handles all-skip batches (loss=0, valid=0 → backward emits zero gradients). +- **Kept separate from the K=2 sibling** for three reasons: + 1. Diagnostic isolation — distinct HEALTH_DIAG slot (`aux_outcome_ce_ema`, lands in Phase F), distinct cubin in profiles, clean attribution. + 2. Sparse-label semantic clarity — the kernel comment documents the ~1-5% active-batch fraction expected at the trade-outcome head. + 3. Future per-class weighting headroom — Profit/Stop/Timeout imbalance can be 3:1 to 10:1; an isolated kernel makes class-weighted CE a surgical modification without touching the K=2 head's contract. +- Cubin: `aux_trade_outcome_loss_reduce_kernel.cubin` (9.9 KB). +- Phase A4 (this commit) **dead code** — no Rust launcher yet. Phase A5 (backward) reads `valid_count_out` for `inv_B = 1/B_valid` scaling. Phase B wires the full forward→loss→backward chain. + +**Discipline**: +- `feedback_no_atomicadd` — single-block shmem-tree reduction. +- `pearl_first_observation_bootstrap` — sentinel 0.0 for valid_count means cold-start (no trade-closes yet) produces zero gradients gracefully. +- `feedback_cpu_is_read_only` — all GPU-side. + +**Remaining vNext work**: +- Phase A5: `aux_trade_outcome_backward_kernel.cu` +- Phase B: 262-dim input concat (`h_s2_aux` || `plan_params`) +- Phase C: 3-slot state assembly (state[121..124] = p_Profit, p_Stop, p_Timeout) +- Phase D: 12-weight W atom-shift (4 actions × 3 outcomes) +- Phase E: dW + Adam for W[4, 3] +- Phase F: validation smoke + +Cargo check clean. Ready for Phase A5.