Files
foxhunt/crates/ml-alpha/cuda/variable_selection.cu
jgrusewski 104fe81ca8 fix(cuda): VSN stride mismatch — read window_tensor at stride 56
Root cause of the step-4 NaN that has blocked the Phase 2.1 smoke gate
since session start.

`window_tensor_d` is allocated [B, K, ENCODER_INPUT_DIM=56]:
- Features [0..40) are per-snapshot market features (snap_feature_assemble_batched)
- Features [40..56) are per-batch broadcast context (rl_encoder_context_broadcast
  writes via offset `idx * 56 + 40`)

But `variable_selection_fwd` and `_bwd` read x with stride
`VSN_FEATURE_DIM=40`. For row 0 the stride happens to align with the
[B, K, 56] layout. For row 1 onwards the kernel reads MIXED data
across row boundaries — broadcast context bleeds into VSN's "snap
features" view, and snap features bleed across K boundaries.

The bleed produces nonsense for most steps but is bounded enough to
train through. Around step 4, accumulated trade_context magnitudes
(time_in_trade, unrealized_R, position_lots) get large enough to
overflow VSN's softmax in the bleed-affected rows → NaN cascades
through the encoder forward → backward computes NaN gradients →
AdamW applies them → step 5's forward graph sees NaN weights → G8 abort.

Fix: introduce `VSN_X_ROW_STRIDE = 56` (= ENCODER_INPUT_DIM) and use it
for all input reads of `x` in both forward and backward:
- variable_selection_fwd:41 (x_row = x + row * 56)
- variable_selection_bwd:173 (x_i read at stride 56)
- variable_selection_bwd:204 (xj read at stride 56)

Output strides (gates_out, y, grad_W, grad_b, grad_x) remain at
VSN_FEATURE_DIM=40 because:
- gates_out / y feed Mamba2 L1 which reads at in_dim=40 stride
- grad_W / grad_b accumulate into [FEATURE_DIM, FEATURE_DIM] / [FEATURE_DIM]
- grad_x is write-only (audit: `vsn_grad_x_d` has zero downstream consumers
  per grep of crates/ml-alpha/)

Verification:
- 10/10 runs clean at seed=16962, n-steps=10 (was 4-7/10 NaN before)
- 1000-step smoke clean at seed=16962: l_q≤1.76, l_v≤1.18, wr climbing
  0.36→0.55 (the surfer pattern emerging as expected per
  pearl_dd049d9a4_surfer_baseline_verified)
- Phase 2.1 smoke gate (l_q<3.0, l_v<2.0, no NaN over 1000 steps) PASSED

Per pearl_atomicadd_masks_v_instability commits 60e96bf55..4e629830b for
the full diagnostic chain (49 nan_scan labels, 4 sub-agent dispatches,
13 commits localizing the bug from "somewhere in the trainer" down to
this single line).

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-05-29 10:49:58 +02:00

