perf(ml-alpha): block-per-batch GRN bwd refactor (Phase B commit 2)

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 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-17 23:53:43 +02:00
parent 494a2e4827
commit 5c2c3b65a8
2 changed files with 256 additions and 170 deletions

View File

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

View File

@@ -385,6 +385,17 @@ pub struct PerceptionTrainer {
cfc_grad_w_rec_scratch_d: CudaSlice<f32>, // [B, n_hid, n_hid]
cfc_grad_b_scratch_d: CudaSlice<f32>, // [B, n_hid]
cfc_grad_tau_scratch_d: CudaSlice<f32>, // [B, n_hid]
// GRN per-batch grad scratch (Phase B commit 2).
grn_grad_w1_scratch_d: CudaSlice<f32>, // [B, 5, HEAD_MID, HIDDEN]
grn_grad_b1_scratch_d: CudaSlice<f32>, // [B, 5, HEAD_MID]
grn_grad_w2_scratch_d: CudaSlice<f32>, // [B, 5, HEAD_MID, HEAD_MID]
grn_grad_b2_scratch_d: CudaSlice<f32>, // [B, 5, HEAD_MID]
grn_grad_w_gate_scratch_d: CudaSlice<f32>, // [B, 5, HEAD_MID]
grn_grad_b_gate_scratch_d: CudaSlice<f32>, // [B, 5]
grn_grad_w_main_scratch_d: CudaSlice<f32>, // [B, 5, HEAD_MID]
grn_grad_b_main_scratch_d: CudaSlice<f32>, // [B, 5]
grn_grad_w_skip_scratch_d: CudaSlice<f32>, // [B, 5, HIDDEN]
grn_grad_b_skip_scratch_d: CudaSlice<f32>, // [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::<f32>(
cfg.n_batch * N_HORIZONS * HEAD_MID_DIM * HIDDEN_DIM)?;
let grn_grad_b1_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?;
let grn_grad_w2_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS * HEAD_MID_DIM * HEAD_MID_DIM)?;
let grn_grad_b2_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?;
let grn_grad_w_gate_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?;
let grn_grad_b_gate_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS)?;
let grn_grad_w_main_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?;
let grn_grad_b_main_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS)?;
let grn_grad_w_skip_scratch_d = stream.alloc_zeros::<f32>(
cfg.n_batch * N_HORIZONS * HIDDEN_DIM)?;
let grn_grad_b_skip_scratch_d = stream.alloc_zeros::<f32>(
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::<f32>(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 +