feat(ml-alpha): wire TFT GRN heads into PerceptionTrainer (Phase 1.7)

Replaces the single-linear heads (5 × 128) with per-horizon TFT GRN:
  trunk h[B, 128]
    → W1 (HIDDEN→HEAD_MID) + GELU      = eta_2 [B, 5, HEAD_MID]
    → W2 (HEAD_MID→HEAD_MID)            = eta_1 [B, 5, HEAD_MID]
    → (W_gate, W_main) split            = (gate_lin, main) [B, 5]
    → W_skip (HIDDEN→1)                 = skip [B, 5]
    → skip + sigma(gate_lin) * main     = logit [B, 5]
    → sigma(logit)                      = p [B, 5]

10 trainable param tensors per horizon (5 weight + 5 bias), 10 AdamW
state objects (biases get wd=0), 6 per-K-position intermediate buffers
sized [K, B, 5, HEAD_MID] (z1, a1, z2) and [K, B, 5] (gate_logit, main,
logit). All saved during fwd, consumed by GRN bwd for the chain rule.

K-loop forward replaces `multi_horizon_heads_batched` with the new
`multi_horizon_heads_grn_fwd_batched` kernel. K-loop backward replaces
`multi_horizon_heads_backward_batched` with `multi_horizon_heads_grn_bwd_batched`.
ISV `lambda_d` is consumed by the GRN bwd kernel to scale ONLY the
trunk gradient (per pearl_adam_normalizes_loss_weights.md — the
effective lever is the trunk gradient, not loss weight).

Eval path also uses the GRN fwd kernel so eval sees the same
distribution downstream layers trained on (same pattern as Phase 1.6
LN wiring).

Xavier-uniform init: scale = sqrt(1 / fan_in) per tensor:
  W1, W_skip:        scale = sqrt(1/HIDDEN)        ≈ 0.088
  W2, W_gate, W_main: scale = sqrt(1/HEAD_MID)     ≈ 0.125

