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>
This commit is contained in:
jgrusewski
2026-05-29 10:49:58 +02:00
parent 4e629830b6
commit 104fe81ca8

View File

@@ -26,6 +26,24 @@
#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]
@@ -38,7 +56,7 @@ extern "C" __global__ void variable_selection_fwd(
int tid = threadIdx.x;
if (row >= n_rows) return;
const float* x_row = x + (long long)row * VSN_FEATURE_DIM;
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];
@@ -170,7 +188,7 @@ extern "C" __global__ void variable_selection_bwd(
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_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;
@@ -201,7 +219,7 @@ extern "C" __global__ void variable_selection_bwd(
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_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;
}