From 5c2c3b65a8ca080f532c46d09ec8c816050220c8 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sun, 17 May 2026 23:53:43 +0200 Subject: [PATCH] perf(ml-alpha): block-per-batch GRN bwd refactor (Phase B commit 2) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit multi_horizon_heads_grn_bwd_batched refactored from grid=(1,1,1) to grid=(B,1,1). Removes the single-SM bottleneck on the second-most-called K-loop kernel (64×/step like cfc_bwd). Adds 10 per-batch grad scratch buffers (one per GRN param tensor) + 10 reduce_axis0 launches collapsing B → final grad after the K-loop: grn_grad_w1_scratch_d [B, 5, HEAD_MID, HIDDEN] grn_grad_b1_scratch_d [B, 5, HEAD_MID] grn_grad_w2_scratch_d [B, 5, HEAD_MID, HEAD_MID] grn_grad_b2_scratch_d [B, 5, HEAD_MID] grn_grad_w_gate_scratch_d [B, 5, HEAD_MID] grn_grad_b_gate_scratch_d [B, 5] grn_grad_w_main_scratch_d [B, 5, HEAD_MID] grn_grad_b_main_scratch_d [B, 5] grn_grad_w_skip_scratch_d [B, 5, HIDDEN] grn_grad_b_skip_scratch_d [B, 5] Total: ~8 MB scratch at B=32. All 9 perception_overfit smokes pass (including stacked_trainer_loss_ shrinks_at_batch_32 which exercises the cross-batch reducer path on both cfc and GRN grads). Co-Authored-By: Claude Opus 4.7 --- crates/ml-alpha/cuda/multi_horizon_heads.cu | 294 ++++++++++---------- crates/ml-alpha/src/trainer/perception.rs | 132 +++++++-- 2 files changed, 256 insertions(+), 170 deletions(-) diff --git a/crates/ml-alpha/cuda/multi_horizon_heads.cu b/crates/ml-alpha/cuda/multi_horizon_heads.cu index 2d89d89c0..469864b33 100644 --- a/crates/ml-alpha/cuda/multi_horizon_heads.cu +++ b/crates/ml-alpha/cuda/multi_horizon_heads.cu @@ -579,37 +579,49 @@ extern "C" __global__ void multi_horizon_heads_grn_fwd_batched( // Thread tid (with tid < 5) owns horizon-scalar grads. // Trunk grad_h tiles i over 2 strides of HEAD_MID. +// Block-per-batch GRN bwd (Phase B). +// grid=(n_batch, 1, 1) block=(HEAD_MID, 1, 1) +// +// Param-grad scratch tensors hold per-batch slices [B, ...]. Thread +// (bi, m) is sole writer to its slice for every K-iteration; the +// K-loop's 64 invocations accumulate via += within (bi, ..., m). +// Caller zeroes scratch once per training step; reduce_axis0 collapses +// B → final grad after the K-loop. +// +// grad_h is per-batch indexed; block bi is sole writer to its +// [bi, :] slice (overwrite + carry). 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] + 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) + float* __restrict__ grad_w1_scratch, // [B, 5, HEAD_MID, HIDDEN] (+=) + float* __restrict__ grad_b1_scratch, // [B, 5, HEAD_MID] (+=) + float* __restrict__ grad_w2_scratch, // [B, 5, HEAD_MID, HEAD_MID] (+=) + float* __restrict__ grad_b2_scratch, // [B, 5, HEAD_MID] (+=) + float* __restrict__ grad_w_gate_scratch, // [B, 5, HEAD_MID] (+=) + float* __restrict__ grad_b_gate_scratch, // [B, 5] (+=) + float* __restrict__ grad_w_main_scratch, // [B, 5, HEAD_MID] (+=) + float* __restrict__ grad_b_main_scratch, // [B, 5] (+=) + float* __restrict__ grad_w_skip_scratch, // [B, 5, HIDDEN] (+=) + float* __restrict__ grad_b_skip_scratch, // [B, 5] (+=) + float* __restrict__ grad_h // [B, HIDDEN] (overwrite + carry) ) { + int bi = blockIdx.x; int tid = threadIdx.x; - if (tid >= HEAD_MID_H) return; + if (bi >= n_batch || tid >= HEAD_MID_H) return; __shared__ float s_lambda[N_HORIZONS_H]; if (tid < N_HORIZONS_H) { @@ -620,139 +632,139 @@ extern "C" __global__ void multi_horizon_heads_grn_bwd_batched( __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_eta1[N_HORIZONS_H * HEAD_MID_H]; + __shared__ float s_d_eta2[N_HORIZONS_H * HEAD_MID_H]; __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; + 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); + // Per-batch base offsets into scratch tensors. + const long long bi_5 = (long long)bi * N_HORIZONS_H; + const long long bi_5_m = bi_5 * HEAD_MID_H; + const long long bi_5_h = bi_5 * HIDDEN_H; - 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]; + // 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 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); + 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]; - s_d_skip[k] = d_skip; - s_d_main[k] = d_main; - s_d_gate_lin[k] = d_gate_lin; + 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); - grad_b_skip[k] += d_skip; - grad_b_main[k] += d_main; - grad_b_gate[k] += d_gate_lin; + s_d_skip[k] = d_skip; + s_d_main[k] = d_main; + s_d_gate_lin[k] = d_gate_lin; + + // Per-batch scratch writes for horizon-scalar biases. + grad_b_skip_scratch[bi_5 + k] += d_skip; + grad_b_main_scratch[bi_5 + k] += d_main; + grad_b_gate_scratch[bi_5 + k] += d_gate_lin; + } + __syncthreads(); + + // Pass 2: per-(k, m_out=m) compute d_eta_1[k, m] + grad_w_{main,gate}. + #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_scratch[bi_5_m + (long long)k * HEAD_MID_H + m] += d_main_k * z2_km; + grad_w_gate_scratch[bi_5_m + (long long)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_scratch[bi_5_m + (long long)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] + // d_eta_2[k, m_in=m] = sum_{m_out} d_eta_1[k, m_out] * W2[k, m_out, m_in=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]; + 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_scratch[bi_5_m * HEAD_MID_H + + ((long long)k * HEAD_MID_H + m) * HEAD_MID_H + n] + += d_eta1_km * a1_kn; } - __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). + 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. + // 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_scratch[bi_5_m + (long long)k * HEAD_MID_H + m] += d_z1_km; + } + __syncthreads(); + + // Pass 5: grad_w1[k, m, i] += d_z1[k, m] * h[i] (per-batch scratch). + 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) { - 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; + grad_w1_scratch[bi_5_m * HIDDEN_H + + ((long long)k * HEAD_MID_H + m) * HIDDEN_H + i] + += s_d_z1[k * HEAD_MID_H + m] * h_bi_i; } - __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. + // Pass 6: trunk grad_w_skip + grad_h. + // grad_w_skip[k, i] += d_skip[k] * h[i] (per-batch scratch) + // 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) { - 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; - } + grad_w_skip_scratch[bi_5_h + (long long)k * HIDDEN_H + i] + += s_d_skip[k] * h_bi_i; - // 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]; + float contrib = s_d_skip[k] * w_skip[k * HIDDEN_H + 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; + 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; } - - // 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(); + 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; } } diff --git a/crates/ml-alpha/src/trainer/perception.rs b/crates/ml-alpha/src/trainer/perception.rs index dfe2deec0..12d1e14d2 100644 --- a/crates/ml-alpha/src/trainer/perception.rs +++ b/crates/ml-alpha/src/trainer/perception.rs @@ -385,6 +385,17 @@ pub struct PerceptionTrainer { cfc_grad_w_rec_scratch_d: CudaSlice, // [B, n_hid, n_hid] cfc_grad_b_scratch_d: CudaSlice, // [B, n_hid] cfc_grad_tau_scratch_d: CudaSlice, // [B, n_hid] + // GRN per-batch grad scratch (Phase B commit 2). + grn_grad_w1_scratch_d: CudaSlice, // [B, 5, HEAD_MID, HIDDEN] + grn_grad_b1_scratch_d: CudaSlice, // [B, 5, HEAD_MID] + grn_grad_w2_scratch_d: CudaSlice, // [B, 5, HEAD_MID, HEAD_MID] + grn_grad_b2_scratch_d: CudaSlice, // [B, 5, HEAD_MID] + grn_grad_w_gate_scratch_d: CudaSlice, // [B, 5, HEAD_MID] + grn_grad_b_gate_scratch_d: CudaSlice, // [B, 5] + grn_grad_w_main_scratch_d: CudaSlice, // [B, 5, HEAD_MID] + grn_grad_b_main_scratch_d: CudaSlice, // [B, 5] + grn_grad_w_skip_scratch_d: CudaSlice, // [B, 5, HIDDEN] + grn_grad_b_skip_scratch_d: CudaSlice, // [B, 5] /// Cross-batch reducer kernel: `[B, N] → [N]` via block tree-reduce. /// Used for every per-batch grad scratch in the refactored bwd path. reduce_axis0_fn: CudaFunction, @@ -703,6 +714,28 @@ impl PerceptionTrainer { let mut opt_heads_b_skip = AdamW::new(dev, N_HORIZONS, cfg.lr_cfc)?; opt_heads_b_skip.wd = 0.0; + // Phase B: GRN per-batch grad scratch. + let grn_grad_w1_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS * HEAD_MID_DIM * HIDDEN_DIM)?; + let grn_grad_b1_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?; + let grn_grad_w2_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS * HEAD_MID_DIM * HEAD_MID_DIM)?; + let grn_grad_b2_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?; + let grn_grad_w_gate_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?; + let grn_grad_b_gate_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS)?; + let grn_grad_w_main_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?; + let grn_grad_b_main_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS)?; + let grn_grad_w_skip_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS * HIDDEN_DIM)?; + let grn_grad_b_skip_scratch_d = stream.alloc_zeros::( + cfg.n_batch * N_HORIZONS)?; + // LayerNorm parameters: gain initialised to 1.0, bias to 0.0. // Adam wd=0 — LN params are scale/shift, never penalised. // Two instances: LN_a between m1 and m2, LN_b between m2 and CfC. @@ -818,6 +851,16 @@ impl PerceptionTrainer { cfc_grad_w_rec_scratch_d, cfc_grad_b_scratch_d, cfc_grad_tau_scratch_d, + grn_grad_w1_scratch_d, + grn_grad_b1_scratch_d, + grn_grad_w2_scratch_d, + grn_grad_b2_scratch_d, + grn_grad_w_gate_scratch_d, + grn_grad_b_gate_scratch_d, + grn_grad_w_main_scratch_d, + grn_grad_b_main_scratch_d, + grn_grad_w_skip_scratch_d, + grn_grad_b_skip_scratch_d, reduce_axis0_fn, _reduce_module: reduce_module, loss_ema_d: stream.alloc_zeros::(N_HORIZONS)?, @@ -1414,27 +1457,29 @@ impl PerceptionTrainer { // Phase 3: attention pool Q grad accumulator. self.stream.memset_zeros(&mut self.grad_attn_q_d) .map_err(|e| anyhow::anyhow!("zero grad_attn_q: {e}"))?; - // GRN heads: 10 grad accumulators. - self.stream.memset_zeros(&mut self.grad_heads_w1_d) - .map_err(|e| anyhow::anyhow!("zero grad_heads_w1: {e}"))?; - self.stream.memset_zeros(&mut self.grad_heads_b1_d) - .map_err(|e| anyhow::anyhow!("zero grad_heads_b1: {e}"))?; - self.stream.memset_zeros(&mut self.grad_heads_w2_d) - .map_err(|e| anyhow::anyhow!("zero grad_heads_w2: {e}"))?; - self.stream.memset_zeros(&mut self.grad_heads_b2_d) - .map_err(|e| anyhow::anyhow!("zero grad_heads_b2: {e}"))?; - self.stream.memset_zeros(&mut self.grad_heads_w_gate_d) - .map_err(|e| anyhow::anyhow!("zero grad_heads_w_gate: {e}"))?; - self.stream.memset_zeros(&mut self.grad_heads_b_gate_d) - .map_err(|e| anyhow::anyhow!("zero grad_heads_b_gate: {e}"))?; - self.stream.memset_zeros(&mut self.grad_heads_w_main_d) - .map_err(|e| anyhow::anyhow!("zero grad_heads_w_main: {e}"))?; - self.stream.memset_zeros(&mut self.grad_heads_b_main_d) - .map_err(|e| anyhow::anyhow!("zero grad_heads_b_main: {e}"))?; - self.stream.memset_zeros(&mut self.grad_heads_w_skip_d) - .map_err(|e| anyhow::anyhow!("zero grad_heads_w_skip: {e}"))?; - self.stream.memset_zeros(&mut self.grad_heads_b_skip_d) - .map_err(|e| anyhow::anyhow!("zero grad_heads_b_skip: {e}"))?; + // GRN per-batch grad scratch (Phase B commit 2): zero ONCE per + // step; K-loop bwd accumulates into these, then reduce_axis0 + // collapses → final grad buffers (OVERWRITE) after the K-loop. + self.stream.memset_zeros(&mut self.grn_grad_w1_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_w1_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_b1_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_b1_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_w2_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_w2_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_b2_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_b2_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_w_gate_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_w_gate_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_b_gate_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_b_gate_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_w_main_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_w_main_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_b_main_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_b_main_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_w_skip_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_w_skip_scratch: {e}"))?; + self.stream.memset_zeros(&mut self.grn_grad_b_skip_scratch_d) + .map_err(|e| anyhow::anyhow!("zero grn_grad_b_skip_scratch: {e}"))?; self.stream.memset_zeros(&mut self.zero_h_d) .map_err(|e| anyhow::anyhow!("zero zero_h: {e}"))?; @@ -1495,10 +1540,11 @@ impl PerceptionTrainer { block_dim: (HEAD_MID_DIM as u32, 1, 1), shared_mem_bytes: 0, }; - // GRN backward launch config: ONE block per launch (single-writer - // discipline), threads tile HEAD_MID. + // GRN backward launch config: block-per-batch (Phase B). One block + // per sample; threads tile HEAD_MID. Per-batch scratch + reducer + // collapses B → final grad after the K-loop. let cfg_grn_bwd = LaunchConfig { - grid_dim: (1, 1, 1), + grid_dim: (b_sz as u32, 1, 1), block_dim: (HEAD_MID_DIM as u32, 1, 1), shared_mem_bytes: 0, }; @@ -1678,6 +1724,9 @@ impl PerceptionTrainer { // ONLY the trunk gradient; per-horizon param grads are // unscaled (per pearl_adam_normalizes_loss_weights.md the // effective lever is the trunk gradient, not loss weight). + // Phase B: per-batch GRN grad scratch + per-batch grad_h_new. + // K-loop's 64 invocations accumulate into the scratch via +=; + // reduce_axis0 collapses → final grad after the K-loop. unsafe { let mut launch = self.stream.launch_builder(&self.heads_grn_bwd_fn); launch @@ -1691,11 +1740,11 @@ impl PerceptionTrainer { .arg(&self.grad_h_carry_d) .arg(&self.lambda_d) .arg(&n_batch_i) - .arg(&mut self.grad_heads_w1_d).arg(&mut self.grad_heads_b1_d) - .arg(&mut self.grad_heads_w2_d).arg(&mut self.grad_heads_b2_d) - .arg(&mut self.grad_heads_w_gate_d).arg(&mut self.grad_heads_b_gate_d) - .arg(&mut self.grad_heads_w_main_d).arg(&mut self.grad_heads_b_main_d) - .arg(&mut self.grad_heads_w_skip_d).arg(&mut self.grad_heads_b_skip_d) + .arg(&mut self.grn_grad_w1_scratch_d).arg(&mut self.grn_grad_b1_scratch_d) + .arg(&mut self.grn_grad_w2_scratch_d).arg(&mut self.grn_grad_b2_scratch_d) + .arg(&mut self.grn_grad_w_gate_scratch_d).arg(&mut self.grn_grad_b_gate_scratch_d) + .arg(&mut self.grn_grad_w_main_scratch_d).arg(&mut self.grn_grad_b_main_scratch_d) + .arg(&mut self.grn_grad_w_skip_scratch_d).arg(&mut self.grn_grad_b_skip_scratch_d) .arg(&mut self.grad_h_new_d); launch.launch(cfg_grn_bwd).context("heads GRN bwd k")?; } @@ -2009,6 +2058,31 @@ impl PerceptionTrainer { &mut self.grad_b_d, "reduce cfc_grad_b")?; reduce_at(HIDDEN_DIM, &self.cfc_grad_tau_scratch_d, &mut self.grad_tau_d, "reduce cfc_grad_tau")?; + + // GRN: 10 reducer launches (Phase B commit 2). + let nh = N_HORIZONS; + let mid = HEAD_MID_DIM; + let h = HIDDEN_DIM; + reduce_at(nh * mid * h, &self.grn_grad_w1_scratch_d, + &mut self.grad_heads_w1_d, "reduce grn_grad_w1")?; + reduce_at(nh * mid, &self.grn_grad_b1_scratch_d, + &mut self.grad_heads_b1_d, "reduce grn_grad_b1")?; + reduce_at(nh * mid * mid, &self.grn_grad_w2_scratch_d, + &mut self.grad_heads_w2_d, "reduce grn_grad_w2")?; + reduce_at(nh * mid, &self.grn_grad_b2_scratch_d, + &mut self.grad_heads_b2_d, "reduce grn_grad_b2")?; + reduce_at(nh * mid, &self.grn_grad_w_gate_scratch_d, + &mut self.grad_heads_w_gate_d, "reduce grn_grad_w_gate")?; + reduce_at(nh, &self.grn_grad_b_gate_scratch_d, + &mut self.grad_heads_b_gate_d, "reduce grn_grad_b_gate")?; + reduce_at(nh * mid, &self.grn_grad_w_main_scratch_d, + &mut self.grad_heads_w_main_d, "reduce grn_grad_w_main")?; + reduce_at(nh, &self.grn_grad_b_main_scratch_d, + &mut self.grad_heads_b_main_d, "reduce grn_grad_b_main")?; + reduce_at(nh * h, &self.grn_grad_w_skip_scratch_d, + &mut self.grad_heads_w_skip_d, "reduce grn_grad_w_skip")?; + reduce_at(nh, &self.grn_grad_b_skip_scratch_d, + &mut self.grad_heads_b_skip_d, "reduce grn_grad_b_skip")?; } // ── 9. Apply AdamW updates on all 17 param groups: CfC×4 +