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:
@@ -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;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user