// 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; } }