Files
foxhunt/crates/ml-alpha/cuda/variable_selection.cu
jgrusewski 9607f33518 perf(ml-alpha): block-per-row VSN bwd refactor (Phase B commit 3)
variable_selection_bwd refactored from grid=(1,1,1) to grid=(B*K,1,1).
VSN's n_rows = B*K positions (one row per (batch, K-position) pair);
block-per-row matches the existing fwd kernel's layout.

Adds 2 per-row grad scratch buffers + 2 reduce_axis0 launches:
  vsn_grad_w_scratch_d  [B*K, FEATURE_DIM, FEATURE_DIM]
  vsn_grad_b_scratch_d  [B*K, FEATURE_DIM]
~210 KB scratch at B=32, K=64.

VSN bwd runs 1×/step (not K×) so the absolute wall-time win here is
small versus commits 1+2. Done for pattern uniformity — every per-batch
or per-row bwd in the trainer now uses scratch+reducer.

All 9 perception_overfit smokes pass.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-05-17 23:57:32 +02:00

221 lines
8.6 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.
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_FEATURE_DIM;
// 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_FEATURE_DIM + 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_FEATURE_DIM + 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;
}
}