diff --git a/crates/ml-alpha/src/heads.rs b/crates/ml-alpha/src/heads.rs index 2860e3bcf..8bf0202be 100644 --- a/crates/ml-alpha/src/heads.rs +++ b/crates/ml-alpha/src/heads.rs @@ -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 { diff --git a/crates/ml-alpha/src/trainer/perception.rs b/crates/ml-alpha/src/trainer/perception.rs index f317620d1..7b889eb37 100644 --- a/crates/ml-alpha/src/trainer/perception.rs +++ b/crates/ml-alpha/src/trainer/perception.rs @@ -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, pub b_d: CudaSlice, pub tau_d: CudaSlice, - pub heads_w_d: CudaSlice, - pub heads_b_d: CudaSlice, + // ── 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, // [N_HORIZONS, HEAD_MID, HIDDEN] + pub heads_b1_d: CudaSlice, // [N_HORIZONS, HEAD_MID] + pub heads_w2_d: CudaSlice, // [N_HORIZONS, HEAD_MID, HEAD_MID] + pub heads_b2_d: CudaSlice, // [N_HORIZONS, HEAD_MID] + pub heads_w_gate_d: CudaSlice, // [N_HORIZONS, HEAD_MID] + pub heads_b_gate_d: CudaSlice, // [N_HORIZONS] + pub heads_w_main_d: CudaSlice, // [N_HORIZONS, HEAD_MID] + pub heads_b_main_d: CudaSlice, // [N_HORIZONS] + pub heads_w_skip_d: CudaSlice, // [N_HORIZONS, HIDDEN] + pub heads_b_skip_d: CudaSlice, // [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, grad_b_d: CudaSlice, grad_tau_d: CudaSlice, - grad_heads_w_d: CudaSlice, - grad_heads_b_d: CudaSlice, + grad_heads_w1_d: CudaSlice, + grad_heads_b1_d: CudaSlice, + grad_heads_w2_d: CudaSlice, + grad_heads_b2_d: CudaSlice, + grad_heads_w_gate_d: CudaSlice, + grad_heads_b_gate_d: CudaSlice, + grad_heads_w_main_d: CudaSlice, + grad_heads_b_main_d: CudaSlice, + grad_heads_w_skip_d: CudaSlice, + grad_heads_b_skip_d: CudaSlice, + // 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, + a1_per_k_d: CudaSlice, + z2_per_k_d: CudaSlice, + // Per-K horizon-scalar intermediates: [K, B, N_HORIZONS]. + gate_logit_per_k_d: CudaSlice, + main_per_k_d: CudaSlice, + logit_per_k_d: CudaSlice, // 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 = (0..N_HORIZONS * n_hid).map(|_| r.gen_range(-head_scale..head_scale)).collect(); - let heads_b: Vec = 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 = (0..N_HORIZONS * HEAD_MID_DIM * n_hid) + .map(|_| r.gen_range(-scale_h_in..scale_h_in)).collect(); + let heads_b1: Vec = vec![0.0; N_HORIZONS * HEAD_MID_DIM]; + let heads_w2: Vec = (0..N_HORIZONS * HEAD_MID_DIM * HEAD_MID_DIM) + .map(|_| r.gen_range(-scale_mid_mid..scale_mid_mid)).collect(); + let heads_b2: Vec = vec![0.0; N_HORIZONS * HEAD_MID_DIM]; + let heads_w_gate: Vec = (0..N_HORIZONS * HEAD_MID_DIM) + .map(|_| r.gen_range(-scale_mid_mid..scale_mid_mid)).collect(); + let heads_b_gate: Vec = vec![0.0; N_HORIZONS]; + let heads_w_main: Vec = (0..N_HORIZONS * HEAD_MID_DIM) + .map(|_| r.gen_range(-scale_mid_mid..scale_mid_mid)).collect(); + let heads_b_main: Vec = vec![0.0; N_HORIZONS]; + let heads_w_skip: Vec = (0..N_HORIZONS * n_hid) + .map(|_| r.gen_range(-scale_h_skip..scale_h_skip)).collect(); + let heads_b_skip: Vec = 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::(n_hid * n_hid)?, grad_b_d: stream.alloc_zeros::(n_hid)?, grad_tau_d: stream.alloc_zeros::(n_hid)?, - grad_heads_w_d: stream.alloc_zeros::(N_HORIZONS * n_hid)?, - grad_heads_b_d: stream.alloc_zeros::(N_HORIZONS)?, + grad_heads_w1_d: stream.alloc_zeros::(N_HORIZONS * HEAD_MID_DIM * n_hid)?, + grad_heads_b1_d: stream.alloc_zeros::(N_HORIZONS * HEAD_MID_DIM)?, + grad_heads_w2_d: stream.alloc_zeros::(N_HORIZONS * HEAD_MID_DIM * HEAD_MID_DIM)?, + grad_heads_b2_d: stream.alloc_zeros::(N_HORIZONS * HEAD_MID_DIM)?, + grad_heads_w_gate_d: stream.alloc_zeros::(N_HORIZONS * HEAD_MID_DIM)?, + grad_heads_b_gate_d: stream.alloc_zeros::(N_HORIZONS)?, + grad_heads_w_main_d: stream.alloc_zeros::(N_HORIZONS * HEAD_MID_DIM)?, + grad_heads_b_main_d: stream.alloc_zeros::(N_HORIZONS)?, + grad_heads_w_skip_d: stream.alloc_zeros::(N_HORIZONS * n_hid)?, + grad_heads_b_skip_d: stream.alloc_zeros::(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::(k * cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?, + a1_per_k_d: stream.alloc_zeros::(k * cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?, + z2_per_k_d: stream.alloc_zeros::(k * cfg.n_batch * N_HORIZONS * HEAD_MID_DIM)?, + gate_logit_per_k_d: stream.alloc_zeros::(k * cfg.n_batch * N_HORIZONS)?, + main_per_k_d: stream.alloc_zeros::(k * cfg.n_batch * N_HORIZONS)?, + logit_per_k_d: stream.alloc_zeros::(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::(); let kb_nh_bytes = b_sz * N_HORIZONS * std::mem::size_of::(); + // 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::(); 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::(); let kb_nh_bytes = b_sz * N_HORIZONS * std::mem::size_of::(); + let kb_nh_mid_bytes = b_sz * N_HORIZONS * HEAD_MID_DIM * std::mem::size_of::(); 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;