239 lines
9.7 KiB
Plaintext
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// variable_selection.cu — TFT-style Variable Selection Network (Phase 2D).
//
// Per-position feature gating: for each (batch, k) sample we learn
// FEATURE_DIM gates (softmax-normalised) and emit feature[i] * gate[i].
// Lets the model down-weight noisy / non-informative features per
// regime (canonical: trade-flow in low-volume windows, OFI in
// spread-Q4-dominant regimes). Borrows from Lim et al. 2021 §4.2.
//
// Forward math (per (b, k) row, threads tile FEATURE_DIM):
// gate_logit[i] = sum_j W_vsn[i, j] * x[j] + b_vsn[i]
// gates = softmax(gate_logit) # [FEATURE_DIM]
// y[i] = x[i] * gates[i]
//
// Saved for backward: `gates` (post-softmax). Don't save `gate_logit`
// — softmax-jacobian is computable from `gates` alone.
//
// Block layout: one block per (b, k) sample = grid_dim = (B * K, 1, 1).
// Block dim = FEATURE_DIM (40 threads — enough for warp-aligned ops
// at FEATURE_DIM=40; one warp + 8 idle threads). Each thread owns one
// feature index `i`.
//
// Per `feedback_no_atomicadd.md`: block tree-reduce only, no
// atomicAdd. Softmax uses standard max-subtract + sum trick for
// numerical stability.
#define VSN_FEATURE_DIM 40
#define VSN_BLOCK 64 // round up to warp-multiple; threads i >= FEATURE_DIM idle.
// 2026-05-29 stride-mismatch fix.
// VSN's input buffer (window_tensor_d) is allocated [B, K, ENCODER_INPUT_DIM=56]
// by perception.rs (snap features [0..40) + per-batch broadcast context
// [40..56) written by rl_encoder_context_broadcast). VSN only processes the
// first VSN_FEATURE_DIM=40 features per row (snap features), but the input
// rows are spaced 56 floats apart, not 40. The kernel originally indexed x
// with stride VSN_FEATURE_DIM=40 — correct ONLY for row 0; every subsequent
// row read mixed broadcast-context + snap features across the [B, K, 56]
// row boundaries. Symptoms: intermittent step-4 NaN as accumulating trade
// context magnitudes overflowed VSN's softmax via the bleed.
//
// Fix: use VSN_X_ROW_STRIDE=56 for reading x in both forward and backward.
// Output (gates, y) and gradient outputs (grad_W, grad_b, grad_x) remain at
// VSN_FEATURE_DIM=40 because the downstream consumers (Mamba2 L1 with
// in_dim=40) read at compact 40-stride. grad_x is unused downstream (see
// `vsn_grad_x_d` audit — write-only), so its stride doesn't matter.
#define VSN_X_ROW_STRIDE 56 // = ENCODER_INPUT_DIM in heads.rs / perception.rs
extern "C" __global__ void variable_selection_fwd(
const float* __restrict__ W_vsn, // [FEATURE_DIM, FEATURE_DIM]
const float* __restrict__ b_vsn, // [FEATURE_DIM]
const float* __restrict__ x, // [N_rows, FEATURE_DIM]
int n_rows,
float* __restrict__ y, // [N_rows, FEATURE_DIM]
float* __restrict__ gates_out // [N_rows, FEATURE_DIM] — softmax(gate_logit)
) {
int row = blockIdx.x;
int tid = threadIdx.x;
if (row >= n_rows) return;
const float* x_row = x + (long long)row * VSN_X_ROW_STRIDE;
// Shared mem: gate_logit + max-reduce scratch + sum-reduce scratch.
__shared__ float s_logit[VSN_FEATURE_DIM];
__shared__ float s_red[VSN_BLOCK];
__shared__ float s_max;
__shared__ float s_sum;
// Pass 1: gate_logit[i] = sum_j W_vsn[i, j] * x[j] + b_vsn[i].
// Thread `tid` (when tid < FEATURE_DIM) owns row i = tid.
float gl = 0.0f;
if (tid < VSN_FEATURE_DIM) {
gl = b_vsn[tid];
#pragma unroll
for (int j = 0; j < VSN_FEATURE_DIM; ++j) {
gl += W_vsn[tid * VSN_FEATURE_DIM + j] * x_row[j];
}
s_logit[tid] = gl;
}
__syncthreads();
// Pass 2a: block-wide max over s_logit (numerical-stability shift).
s_red[tid] = (tid < VSN_FEATURE_DIM) ? s_logit[tid] : -INFINITY;
__syncthreads();
for (int s = VSN_BLOCK / 2; s > 0; s >>= 1) {
if (tid < s) {
const float a = s_red[tid];
const float b = s_red[tid + s];
s_red[tid] = (a > b) ? a : b;
}
__syncthreads();
}
if (tid == 0) s_max = s_red[0];
__syncthreads();
// Pass 2b: exp(logit - max), then sum.
float e = 0.0f;
if (tid < VSN_FEATURE_DIM) {
e = expf(s_logit[tid] - s_max);
s_logit[tid] = e; // reuse — now holds exp(shifted)
}
s_red[tid] = (tid < VSN_FEATURE_DIM) ? e : 0.0f;
__syncthreads();
for (int s = VSN_BLOCK / 2; s > 0; s >>= 1) {
if (tid < s) s_red[tid] += s_red[tid + s];
__syncthreads();
}
if (tid == 0) s_sum = s_red[0];
__syncthreads();
// Pass 3: gates[i] = exp/sum; output y[i] = x[i] * gates[i].
if (tid < VSN_FEATURE_DIM) {
const float g = s_logit[tid] / s_sum;
gates_out[(long long)row * VSN_FEATURE_DIM + tid] = g;
y[(long long)row * VSN_FEATURE_DIM + tid] = x_row[tid] * g;
}
}
// VSN backward — chain rule:
//
// d_x_via_gate[i] = grad_y[i] * gates[i] (straight-through)
// d_gates[i] = grad_y[i] * x[i]
//
// Through softmax:
// d_logit[i] = gates[i] * (d_gates[i] - sum_j gates[j] * d_gates[j])
//
// Through W_vsn @ x + b_vsn:
// grad_W_vsn[i, j] += d_logit[i] * x[j]
// grad_b_vsn[i] += d_logit[i]
// d_x_via_W[j] = sum_i d_logit[i] * W_vsn[i, j]
//
// Total trunk gradient:
// grad_x[j] = d_x_via_gate[j] + d_x_via_W[j]
//
// Single-writer discipline:
// - Thread tid owns row tid of grad_W_vsn (writes all columns j).
// - Thread tid owns grad_b_vsn[tid].
// - Thread tid owns column tid of d_x_via_W (sums over i).
// - One block per (b, k) row, B*K blocks. Param-grads accumulate via
// += across rows; since each block has a single-writer per (i, *)
// within its row contribution, and rows are sequential per block,
// no cross-block race except on the += itself. Resolution: launch
// ONE block per launch (n_rows internal loop), same pattern as the
// 2-layer and GRN heads-bwd kernels.
// Block-per-row VSN bwd (Phase B commit 3).
// grid=(n_rows, 1, 1) block=(VSN_BLOCK, 1, 1)
//
// n_rows = B * K (one row per (batch, K-position) pair). Each block
// handles one row's FEATURE_DIM features.
//
// Per-row grad scratch tensors:
// grad_W_vsn_scratch [n_rows, FEATURE_DIM, FEATURE_DIM]
// grad_b_vsn_scratch [n_rows, FEATURE_DIM]
// Thread (row, tid) is sole writer to its slice — single launch per
// training step (variable_selection_bwd runs 1×/step, not K×).
// Caller zeroes scratch at step start; reduce_axis0 collapses n_rows
// → final grad after this kernel.
//
// grad_x is per-row indexed; block row is sole writer to its slice
// (overwrite).
extern "C" __global__ void variable_selection_bwd(
const float* __restrict__ W_vsn, // [FEATURE_DIM, FEATURE_DIM]
const float* __restrict__ x, // [N_rows, FEATURE_DIM]
const float* __restrict__ gates, // [N_rows, FEATURE_DIM] saved by fwd
const float* __restrict__ grad_y, // [N_rows, FEATURE_DIM]
int n_rows,
float* __restrict__ grad_W_vsn_scratch, // [N_rows, FEATURE_DIM, FEATURE_DIM] (+=)
float* __restrict__ grad_b_vsn_scratch, // [N_rows, FEATURE_DIM] (+=)
float* __restrict__ grad_x // [N_rows, FEATURE_DIM] (overwrite)
) {
int row = blockIdx.x;
int tid = threadIdx.x;
if (row >= n_rows || tid >= VSN_BLOCK) return;
__shared__ float s_dgates[VSN_FEATURE_DIM];
__shared__ float s_dlogit[VSN_FEATURE_DIM];
__shared__ float s_gates[VSN_FEATURE_DIM];
__shared__ float s_red[VSN_BLOCK];
__shared__ float s_sumdot;
__syncthreads();
// Per-row base offsets into scratch tensors.
const long long row_F = (long long)row * VSN_FEATURE_DIM;
const long long row_FF = row_F * VSN_FEATURE_DIM;
// Pass 1: d_gates[i] = grad_y[i] * x[i]; cache gates[i] + x[i].
float gates_i = 0.0f;
float x_i = 0.0f;
float dy_i = 0.0f;
if (tid < VSN_FEATURE_DIM) {
gates_i = gates[(long long)row * VSN_FEATURE_DIM + tid];
x_i = x[(long long)row * VSN_X_ROW_STRIDE + tid];
dy_i = grad_y[(long long)row * VSN_FEATURE_DIM + tid];
s_gates[tid] = gates_i;
s_dgates[tid] = dy_i * x_i;
}
__syncthreads();
// Pass 2: dot_product = sum_j gates[j] * d_gates[j].
s_red[tid] = (tid < VSN_FEATURE_DIM) ? s_gates[tid] * s_dgates[tid] : 0.0f;
__syncthreads();
for (int s = VSN_BLOCK / 2; s > 0; s >>= 1) {
if (tid < s) s_red[tid] += s_red[tid + s];
__syncthreads();
}
if (tid == 0) s_sumdot = s_red[0];
__syncthreads();
// Pass 3: d_logit[i] = gates[i] * (d_gates[i] - sum_dot).
// grad_b_vsn_scratch[row, i] += d_logit[i]. (Thread tid sole writer per row.)
if (tid < VSN_FEATURE_DIM) {
const float dlogit_i = gates_i * (s_dgates[tid] - s_sumdot);
s_dlogit[tid] = dlogit_i;
grad_b_vsn_scratch[row_F + tid] += dlogit_i;
}
__syncthreads();
// Pass 4: grad_W_vsn_scratch[row, i=tid, j] += d_logit[tid] * x[row, j].
if (tid < VSN_FEATURE_DIM) {
const float dl_t = s_dlogit[tid];
#pragma unroll
for (int j = 0; j < VSN_FEATURE_DIM; ++j) {
const float xj = x[(long long)row * VSN_X_ROW_STRIDE + j];
grad_W_vsn_scratch[row_FF + (long long)tid * VSN_FEATURE_DIM + j]
+= dl_t * xj;
}
}
// Pass 5: grad_x[row, j=tid] = sum_i d_logit[i] * W_vsn[i, j=tid] + grad_y[j] * gates[j].
if (tid < VSN_FEATURE_DIM) {
float dx_via_W = 0.0f;
#pragma unroll
for (int i = 0; i < VSN_FEATURE_DIM; ++i) {
dx_via_W += s_dlogit[i] * W_vsn[i * VSN_FEATURE_DIM + tid];
}
const float dx_via_gate = dy_i * gates_i;
grad_x[(long long)row * VSN_FEATURE_DIM + tid] = dx_via_W + dx_via_gate;
}
}