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>
221 lines
8.6 KiB
Plaintext
221 lines
8.6 KiB
Plaintext
// 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;
|
||
}
|
||
}
|