Parameter count per horizon: W1 (8192) + b1 (64) + W2 (4096) + b2 (64)
  + W_gate (64) + b_gate (1) + W_main (64) + b_main (1) + W_skip (128)
  + b_skip (1) = 12,675 params × 5 horizons = 63,375 head params total
  (vs single-linear's 5*(128+1)=645). Trunk grad dimensionality is
  unchanged.

Synthetic overfit smoke: loss 0.32 → 0.0000 by step 50 (faster than
single-linear's 0.37 → 0.0006 in 250 steps); all 7 perception_overfit
tests PASS. Confirms the full Mamba2 + LN + CfC + GRN backward chain
is wired correctly and all 17 AdamW optimisers move weights.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-17 21:54:52 +02:00
parent 5e23005dea
commit b6dfd89903
2 changed files with 291 additions and 60 deletions

View File

@@ -19,6 +19,10 @@ pub const N_HORIZONS: usize = 5;
pub const HORIZONS: [usize; N_HORIZONS] = [30, 100, 300, 1000, 6000];
pub const HIDDEN_DIM: usize = 128;
pub const PROJ_DIM: usize = 8;
/// GRN intermediate dimension (Phase 1.7). Two linear layers + GLU gate
/// + skip projection all operate at this width per horizon. Matches the
/// HEAD_MID_H define in `cuda/multi_horizon_heads.cu`.
pub const HEAD_MID_DIM: usize = 64;
#[derive(Clone, Debug)]
pub struct HeadsWeights {

View File

@@ -46,7 +46,7 @@ use rand::{Rng, SeedableRng};
use rand_chacha::ChaCha8Rng;
use crate::cfc::snap_features::{Mbp10RawInput, ES_TICK_SIZE, FEATURE_DIM, REGIME_DIM};
use crate::heads::{HIDDEN_DIM, N_HORIZONS};
use crate::heads::{HEAD_MID_DIM, HIDDEN_DIM, N_HORIZONS};
use crate::mamba2_block::{
Mamba2AdamW, Mamba2AdamWConfig, Mamba2BackwardGradsBuffers, Mamba2BackwardScratch,
Mamba2Block, Mamba2BlockConfig, Mamba2BlockForwardScratch,
@@ -138,8 +138,8 @@ pub struct PerceptionTrainer {
/// unbatched kernels (the inner per-thread loop runs once).
step_batched_fn: CudaFunction,
step_bwd_batched_fn: CudaFunction,
heads_batched_fn: CudaFunction,
heads_bwd_batched_fn: CudaFunction,
heads_grn_fwd_fn: CudaFunction,
heads_grn_bwd_fn: CudaFunction,
transpose_3d_fn: CudaFunction,
// Mamba2 encoder block + its optimizer
@@ -176,14 +176,35 @@ pub struct PerceptionTrainer {
pub w_rec_d: CudaSlice<f32>,
pub b_d: CudaSlice<f32>,
pub tau_d: CudaSlice<f32>,
pub heads_w_d: CudaSlice<f32>,
pub heads_b_d: CudaSlice<f32>,
// ── TFT GRN heads (Phase 1.7) ──
// Per-horizon Gated Residual Network: 2-layer GELU MLP body
// (eta_2 → eta_1) + GLU gate + skip-projection. Gives per-horizon
// "linear vs deeper-transform" specialisation, matching the
// regime-conditional alpha pattern (pearl_snapshot_alpha_is_regime_conditional).
pub heads_w1_d: CudaSlice<f32>, // [N_HORIZONS, HEAD_MID, HIDDEN]
pub heads_b1_d: CudaSlice<f32>, // [N_HORIZONS, HEAD_MID]
pub heads_w2_d: CudaSlice<f32>, // [N_HORIZONS, HEAD_MID, HEAD_MID]
pub heads_b2_d: CudaSlice<f32>, // [N_HORIZONS, HEAD_MID]
pub heads_w_gate_d: CudaSlice<f32>, // [N_HORIZONS, HEAD_MID]
pub heads_b_gate_d: CudaSlice<f32>, // [N_HORIZONS]
pub heads_w_main_d: CudaSlice<f32>, // [N_HORIZONS, HEAD_MID]
pub heads_b_main_d: CudaSlice<f32>, // [N_HORIZONS]
pub heads_w_skip_d: CudaSlice<f32>, // [N_HORIZONS, HIDDEN]
pub heads_b_skip_d: CudaSlice<f32>, // [N_HORIZONS]
pub opt_w_in: AdamW,
pub opt_w_rec: AdamW,
pub opt_b: AdamW,
pub opt_tau: AdamW,
pub opt_heads_w: AdamW,
pub opt_heads_b: AdamW,
pub opt_heads_w1: AdamW,
pub opt_heads_b1: AdamW,
pub opt_heads_w2: AdamW,
pub opt_heads_b2: AdamW,
pub opt_heads_w_gate: AdamW,
pub opt_heads_b_gate: AdamW,
pub opt_heads_w_main: AdamW,
pub opt_heads_b_main: AdamW,
pub opt_heads_w_skip: AdamW,
pub opt_heads_b_skip: AdamW,
// Pre-allocated BATCHED snap_feature staging — sized for max B*K
// snapshots. All staging buffers are MAPPED-PINNED (the only
@@ -316,8 +337,26 @@ pub struct PerceptionTrainer {
grad_w_rec_d: CudaSlice<f32>,
grad_b_d: CudaSlice<f32>,
grad_tau_d: CudaSlice<f32>,
grad_heads_w_d: CudaSlice<f32>,
grad_heads_b_d: CudaSlice<f32>,
grad_heads_w1_d: CudaSlice<f32>,
grad_heads_b1_d: CudaSlice<f32>,
grad_heads_w2_d: CudaSlice<f32>,
grad_heads_b2_d: CudaSlice<f32>,
grad_heads_w_gate_d: CudaSlice<f32>,
grad_heads_b_gate_d: CudaSlice<f32>,
grad_heads_w_main_d: CudaSlice<f32>,
grad_heads_b_main_d: CudaSlice<f32>,
grad_heads_w_skip_d: CudaSlice<f32>,
grad_heads_b_skip_d: CudaSlice<f32>,
// GRN forward intermediates saved for backward, per K position.
// Stored as [K, B, N_HORIZONS, HEAD_MID] in row-major; per-K slot
// is [B, N_HORIZONS, HEAD_MID] contiguous.
z1_per_k_d: CudaSlice<f32>,
a1_per_k_d: CudaSlice<f32>,
z2_per_k_d: CudaSlice<f32>,
// Per-K horizon-scalar intermediates: [K, B, N_HORIZONS].
gate_logit_per_k_d: CudaSlice<f32>,
main_per_k_d: CudaSlice<f32>,
logit_per_k_d: CudaSlice<f32>,
// Mapped-pinned staging for label uploads (the only remaining
// single-call host→device path; staging covers all K*B labels per
@@ -379,8 +418,12 @@ impl PerceptionTrainer {
let bce_fn = bce_module.load_function("bce_multi_horizon_forward_backward")?;
let step_batched_fn = step_module.load_function("cfc_step_batched")?;
let step_bwd_batched_fn = step_module.load_function("cfc_step_backward_batched")?;
let heads_batched_fn = heads_module.load_function("multi_horizon_heads_batched")?;
let heads_bwd_batched_fn = heads_module.load_function("multi_horizon_heads_backward_batched")?;
let heads_grn_fwd_fn = heads_module
.load_function("multi_horizon_heads_grn_fwd_batched")
.context("heads GRN fwd symbol")?;
let heads_grn_bwd_fn = heads_module
.load_function("multi_horizon_heads_grn_bwd_batched")
.context("heads GRN bwd symbol")?;
let transpose_3d_fn = step_module.load_function("transpose_3d_swap_01")?;
// Mamba2 encoder: in_dim=FEATURE_DIM (32), hidden=HIDDEN_DIM (128).
@@ -436,16 +479,41 @@ impl PerceptionTrainer {
(-4.6 + u * 11.5).exp()
})
.collect();
let head_scale = scale_rec;
let heads_w: Vec<f32> = (0..N_HORIZONS * n_hid).map(|_| r.gen_range(-head_scale..head_scale)).collect();
let heads_b: Vec<f32> = vec![0.0; N_HORIZONS];
// ── TFT GRN heads init ──
// Xavier-uniform per-tensor: scale = sqrt(1 / fan_in).
let scale_h_in = (1.0_f32 / n_hid as f32).sqrt(); // HIDDEN → HEAD_MID
let scale_mid_mid = (1.0_f32 / HEAD_MID_DIM as f32).sqrt(); // HEAD_MID → HEAD_MID / → 1
let scale_h_skip = (1.0_f32 / n_hid as f32).sqrt(); // HIDDEN → 1
let heads_w1: Vec<f32> = (0..N_HORIZONS * HEAD_MID_DIM * n_hid)
.map(|_| r.gen_range(-scale_h_in..scale_h_in)).collect();
let heads_b1: Vec<f32> = vec![0.0; N_HORIZONS * HEAD_MID_DIM];
let heads_w2: Vec<f32> = (0..N_HORIZONS * HEAD_MID_DIM * HEAD_MID_DIM)
.map(|_| r.gen_range(-scale_mid_mid..scale_mid_mid)).collect();
let heads_b2: Vec<f32> = vec![0.0; N_HORIZONS * HEAD_MID_DIM];
let heads_w_gate: Vec<f32> = (0..N_HORIZONS * HEAD_MID_DIM)
.map(|_| r.gen_range(-scale_mid_mid..scale_mid_mid)).collect();
let heads_b_gate: Vec<f32> = vec![0.0; N_HORIZONS];
let heads_w_main: Vec<f32> = (0..N_HORIZONS * HEAD_MID_DIM)
.map(|_| r.gen_range(-scale_mid_mid..scale_mid_mid)).collect();
let heads_b_main: Vec<f32> = vec![0.0; N_HORIZONS];
let heads_w_skip: Vec<f32> = (0..N_HORIZONS * n_hid)
.map(|_| r.gen_range(-scale_h_skip..scale_h_skip)).collect();
let heads_b_skip: Vec<f32> = vec![0.0; N_HORIZONS];
let w_in_d = upload(&stream, &w_in)?;
let w_rec_d = upload(&stream, &w_rec)?;
let b_d = upload(&stream, &b)?;
let tau_d = upload(&stream, &tau)?;
let heads_w_d = upload(&stream, &heads_w)?;
let heads_b_d = upload(&stream, &heads_b)?;
let heads_w1_d = upload(&stream, &heads_w1)?;
let heads_b1_d = upload(&stream, &heads_b1)?;
let heads_w2_d = upload(&stream, &heads_w2)?;
let heads_b2_d = upload(&stream, &heads_b2)?;
let heads_w_gate_d = upload(&stream, &heads_w_gate)?;
let heads_b_gate_d = upload(&stream, &heads_b_gate)?;
let heads_w_main_d = upload(&stream, &heads_w_main)?;
let heads_b_main_d = upload(&stream, &heads_b_main)?;
let heads_w_skip_d = upload(&stream, &heads_w_skip)?;
let heads_b_skip_d = upload(&stream, &heads_b_skip)?;
let opt_w_in = AdamW::new(dev, n_hid * n_in, cfg.lr_cfc)?;
let opt_w_rec = AdamW::new(dev, n_hid * n_hid, cfg.lr_cfc)?;
@@ -455,9 +523,23 @@ impl PerceptionTrainer {
// decay so the per-cell decay constants drift smoothly.
let mut opt_tau = AdamW::new(dev, n_hid, cfg.lr_cfc * 0.1)?;
opt_tau.wd = 0.0;
let opt_heads_w = AdamW::new(dev, N_HORIZONS * n_hid, cfg.lr_cfc)?;
let mut opt_heads_b = AdamW::new(dev, N_HORIZONS, cfg.lr_cfc)?;
opt_heads_b.wd = 0.0;
// GRN heads: 10 AdamW (5 weight + 5 bias). Biases get wd=0 per
// the existing pattern; weights keep AdamW default wd.
let opt_heads_w1 = AdamW::new(dev, N_HORIZONS * HEAD_MID_DIM * n_hid, cfg.lr_cfc)?;
let mut opt_heads_b1 = AdamW::new(dev, N_HORIZONS * HEAD_MID_DIM, cfg.lr_cfc)?;
opt_heads_b1.wd = 0.0;
let opt_heads_w2 = AdamW::new(dev, N_HORIZONS * HEAD_MID_DIM * HEAD_MID_DIM, cfg.lr_cfc)?;
let mut opt_heads_b2 = AdamW::new(dev, N_HORIZONS * HEAD_MID_DIM, cfg.lr_cfc)?;
opt_heads_b2.wd = 0.0;
let opt_heads_w_gate = AdamW::new(dev, N_HORIZONS * HEAD_MID_DIM, cfg.lr_cfc)?;
let mut opt_heads_b_gate = AdamW::new(dev, N_HORIZONS, cfg.lr_cfc)?;
opt_heads_b_gate.wd = 0.0;
let opt_heads_w_main = AdamW::new(dev, N_HORIZONS * HEAD_MID_DIM, cfg.lr_cfc)?;
let mut opt_heads_b_main = AdamW::new(dev, N_HORIZONS, cfg.lr_cfc)?;
opt_heads_b_main.wd = 0.0;
let opt_heads_w_skip = AdamW::new(dev, N_HORIZONS * n_hid, cfg.lr_cfc)?;
let mut opt_heads_b_skip = AdamW::new(dev, N_HORIZONS, cfg.lr_cfc)?;
opt_heads_b_skip.wd = 0.0;
// LayerNorm parameters: gain initialised to 1.0, bias to 0.0.
// Adam wd=0 — LN params are scale/shift, never penalised.
@@ -514,8 +596,24 @@ impl PerceptionTrainer {
grad_w_rec_d: stream.alloc_zeros::<f32>(n_hid * n_hid)?,
grad_b_d: stream.alloc_zeros::<f32>(n_hid)?,
grad_tau_d: stream.alloc_zeros::<f32>(n_hid)?,
grad_heads_w_d: stream.alloc_zeros::<f32>(N_HORIZONS * n_hid)?,
grad_heads_b_d: stream.alloc_zeros::<f32>(N_HORIZONS)?,
grad_heads_w1_d: stream.alloc_zeros::<f32>(N_HORIZONS * HEAD_MID_DIM * n_hid)?,
grad_heads_b1_d: stream.alloc_zeros::<f32>(N_HORIZONS * HEAD_MID_DIM)?,
grad_heads_w2_d: stream.alloc_zeros::<f32>(N_HORIZONS * HEAD_MID_DIM * HEAD_MID_DIM)?,
grad_heads_b2_d: stream.alloc_zeros::<f32>(N_HORIZONS * HEAD_MID_DIM)?,
grad_heads_w_gate_d: stream.alloc_zeros::<f32>(N_HORIZONS * HEAD_MID_DIM)?,
grad_heads_b_gate_d: stream.alloc_zeros::<f32>(N_HORIZONS)?,
grad_heads_w_main_d: stream.alloc_zeros::<f32>(N_HORIZONS * HEAD_MID_DIM)?,
grad_heads_b_main_d: stream.alloc_zeros::<f32>(N_HORIZONS)?,
grad_heads_w_skip_d: stream.alloc_zeros::<f32>(N_HORIZONS * n_hid)?,
grad_heads_b_skip_d: stream.alloc_zeros::<f32>(N_HORIZONS)?,
// GRN per-K intermediate buffers (sized for [K, B, 5, HEAD_MID]
// and [K, B, 5]) — the K-loop reads/writes per-K slot offsets.
z1_per_k_d: stream.alloc_zeros::<f32>(k * cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?,
a1_per_k_d: stream.alloc_zeros::<f32>(k * cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?,
z2_per_k_d: stream.alloc_zeros::<f32>(k * cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?,
gate_logit_per_k_d: stream.alloc_zeros::<f32>(k * cfg.n_batch * N_HORIZONS)?,
main_per_k_d: stream.alloc_zeros::<f32>(k * cfg.n_batch * N_HORIZONS)?,
logit_per_k_d: stream.alloc_zeros::<f32>(k * cfg.n_batch * N_HORIZONS)?,
stg_labels: unsafe { MappedF32Buffer::new(k * cfg.n_batch * N_HORIZONS) }.map_err(|e| anyhow::anyhow!("stg_labels: {e}"))?,
train_graph: None,
cublas_warmed: false,
@@ -559,8 +657,16 @@ impl PerceptionTrainer {
w_rec_d,
b_d,
tau_d,
heads_w_d,
heads_b_d,
heads_w1_d,
heads_b1_d,
heads_w2_d,
heads_b2_d,
heads_w_gate_d,
heads_b_gate_d,
heads_w_main_d,
heads_b_main_d,
heads_w_skip_d,
heads_b_skip_d,
stream,
_snap_module: snap_module,
_step_module: step_module,
@@ -570,8 +676,8 @@ impl PerceptionTrainer {
bce_fn,
step_batched_fn,
step_bwd_batched_fn,
heads_batched_fn,
heads_bwd_batched_fn,
heads_grn_fwd_fn,
heads_grn_bwd_fn,
transpose_3d_fn,
mamba2,
mamba2_adamw,
@@ -586,8 +692,16 @@ impl PerceptionTrainer {
opt_w_rec,
opt_b,
opt_tau,
opt_heads_w,
opt_heads_b,
opt_heads_w1,
opt_heads_b1,
opt_heads_w2,
opt_heads_b2,
opt_heads_w_gate,
opt_heads_b_gate,
opt_heads_w_main,
opt_heads_b_main,
opt_heads_w_skip,
opt_heads_b_skip,
})
}
@@ -598,8 +712,16 @@ impl PerceptionTrainer {
self.opt_w_rec.lr = lr;
self.opt_b.lr = lr;
self.opt_tau.lr = lr * 0.1;
self.opt_heads_w.lr = lr;
self.opt_heads_b.lr = lr;
self.opt_heads_w1.lr = lr;
self.opt_heads_b1.lr = lr;
self.opt_heads_w2.lr = lr;
self.opt_heads_b2.lr = lr;
self.opt_heads_w_gate.lr = lr;
self.opt_heads_b_gate.lr = lr;
self.opt_heads_w_main.lr = lr;
self.opt_heads_b_main.lr = lr;
self.opt_heads_w_skip.lr = lr;
self.opt_heads_b_skip.lr = lr;
}
/// Set Mamba2 AdamW learning rate.
@@ -956,10 +1078,27 @@ impl PerceptionTrainer {
.map_err(|e| anyhow::anyhow!("zero grad_b: {e}"))?;
self.stream.memset_zeros(&mut self.grad_tau_d)
.map_err(|e| anyhow::anyhow!("zero grad_tau: {e}"))?;
self.stream.memset_zeros(&mut self.grad_heads_w_d)
.map_err(|e| anyhow::anyhow!("zero grad_heads_w: {e}"))?;
self.stream.memset_zeros(&mut self.grad_heads_b_d)
.map_err(|e| anyhow::anyhow!("zero grad_heads_b: {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}"))?;
self.stream.memset_zeros(&mut self.zero_h_d)
.map_err(|e| anyhow::anyhow!("zero zero_h: {e}"))?;
@@ -1005,8 +1144,26 @@ impl PerceptionTrainer {
// Per-K slot strides for [K, B, H] / [K, B, N_HORIZONS] layout.
let kb_hid_bytes = b_sz * HIDDEN_DIM * std::mem::size_of::<f32>();
let kb_nh_bytes = b_sz * N_HORIZONS * std::mem::size_of::<f32>();
// GRN per-K intermediate stride: [B, N_HORIZONS, HEAD_MID].
let kb_nh_mid_bytes = b_sz * N_HORIZONS * HEAD_MID_DIM * std::mem::size_of::<f32>();
let n_batch_i = b_sz as i32;
// GRN forward launch config: block-per-sample, threads tile HEAD_MID.
let cfg_grn_fwd = LaunchConfig {
grid_dim: (b_sz as u32, 1, 1),
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.
let cfg_grn_bwd = LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (HEAD_MID_DIM as u32, 1, 1),
shared_mem_bytes: 0,
};
let _ = cfg_heads_batched_fwd; // GRN fwd uses cfg_grn_fwd now.
let _ = cfg_heads_bwd; // GRN bwd uses cfg_grn_bwd now.
// ── 4. Stream-ordered forward K loop. CfC is RECURRENT —
// h_old at step k for sample b IS h_new at step k-1 for
// the same b. Per-K slot is contiguous [B, H] in the
@@ -1023,6 +1180,12 @@ impl PerceptionTrainer {
};
let (h_per_k_base, _g_hpk) = self.h_new_per_k_d.device_ptr_mut(&self.stream);
let (probs_base, _g_probs) = self.probs_per_k_d.device_ptr_mut(&self.stream);
let (z1_base, _g_z1) = self.z1_per_k_d.device_ptr_mut(&self.stream);
let (a1_base, _g_a1) = self.a1_per_k_d.device_ptr_mut(&self.stream);
let (z2_base, _g_z2) = self.z2_per_k_d.device_ptr_mut(&self.stream);
let (gate_base, _g_gate) = self.gate_logit_per_k_d.device_ptr_mut(&self.stream);
let (main_base, _g_main) = self.main_per_k_d.device_ptr_mut(&self.stream);
let (logit_base, _g_logit) = self.logit_per_k_d.device_ptr_mut(&self.stream);
for k in 0..k_seq {
let h_new_k_ptr = h_per_k_base + (k * kb_hid_bytes) as u64;
@@ -1033,6 +1196,12 @@ impl PerceptionTrainer {
};
let x_k_ptr = henr_t_base + (k * kb_hid_bytes) as u64;
let probs_k_ptr = probs_base + (k * kb_nh_bytes) as u64;
let z1_k_ptr = z1_base + (k * kb_nh_mid_bytes) as u64;
let a1_k_ptr = a1_base + (k * kb_nh_mid_bytes) as u64;
let z2_k_ptr = z2_base + (k * kb_nh_mid_bytes) as u64;
let gate_k_ptr = gate_base + (k * kb_nh_bytes) as u64;
let main_k_ptr = main_base + (k * kb_nh_bytes) as u64;
let logit_k_ptr = logit_base + (k * kb_nh_bytes) as u64;
// cfc_step_batched(W, x[B*n_in], h_old[B*n_hid], …, → h_new[B*n_hid]).
unsafe {
@@ -1045,16 +1214,24 @@ impl PerceptionTrainer {
launch.launch(cfg_cfc).context("cfc fwd k batched")?;
}
// multi_horizon_heads_batched(W, b, h[B*128], B, → probs[B*5]).
// multi_horizon_heads_grn_fwd_batched(10 params + h[B*128] + n_batch
// → probs[B*5] + z1/a1/z2 [B,5,HEAD_MID] + gate_logit/main/logit [B,5]).
unsafe {
let mut launch = self.stream.launch_builder(&self.heads_batched_fn);
let mut launch = self.stream.launch_builder(&self.heads_grn_fwd_fn);
launch
.arg(&self.heads_w_d).arg(&self.heads_b_d)
.arg(&h_new_k_ptr).arg(&n_batch_i).arg(&probs_k_ptr);
launch.launch(cfg_heads_batched_fwd).context("heads fwd k batched")?;
.arg(&self.heads_w1_d).arg(&self.heads_b1_d)
.arg(&self.heads_w2_d).arg(&self.heads_b2_d)
.arg(&self.heads_w_gate_d).arg(&self.heads_b_gate_d)
.arg(&self.heads_w_main_d).arg(&self.heads_b_main_d)
.arg(&self.heads_w_skip_d).arg(&self.heads_b_skip_d)
.arg(&h_new_k_ptr).arg(&n_batch_i)
.arg(&probs_k_ptr)
.arg(&z1_k_ptr).arg(&a1_k_ptr).arg(&z2_k_ptr)
.arg(&gate_k_ptr).arg(&main_k_ptr).arg(&logit_k_ptr);
launch.launch(cfg_grn_fwd).context("heads GRN fwd k")?;
}
}
drop((_g_hpk, _g_probs));
drop((_g_hpk, _g_probs, _g_z1, _g_a1, _g_z2, _g_gate, _g_main, _g_logit));
// ── 5. Fused multi-horizon BCE over the full [K*B, N_HORIZONS]
// grid. The kernel doesn't distinguish position-vs-batch
@@ -1117,12 +1294,22 @@ impl PerceptionTrainer {
let (probs_base_bwd, _g_probs_bwd) = self.probs_per_k_d.device_ptr_mut(&self.stream);
let (gprobs_base_bwd, _g_gprobs_bwd) = self.grad_probs_per_k_d.device_ptr_mut(&self.stream);
let (grad_henr_t_base, _g_ghen_t_mut) = self.grad_h_enriched_seq_t_d.data_mut().device_ptr_mut(&self.stream);
let (z1_base_bwd, _g_z1_bwd) = self.z1_per_k_d.device_ptr_mut(&self.stream);
let (a1_base_bwd, _g_a1_bwd) = self.a1_per_k_d.device_ptr_mut(&self.stream);
let (z2_base_bwd, _g_z2_bwd) = self.z2_per_k_d.device_ptr_mut(&self.stream);
let (gate_base_bwd, _g_gate_bwd) = self.gate_logit_per_k_d.device_ptr_mut(&self.stream);
let (main_base_bwd, _g_main_bwd) = self.main_per_k_d.device_ptr_mut(&self.stream);
for k in (0..k_seq).rev() {
let h_new_k_ptr = h_per_k_base_bwd + (k * kb_hid_bytes) as u64;
let x_k_ptr = henr_t_base_bwd + (k * kb_hid_bytes) as u64;
let probs_k_ptr = probs_base_bwd + (k * kb_nh_bytes) as u64;
let gprobs_k_ptr = gprobs_base_bwd + (k * kb_nh_bytes) as u64;
let z1_k_ptr = z1_base_bwd + (k * kb_nh_mid_bytes) as u64;
let a1_k_ptr = a1_base_bwd + (k * kb_nh_mid_bytes) as u64;
let z2_k_ptr = z2_base_bwd + (k * kb_nh_mid_bytes) as u64;
let gate_k_ptr = gate_base_bwd + (k * kb_nh_bytes) as u64;
let main_k_ptr = main_base_bwd + (k * kb_nh_bytes) as u64;
let h_old_k_ptr = if k == 0 {
zero_h_ptr_bwd
} else {
@@ -1130,23 +1317,31 @@ impl PerceptionTrainer {
};
let grad_henr_k_ptr = grad_henr_t_base + (k * kb_hid_bytes) as u64;
// heads_bwd_batched writes grad_h_new[B, 128] = heads-side + carry.
// `lambda_d` scales ONLY the trunk gradient — the heads'
// own weight/bias gradients are unscaled.
// GRN backward: chain rule through sigmoid → (skip + sigmoid(gate)*main)
// → GLU → linear → GELU → linear → trunk. `lambda_d` scales
// 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).
unsafe {
let mut launch = self.stream.launch_builder(&self.heads_bwd_batched_fn);
let mut launch = self.stream.launch_builder(&self.heads_grn_bwd_fn);
launch
.arg(&self.heads_w_d)
.arg(&probs_k_ptr)
.arg(&self.heads_w1_d).arg(&self.heads_w2_d)
.arg(&self.heads_w_gate_d).arg(&self.heads_w_main_d)
.arg(&self.heads_w_skip_d)
.arg(&probs_k_ptr).arg(&gprobs_k_ptr)
.arg(&z1_k_ptr).arg(&a1_k_ptr).arg(&z2_k_ptr)
.arg(&gate_k_ptr).arg(&main_k_ptr)
.arg(&h_new_k_ptr)
.arg(&gprobs_k_ptr)
.arg(&self.grad_h_carry_d)
.arg(&self.lambda_d)
.arg(&n_batch_i)
.arg(&mut self.grad_heads_w_d)
.arg(&mut self.grad_heads_b_d)
.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.grad_h_new_d);
launch.launch(cfg_heads_bwd).context("heads bwd k batched")?;
launch.launch(cfg_grn_bwd).context("heads GRN bwd k")?;
}
// cfc_step_bwd_batched: writes grad_h_carry (= grad_h_old_out for
@@ -1165,7 +1360,8 @@ impl PerceptionTrainer {
launch.launch(cfg_cfc_bwd).context("cfc bwd k batched")?;
}
}
drop((_g_hpk_bwd, _g_probs_bwd, _g_gprobs_bwd, _g_ghen_t_mut));
drop((_g_hpk_bwd, _g_probs_bwd, _g_gprobs_bwd, _g_ghen_t_mut,
_g_z1_bwd, _g_a1_bwd, _g_z2_bwd, _g_gate_bwd, _g_main_bwd));
// ── 7. Stream stays async — bwd loop kernels, transposes, and
// Mamba2 bwd are all sequential on the same stream, no
@@ -1270,13 +1466,22 @@ impl PerceptionTrainer {
)
.context("mamba2 backward_from_h_enriched_seq_full_into")?;
// ── 9. Apply AdamW updates on all 9 param groups (CfC×4 + heads×2 + LN×2 + Mamba2 grouped).
// ── 9. Apply AdamW updates on all 17 param groups: CfC×4 +
// GRN heads×10 + LN×2 + Mamba2 grouped.
self.opt_w_in.step(&mut self.w_in_d, &self.grad_w_in_d)?;
self.opt_w_rec.step(&mut self.w_rec_d, &self.grad_w_rec_d)?;
self.opt_b.step(&mut self.b_d, &self.grad_b_d)?;
self.opt_tau.step(&mut self.tau_d, &self.grad_tau_d)?;
self.opt_heads_w.step(&mut self.heads_w_d, &self.grad_heads_w_d)?;
self.opt_heads_b.step(&mut self.heads_b_d, &self.grad_heads_b_d)?;
self.opt_heads_w1.step(&mut self.heads_w1_d, &self.grad_heads_w1_d)?;
self.opt_heads_b1.step(&mut self.heads_b1_d, &self.grad_heads_b1_d)?;
self.opt_heads_w2.step(&mut self.heads_w2_d, &self.grad_heads_w2_d)?;
self.opt_heads_b2.step(&mut self.heads_b2_d, &self.grad_heads_b2_d)?;
self.opt_heads_w_gate.step(&mut self.heads_w_gate_d, &self.grad_heads_w_gate_d)?;
self.opt_heads_b_gate.step(&mut self.heads_b_gate_d, &self.grad_heads_b_gate_d)?;
self.opt_heads_w_main.step(&mut self.heads_w_main_d, &self.grad_heads_w_main_d)?;
self.opt_heads_b_main.step(&mut self.heads_b_main_d, &self.grad_heads_b_main_d)?;
self.opt_heads_w_skip.step(&mut self.heads_w_skip_d, &self.grad_heads_w_skip_d)?;
self.opt_heads_b_skip.step(&mut self.heads_b_skip_d, &self.grad_heads_b_skip_d)?;
self.opt_ln_gain.step(&mut self.ln_gain_d, &self.grad_ln_gain_d)?;
self.opt_ln_bias.step(&mut self.ln_bias_d, &self.grad_ln_bias_d)?;
// GPU-resident grad-clip + AdamW (zero host roundtrips).
@@ -1512,11 +1717,14 @@ impl PerceptionTrainer {
let cfg_cfc = LaunchConfig {
grid_dim: (grid_dim, 1, 1), block_dim: (block_dim, 1, 1), shared_mem_bytes: 0,
};
let cfg_heads = LaunchConfig {
grid_dim: (1, 1, 1), block_dim: (128, 1, 1), shared_mem_bytes: 0,
let cfg_grn_fwd = LaunchConfig {
grid_dim: (b_sz as u32, 1, 1),
block_dim: (HEAD_MID_DIM as u32, 1, 1),
shared_mem_bytes: 0,
};
let kb_hid_bytes = b_sz * HIDDEN_DIM * std::mem::size_of::<f32>();
let kb_nh_bytes = b_sz * N_HORIZONS * std::mem::size_of::<f32>();
let kb_nh_mid_bytes = b_sz * N_HORIZONS * HEAD_MID_DIM * std::mem::size_of::<f32>();
let n_batch_i = b_sz as i32;
let zero_h_ptr = {
let (p, _g) = self.zero_h_d.device_ptr(&self.stream);
@@ -1528,6 +1736,12 @@ impl PerceptionTrainer {
};
let (h_per_k_base, _g_hpk) = self.h_new_per_k_d.device_ptr_mut(&self.stream);
let (probs_base, _g_probs) = self.probs_per_k_d.device_ptr_mut(&self.stream);
let (z1_base, _g_z1) = self.z1_per_k_d.device_ptr_mut(&self.stream);
let (a1_base, _g_a1) = self.a1_per_k_d.device_ptr_mut(&self.stream);
let (z2_base, _g_z2) = self.z2_per_k_d.device_ptr_mut(&self.stream);
let (gate_base, _g_gate) = self.gate_logit_per_k_d.device_ptr_mut(&self.stream);
let (main_base, _g_main) = self.main_per_k_d.device_ptr_mut(&self.stream);
let (logit_base, _g_logit) = self.logit_per_k_d.device_ptr_mut(&self.stream);
for k in 0..k_seq {
let h_new_k_ptr = h_per_k_base + (k * kb_hid_bytes) as u64;
let h_old_k_ptr = if k == 0 {
@@ -1537,6 +1751,12 @@ impl PerceptionTrainer {
};
let x_k_ptr = henr_t_base + (k * kb_hid_bytes) as u64;
let probs_k_ptr = probs_base + (k * kb_nh_bytes) as u64;
let z1_k_ptr = z1_base + (k * kb_nh_mid_bytes) as u64;
let a1_k_ptr = a1_base + (k * kb_nh_mid_bytes) as u64;
let z2_k_ptr = z2_base + (k * kb_nh_mid_bytes) as u64;
let gate_k_ptr = gate_base + (k * kb_nh_bytes) as u64;
let main_k_ptr = main_base + (k * kb_nh_bytes) as u64;
let logit_k_ptr = logit_base + (k * kb_nh_bytes) as u64;
unsafe {
let mut launch = self.stream.launch_builder(&self.step_batched_fn);
launch
@@ -1547,14 +1767,21 @@ impl PerceptionTrainer {
launch.launch(cfg_cfc).context("eval cfc fwd")?;
}
unsafe {
let mut launch = self.stream.launch_builder(&self.heads_batched_fn);
let mut launch = self.stream.launch_builder(&self.heads_grn_fwd_fn);
launch
.arg(&self.heads_w_d).arg(&self.heads_b_d)
.arg(&h_new_k_ptr).arg(&n_batch_i).arg(&probs_k_ptr);
launch.launch(cfg_heads).context("eval heads fwd")?;
.arg(&self.heads_w1_d).arg(&self.heads_b1_d)
.arg(&self.heads_w2_d).arg(&self.heads_b2_d)
.arg(&self.heads_w_gate_d).arg(&self.heads_b_gate_d)
.arg(&self.heads_w_main_d).arg(&self.heads_b_main_d)
.arg(&self.heads_w_skip_d).arg(&self.heads_b_skip_d)
.arg(&h_new_k_ptr).arg(&n_batch_i)
.arg(&probs_k_ptr)
.arg(&z1_k_ptr).arg(&a1_k_ptr).arg(&z2_k_ptr)
.arg(&gate_k_ptr).arg(&main_k_ptr).arg(&logit_k_ptr);
launch.launch(cfg_grn_fwd).context("eval heads GRN fwd")?;
}
}
drop((_g_hpk, _g_probs));
drop((_g_hpk, _g_probs, _g_z1, _g_a1, _g_z2, _g_gate, _g_main, _g_logit));
// BCE for loss reporting (uses same per-horizon weights as train).
let n_pos_i = (k_seq * b_sz) as i32;