feat(ml-alpha): TFT GRN forward+backward kernels for multi-horizon heads (Phase 1.7a)

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 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-17 21:47:09 +02:00
parent 010445b5df
commit 5e23005dea
2 changed files with 336 additions and 4 deletions

View File

@@ -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");

View File

@@ -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();
}
}