From 5e23005deaddb753e67b855311a4ca25e6076a51 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sun, 17 May 2026 21:47:09 +0200 Subject: [PATCH] feat(ml-alpha): TFT GRN forward+backward kernels for multi-horizon heads (Phase 1.7a) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Per-horizon GRN structure (Lim et al. 2021 §3.3 adapted to scalar output): eta_2[k, m] = GELU(W1[k, m, :] @ h + b1[k, m]) # [HIDDEN] → [HEAD_MID] eta_1[k, m] = W2[k, m, :] @ eta_2[k, :] + b2[k, m] # [HEAD_MID] → [HEAD_MID] gate_lin[k] = W_gate[k, :] @ eta_1[k, :] + b_gate[k] # → scalar main[k] = W_main[k, :] @ eta_1[k, :] + b_main[k] # → scalar skip[k] = W_skip[k, :] @ h + b_skip[k] # [HIDDEN] → scalar logit[k] = skip[k] + sigmoid(gate_lin[k]) * main[k] p[k] = sigmoid(logit[k]) Gated residual lets each per-horizon head learn "linear vs deeper-transform" gating, matching the regime-conditional alpha pattern from pearl_snapshot_alpha_is_regime_conditional (~20% of book states carry the edge; spread-Q4 hits 75% acc, middle quintiles below chance). Backward chain rule covers all 10 parameter tensors + the trunk gradient (skip-path direct + main-path through W2→GELU→W1, lambda-scaled). Single-writer discipline (no atomicAdd per feedback_no_atomicadd.md): - Thread m owns row m of grad_w1 (col i in 0..HIDDEN), row m of grad_w2 (col m_in in 0..HEAD_MID), and column m of d_eta_2. - Threads 0..4 own per-horizon scalar grads (skip/gate/main biases). - Trunk grad_h tiles i over 2 strides of HEAD_MID for HIDDEN=128 coverage. Shared mem: ~6.5KB (s_a1 + s_z2 + s_d_eta1 + s_d_eta2 + s_d_z1 + scalars), well within 48KB limit. Existing 2-layer MLP kernels (Tasks 1.3/1.4) stay in the cubin as ablation baseline; the wired path becomes GRN once perception.rs lands. build.rs cache-bust → v7. Co-Authored-By: Claude Opus 4.7 --- crates/ml-alpha/build.rs | 9 +- crates/ml-alpha/cuda/multi_horizon_heads.cu | 331 ++++++++++++++++++++ 2 files changed, 336 insertions(+), 4 deletions(-) diff --git a/crates/ml-alpha/build.rs b/crates/ml-alpha/build.rs index 416198024..377a2facc 100644 --- a/crates/ml-alpha/build.rs +++ b/crates/ml-alpha/build.rs @@ -20,10 +20,11 @@ const KERNELS: &[&str] = &[ "layer_norm", // Phase 1: trunk pre-CfC normalisation ]; -// Cache bust v6 (2026-05-17): Phase 1 model capacity — LayerNorm -// kernel added (cuda/layer_norm.cu) + multi_horizon_heads.cu will get -// 2-layer MLP heads. Old cubins don't have the new symbols. Force -// fresh nvcc compile on the cluster's /cargo-target PVC. +// Cache bust v7 (2026-05-17): Phase 1.7 GRN heads — appended +// multi_horizon_heads_grn_fwd_batched + _bwd_batched kernels to +// cuda/multi_horizon_heads.cu (TFT gated residual: 2-layer MLP body + +// GLU gate + skip-projection from trunk). Old cubins don't have the +// new symbols. Force fresh nvcc compile. fn main() { println!("cargo:rerun-if-changed=build.rs"); diff --git a/crates/ml-alpha/cuda/multi_horizon_heads.cu b/crates/ml-alpha/cuda/multi_horizon_heads.cu index f8194dfdf..2d89d89c0 100644 --- a/crates/ml-alpha/cuda/multi_horizon_heads.cu +++ b/crates/ml-alpha/cuda/multi_horizon_heads.cu @@ -425,3 +425,334 @@ extern "C" __global__ void multi_horizon_heads_2layer_bwd_batched( __syncthreads(); } } + +// ───────────────────────────────────────────────────────────────────── +// TFT Gated Residual Network (GRN) heads — Phase 1.7 (2026-05-17) +// ───────────────────────────────────────────────────────────────────── +// +// Per-horizon GRN structure (Lim et al. 2021 §3.3, adapted to scalar +// output per horizon): +// +// eta_2[k, m_out] = GELU(W1[k, m_out, i] @ h[i] + b1[k, m_out]) # [HIDDEN] → [HEAD_MID] +// eta_1[k, m_out] = W2[k, m_out, m_in] @ eta_2[k, m_in] + b2[...] # [HEAD_MID] → [HEAD_MID] +// gate_lin[k] = W_gate[k, m] @ eta_1[k, m] + b_gate[k] # [HEAD_MID] → scalar +// main[k] = W_main[k, m] @ eta_1[k, m] + b_main[k] +// skip[k] = W_skip[k, i] @ h[i] + b_skip[k] # [HIDDEN] → scalar +// logit[k] = skip[k] + sigmoid(gate_lin[k]) * main[k] +// p[k] = sigmoid(logit[k]) +// +// Naming convention follows the existing 2-layer kernel: w1 = first +// linear (HIDDEN → HEAD_MID, GELU-activated), w2 = second linear +// (HEAD_MID → HEAD_MID, no activation; produces eta_1). +// +// Thread layout (block_dim = HEAD_MID = 64): +// m = threadIdx.x -- owns "output_row" dim of (k, m_out=m) for w1/w2/eta_1/eta_2 +// for skip / gate / main scalars: thread tid<5 handles horizon k=tid + +extern "C" __global__ void multi_horizon_heads_grn_fwd_batched( + const float* __restrict__ w1, // [5, HEAD_MID, HIDDEN] + const float* __restrict__ b1, // [5, HEAD_MID] + const float* __restrict__ w2, // [5, HEAD_MID_out, HEAD_MID_in] + const float* __restrict__ b2, // [5, HEAD_MID_out] + const float* __restrict__ w_gate, // [5, HEAD_MID] + const float* __restrict__ b_gate, // [5] + const float* __restrict__ w_main, // [5, HEAD_MID] + const float* __restrict__ b_main, // [5] + const float* __restrict__ w_skip, // [5, HIDDEN] + const float* __restrict__ b_skip, // [5] + const float* __restrict__ h, // [B, HIDDEN] + int n_batch, + float* __restrict__ probs, // [B, 5] + float* __restrict__ z1_out, // [B, 5, HEAD_MID] — pre-GELU z1 (= W1 @ h + b1) + float* __restrict__ a1_out, // [B, 5, HEAD_MID] — post-GELU eta_2 + float* __restrict__ z2_out, // [B, 5, HEAD_MID] — eta_1 (= W2 @ eta_2 + b2) + float* __restrict__ gate_logit_out, // [B, 5] — pre-sigmoid gate scalar + float* __restrict__ main_out, // [B, 5] + float* __restrict__ logit_out // [B, 5] — pre-final-sigmoid logit (skip + sigma(gate)*main) +) { + int b_idx = blockIdx.x; + int m = threadIdx.x; + if (b_idx >= n_batch || m >= HEAD_MID_H) return; + + const float* h_row = h + (long long)b_idx * HIDDEN_H; + + __shared__ float s_a1[N_HORIZONS_H * HEAD_MID_H]; // post-GELU eta_2 + __shared__ float s_z2[N_HORIZONS_H * HEAD_MID_H]; // eta_1 + + // Pass 1: z1[k, m] = W1[k, m, :] @ h + b1[k, m]; a1 = GELU(z1). + // Thread m owns row m (output dim) for all k. Sequential over HIDDEN. + #pragma unroll + for (int k = 0; k < N_HORIZONS_H; ++k) { + float z = b1[k * HEAD_MID_H + m]; + for (int i = 0; i < HIDDEN_H; ++i) { + z += w1[((long long)(k * HEAD_MID_H) + m) * HIDDEN_H + i] * h_row[i]; + } + const float a = gelu_act(z); + s_a1[k * HEAD_MID_H + m] = a; + z1_out[((long long)b_idx * N_HORIZONS_H + k) * HEAD_MID_H + m] = z; + a1_out[((long long)b_idx * N_HORIZONS_H + k) * HEAD_MID_H + m] = a; + } + __syncthreads(); + + // Pass 2: z2[k, m] = W2[k, m_out=m, m_in=n] @ a1[k, n] + b2[k, m]. + // Pure linear — no activation. Thread m owns output row m for all k. + #pragma unroll + for (int k = 0; k < N_HORIZONS_H; ++k) { + float z = b2[k * HEAD_MID_H + m]; + for (int n = 0; n < HEAD_MID_H; ++n) { + z += w2[((long long)(k * HEAD_MID_H) + m) * HEAD_MID_H + n] + * s_a1[k * HEAD_MID_H + n]; + } + s_z2[k * HEAD_MID_H + m] = z; + z2_out[((long long)b_idx * N_HORIZONS_H + k) * HEAD_MID_H + m] = z; + } + __syncthreads(); + + // Pass 3: per-horizon scalars. Threads 0..4 each handle one horizon. + // gate_lin[k] = W_gate[k, :] @ z2[k, :] + b_gate[k] + // main[k] = W_main[k, :] @ z2[k, :] + b_main[k] + // skip[k] = W_skip[k, :] @ h + b_skip[k] + // logit[k] = skip[k] + sigmoid(gate_lin[k]) * main[k] + // p[k] = sigmoid(logit[k]) + if (m < N_HORIZONS_H) { + const int k = m; + float gl = b_gate[k]; + float mv = b_main[k]; + #pragma unroll + for (int n = 0; n < HEAD_MID_H; ++n) { + const float z2_kn = s_z2[k * HEAD_MID_H + n]; + gl += w_gate[k * HEAD_MID_H + n] * z2_kn; + mv += w_main[k * HEAD_MID_H + n] * z2_kn; + } + float sk = b_skip[k]; + for (int i = 0; i < HIDDEN_H; ++i) { + sk += w_skip[k * HIDDEN_H + i] * h_row[i]; + } + const float sigmoid_gl = 1.0f / (1.0f + expf(-gl)); + const float logit = sk + sigmoid_gl * mv; + const float p = 1.0f / (1.0f + expf(-logit)); + + gate_logit_out[(long long)b_idx * N_HORIZONS_H + k] = gl; + main_out[(long long)b_idx * N_HORIZONS_H + k] = mv; + logit_out[(long long)b_idx * N_HORIZONS_H + k] = logit; + probs[(long long)b_idx * N_HORIZONS_H + k] = p; + } +} + +// GRN backward — chain rule: +// +// d_logit = grad_probs * p * (1 - p) +// d_skip = d_logit +// d_main = d_logit * sigma(gate_lin) +// d_gate_a = d_logit * main +// d_gate_lin = d_gate_a * sigma(gate_lin) * (1 - sigma(gate_lin)) +// +// Per-horizon scalar grads: +// grad_w_skip[k, i] += d_skip[k] * h[i] +// grad_b_skip[k] += d_skip[k] +// grad_w_main[k, m] += d_main[k] * eta_1[k, m] +// grad_b_main[k] += d_main[k] +// grad_w_gate[k, m] += d_gate_lin[k] * eta_1[k, m] +// grad_b_gate[k] += d_gate_lin[k] +// +// Through eta_1: +// d_eta_1[k, m] = d_main * w_main[k, m] + d_gate_lin * w_gate[k, m] +// grad_w2[k, m_out, m_in] += d_eta_1[k, m_out] * eta_2[k, m_in] +// grad_b2[k, m_out] += d_eta_1[k, m_out] +// d_eta_2[k, m_in] = sum_{m_out} d_eta_1[k, m_out] * w2[k, m_out, m_in] +// +// Through GELU + W1: +// d_z1[k, m] = d_eta_2[k, m] * gelu_prime(z1[k, m]) +// grad_w1[k, m, i] += d_z1[k, m] * h[i] +// grad_b1[k, m] += d_z1[k, m] +// +// Trunk gradient (lambda-scaled, sums two paths): +// grad_h[i] = sum_k lambda[k] * ( +// d_skip[k] * w_skip[k, i] (skip path) +// + sum_m d_z1[k, m] * w1[k, m, i] (main path through W2→W1) +// ) + carry[i] +// +// Single-writer discipline: +// Thread m (= m_out for w2, m for w1) owns row m of grad_w2 (writes +// cols m_in for all m_in) and row m of grad_w1 (writes cols i for +// all i). Thread m also owns col m of d_eta_2 (sums over m_out). +// Thread tid (with tid < 5) owns horizon-scalar grads. +// Trunk grad_h tiles i over 2 strides of HEAD_MID. + +extern "C" __global__ void multi_horizon_heads_grn_bwd_batched( + const float* __restrict__ w1, // [5, HEAD_MID, HIDDEN] + const float* __restrict__ w2, // [5, HEAD_MID_out, HEAD_MID_in] + const float* __restrict__ w_gate, // [5, HEAD_MID] + const float* __restrict__ w_main, // [5, HEAD_MID] + const float* __restrict__ w_skip, // [5, HIDDEN] + const float* __restrict__ probs, // [B, 5] + const float* __restrict__ grad_probs, // [B, 5] + const float* __restrict__ z1, // [B, 5, HEAD_MID] + const float* __restrict__ a1, // [B, 5, HEAD_MID] + const float* __restrict__ z2, // [B, 5, HEAD_MID] + const float* __restrict__ gate_logit, // [B, 5] + const float* __restrict__ main_val, // [B, 5] + const float* __restrict__ h, // [B, HIDDEN] + const float* __restrict__ grad_h_carry, // [B, HIDDEN] (nullptr OK) + const float* __restrict__ lambda, // [5] + int n_batch, + float* __restrict__ grad_w1, // [5, HEAD_MID, HIDDEN] (+=) + float* __restrict__ grad_b1, // [5, HEAD_MID] (+=) + float* __restrict__ grad_w2, // [5, HEAD_MID_out, HEAD_MID_in] (+=) + float* __restrict__ grad_b2, // [5, HEAD_MID_out] (+=) + float* __restrict__ grad_w_gate, // [5, HEAD_MID] (+=) + float* __restrict__ grad_b_gate, // [5] (+=) + float* __restrict__ grad_w_main, // [5, HEAD_MID] (+=) + float* __restrict__ grad_b_main, // [5] (+=) + float* __restrict__ grad_w_skip, // [5, HIDDEN] (+=) + float* __restrict__ grad_b_skip, // [5] (+=) + float* __restrict__ grad_h // [B, HIDDEN] (overwrite + carry) +) { + int tid = threadIdx.x; + if (tid >= HEAD_MID_H) return; + + __shared__ float s_lambda[N_HORIZONS_H]; + if (tid < N_HORIZONS_H) { + const float l = lambda[tid]; + s_lambda[tid] = (l > 0.0f) ? l : 1.0f; // sentinel: zero buffer → 1.0 + } + + __shared__ float s_d_skip[N_HORIZONS_H]; // = d_logit (skip-path grad) + __shared__ float s_d_main[N_HORIZONS_H]; + __shared__ float s_d_gate_lin[N_HORIZONS_H]; + __shared__ float s_d_eta1[N_HORIZONS_H * HEAD_MID_H]; // d_eta_1 = d_z2 + __shared__ float s_d_eta2[N_HORIZONS_H * HEAD_MID_H]; // d_eta_2 = d_a1 + __shared__ float s_d_z1[N_HORIZONS_H * HEAD_MID_H]; + __syncthreads(); + + for (int bi = 0; bi < n_batch; ++bi) { + const int m = tid; + const float* h_row = h + (long long)bi * HIDDEN_H; + + // Pass 1: per-horizon scalar grads. Threads 0..4 only. + if (tid < N_HORIZONS_H) { + const int k = tid; + const float p = probs[(long long)bi * N_HORIZONS_H + k]; + const float dp = grad_probs[(long long)bi * N_HORIZONS_H + k]; + const float d_logit = dp * p * (1.0f - p); + + const float gl = gate_logit[(long long)bi * N_HORIZONS_H + k]; + const float sigmoid_gl = 1.0f / (1.0f + expf(-gl)); + const float mv = main_val[(long long)bi * N_HORIZONS_H + k]; + + const float d_skip = d_logit; + const float d_main = d_logit * sigmoid_gl; + const float d_gate_a = d_logit * mv; + const float d_gate_lin = d_gate_a * sigmoid_gl * (1.0f - sigmoid_gl); + + s_d_skip[k] = d_skip; + s_d_main[k] = d_main; + s_d_gate_lin[k] = d_gate_lin; + + grad_b_skip[k] += d_skip; + grad_b_main[k] += d_main; + grad_b_gate[k] += d_gate_lin; + } + __syncthreads(); + + // Pass 2: per-(k, m_out=m) compute d_eta_1[k, m] and accumulate + // grad_w_main[k, m] / grad_w_gate[k, m]. Single-writer per (k, m). + #pragma unroll + for (int k = 0; k < N_HORIZONS_H; ++k) { + const float d_main_k = s_d_main[k]; + const float d_gate_lin_k = s_d_gate_lin[k]; + const float w_main_km = w_main[k * HEAD_MID_H + m]; + const float w_gate_km = w_gate[k * HEAD_MID_H + m]; + const float z2_km = z2[((long long)bi * N_HORIZONS_H + k) * HEAD_MID_H + m]; + + const float d_eta1_km = d_main_k * w_main_km + d_gate_lin_k * w_gate_km; + s_d_eta1[k * HEAD_MID_H + m] = d_eta1_km; + + grad_w_main[k * HEAD_MID_H + m] += d_main_k * z2_km; + grad_w_gate[k * HEAD_MID_H + m] += d_gate_lin_k * z2_km; + + // grad_b2[k, m] += d_eta_1[k, m] -- thread m sole writer of row m. + grad_b2[k * HEAD_MID_H + m] += d_eta1_km; + } + __syncthreads(); + + // Pass 3: through W2. + // grad_w2[k, m_out=m, m_in] += d_eta_1[k, m] * eta_2[k, m_in] + // thread m loops over m_in and is sole writer of row m. + // d_eta_2[k, m_in=m] = sum_{m_out} d_eta_1[k, m_out] * W2[k, m_out, m_in=m] + // thread m owns column m. + #pragma unroll + for (int k = 0; k < N_HORIZONS_H; ++k) { + const float d_eta1_km = s_d_eta1[k * HEAD_MID_H + m]; + // grad_w2[k, m_out=m, m_in=n] += d_eta_1[k, m] * eta_2[k, n] + for (int n = 0; n < HEAD_MID_H; ++n) { + const float a1_kn = a1[((long long)bi * N_HORIZONS_H + k) * HEAD_MID_H + n]; + grad_w2[((long long)(k * HEAD_MID_H) + m) * HEAD_MID_H + n] += d_eta1_km * a1_kn; + } + + // d_eta_2[k, m_in=m] = sum_{m_out} d_eta_1[k, m_out] * w2[k, m_out, m] + float d_eta2_km = 0.0f; + for (int m_out = 0; m_out < HEAD_MID_H; ++m_out) { + d_eta2_km += s_d_eta1[k * HEAD_MID_H + m_out] + * w2[((long long)(k * HEAD_MID_H) + m_out) * HEAD_MID_H + m]; + } + s_d_eta2[k * HEAD_MID_H + m] = d_eta2_km; + } + __syncthreads(); + + // Pass 4: through GELU. Thread m owns col m of d_z1. + // d_z1[k, m] = d_eta_2[k, m] * gelu_prime(z1[k, m]) + // grad_b1[k, m] += d_z1[k, m] + #pragma unroll + for (int k = 0; k < N_HORIZONS_H; ++k) { + const float d_eta2_km = s_d_eta2[k * HEAD_MID_H + m]; + const float z1_km = z1[((long long)bi * N_HORIZONS_H + k) * HEAD_MID_H + m]; + const float d_z1_km = d_eta2_km * gelu_prime(z1_km); + s_d_z1[k * HEAD_MID_H + m] = d_z1_km; + grad_b1[k * HEAD_MID_H + m] += d_z1_km; + } + __syncthreads(); + + // Pass 5: grad_w1[k, m, i] += d_z1[k, m] * h[i] + // Thread m owns row m for all k, i (HIDDEN cols sequential). + for (int i = 0; i < HIDDEN_H; ++i) { + const float h_bi_i = h_row[i]; + #pragma unroll + for (int k = 0; k < N_HORIZONS_H; ++k) { + grad_w1[((long long)(k * HEAD_MID_H) + m) * HIDDEN_H + i] += + s_d_z1[k * HEAD_MID_H + m] * h_bi_i; + } + } + + // Pass 6: trunk grad_w_skip + grad_h. + // grad_w_skip[k, i] += d_skip[k] * h[i] + // grad_h[bi, i] = sum_k lambda[k] * ( + // d_skip[k] * w_skip[k, i] # skip path + // + sum_m d_z1[k, m] * w1[k, m, i] # main path + // ) + carry[i] + // + // Tile i over 2 strides of HEAD_MID. Thread tid handles i=tid and i=tid+HEAD_MID. + #pragma unroll + for (int tile = 0; tile < HIDDEN_H / HEAD_MID_H; ++tile) { + const int i = tile * HEAD_MID_H + tid; + const float h_bi_i = h_row[i]; + float acc_grad_h = 0.0f; + #pragma unroll + for (int k = 0; k < N_HORIZONS_H; ++k) { + grad_w_skip[k * HIDDEN_H + i] += s_d_skip[k] * h_bi_i; + + float contrib = s_d_skip[k] * w_skip[k * HIDDEN_H + i]; + #pragma unroll + for (int mm = 0; mm < HEAD_MID_H; ++mm) { + contrib += s_d_z1[k * HEAD_MID_H + mm] + * w1[((long long)(k * HEAD_MID_H) + mm) * HIDDEN_H + i]; + } + acc_grad_h += s_lambda[k] * contrib; + } + const float carry = (grad_h_carry != nullptr) + ? grad_h_carry[(long long)bi * HIDDEN_H + i] : 0.0f; + grad_h[(long long)bi * HIDDEN_H + i] = acc_grad_h + carry; + } + __syncthreads(); + } +}