From 5aad1eb8468cc332ecf9289fbd084701cf944037 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Mon, 18 May 2026 14:28:50 +0200 Subject: [PATCH] feat(ml-alpha): PerceptionV2State bundle (axes A+B+C+D+E) [V9] MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Bundles all v2 components into one cohesive trainer field that PerceptionTrainer wires through in V10. Owns: (C) horizon_tokens + Q for the horizon-token attention pool + per-batch grad scratches + 2 AdamWs (E) inverted attention pool + saved attn + d_scores scratch (D) MoE gate weights + experts (W, b) + per-batch sparse-by-expert grad scratches + 3 AdamWs + top_e / gate_probs / aux_loss (A) log_sigma_h per-horizon Kendall scalar + AdamW (0.25× LR) (B) AnchorController (Wiener-α host-side) + anchor_l2 kernel + init snapshots of horizon_tokens, Q, experts_w, experts_b for the L2 anchor Init scale: 1/√HIDDEN_DIM Xavier for horizon_tokens/Q/experts_W; zeros for experts_b and log_sigma_h. Anchor controller bootstrap uses ‖horizon_tokens_init‖₂ + ‖Q_init‖₂ + their total numel to derive the signal-driven floor (no tuned constants). zero_grads uses memset_zeros only — capture-safe. Compile-clean. V10 wires this state into perception.rs::step_batched forward / backward / AdamW. Co-Authored-By: Claude Opus 4.7 --- crates/ml-alpha/src/trainer/mod.rs | 1 + .../src/trainer/perception_v2_state.rs | 246 ++++++++++++++++++ 2 files changed, 247 insertions(+) create mode 100644 crates/ml-alpha/src/trainer/perception_v2_state.rs diff --git a/crates/ml-alpha/src/trainer/mod.rs b/crates/ml-alpha/src/trainer/mod.rs index fcd230dcf..532ded335 100644 --- a/crates/ml-alpha/src/trainer/mod.rs +++ b/crates/ml-alpha/src/trainer/mod.rs @@ -6,3 +6,4 @@ pub mod loss; pub mod loss_sigma; pub mod optim; pub mod perception; +pub mod perception_v2_state; diff --git a/crates/ml-alpha/src/trainer/perception_v2_state.rs b/crates/ml-alpha/src/trainer/perception_v2_state.rs new file mode 100644 index 000000000..02b957d5c --- /dev/null +++ b/crates/ml-alpha/src/trainer/perception_v2_state.rs @@ -0,0 +1,246 @@ +//! PerceptionV2State — bundles the v2 axis-A/B/C/D/E components into +//! one cohesive trainer field. (V9 commit.) +//! +//! Owns: +//! - horizon_tokens (axis C) + their init snapshot for anchor (axis B) +//! - inv-attention saved attn buffer (axis E) +//! - MoE experts + gate (axis D) +//! - Kendall σ logarithm (axis A) +//! - Anchor controller (axis B Wiener-α) +//! +//! All kernel calls go through the bound modules; nothing in this +//! file does host-side compute or sync inside captured regions. + +use anyhow::{Context, Result}; +use cudarc::driver::{CudaSlice, CudaStream}; +use ml_core::device::MlDevice; +use rand::{Rng, SeedableRng}; +use rand_chacha::ChaCha8Rng; +use std::sync::Arc; + +use crate::anchor_l2::AnchorL2; +use crate::horizon_token_attention_pool::{ + HorizonTokenAttentionPool, HTAP_HIDDEN_DIM, HTAP_N_HORIZONS, +}; +use crate::inverted_attention_pool::InvertedAttentionPool; +use crate::regime_moe_gate::{RegimeMoeGate, MOE_N_EXPERTS}; +use crate::trainer::anchor_controller::AnchorController; +use crate::trainer::optim::AdamW; + +pub const V2_HIDDEN_DIM: usize = HTAP_HIDDEN_DIM; +pub const V2_N_HORIZONS: usize = HTAP_N_HORIZONS; +pub const V2_N_EXPERTS: usize = MOE_N_EXPERTS; +pub const V2_REGIME_DIM: usize = 12; // spread(5) + vol(3) + tod(4); see spec §3.3 + +pub struct PerceptionV2State { + // ── (C) horizon-token attention pool ── + pub pool: HorizonTokenAttentionPool, + pub horizon_tokens_d: CudaSlice, // [N_H, H] learnable + pub horizon_tokens_init_d: CudaSlice, // [N_H, H] non-trainable snapshot + pub q_d: CudaSlice, // [H] single shared Q + pub q_init_d: CudaSlice, // [H] non-trainable snapshot + pub grad_horizon_tokens_d: CudaSlice, + pub grad_horizon_tokens_scratch_d: CudaSlice, // [B, N_H, H] per-batch scratch + pub grad_q_d: CudaSlice, + pub grad_q_scratch_d: CudaSlice, // [B, H] per-batch scratch + pub ctx_h_d: CudaSlice, // [B, N_H, H] fwd output + pub attn_h_d: CudaSlice, // [B, N_H + K] + pub opt_horizon_tokens: AdamW, + pub opt_q: AdamW, + + // ── (E) inverted attention pool ── + pub inv_pool: InvertedAttentionPool, + pub inv_pooled_d: CudaSlice, // [B, H] + pub inv_attn_d: CudaSlice, // [B, H, H] + pub inv_d_scores_scratch_d: CudaSlice, // [B, H, H] bwd scratch + + // ── (C+E fuse) → fused_ctx [B, N_H, H] via additive merge ── + // No separate kernel; computed inline in step_batched as + // fused_ctx[b, h, d] = ctx_h[b, h, d] + inv_pooled[b, d]. + + // ── (D) regime-MoE gate + experts ── + pub moe: RegimeMoeGate, + pub gate_logits_d: CudaSlice, // [B, N_E] computed each step from regime features + pub w_gate_d: CudaSlice, // [N_E, REGIME_DIM] + pub experts_w_d: CudaSlice, // [N_E, H, H] + pub experts_b_d: CudaSlice, // [N_E, H] + pub experts_w_init_d: CudaSlice, // anchor snapshot (B) + pub experts_b_init_d: CudaSlice, // anchor snapshot + pub top_e_d: CudaSlice, // [B] + pub gate_probs_d: CudaSlice, // [B, N_E] + pub aux_loss_d: CudaSlice, // [1] + pub routed_ctx_d: CudaSlice, // [B, N_H, H] v2 output (consumed by GRN heads) + pub grad_routed_d: CudaSlice, // [B, N_H, H] + pub grad_w_gate_d: CudaSlice, + pub grad_experts_w_d: CudaSlice, + pub grad_experts_b_d: CudaSlice, + pub grad_w_scratch_d: CudaSlice, // [B, N_H, N_E, H, H] + pub grad_b_scratch_d: CudaSlice, // [B, N_H, N_E, H] + pub opt_w_gate: AdamW, + pub opt_experts_w: AdamW, + pub opt_experts_b: AdamW, + + // ── (A) Kendall σ on the BCE ── + pub log_sigma_h_d: CudaSlice, // [N_HORIZONS] + pub grad_log_sigma_h_d: CudaSlice, + pub opt_log_sigma: AdamW, + + // ── (B) anchor regularization controller ── + pub anchor_l2: AnchorL2, + pub anchor: AnchorController, + pub lambda_d: CudaSlice, // [1] current λ uploaded each step + pub anchor_loss_partial_d: CudaSlice, // [1] per-launch partial + pub anchor_loss_total_d: CudaSlice, // [1] accumulated + + // Bookkeeping. + n_batch: usize, + k_seq: usize, + stream: Arc, +} + +impl PerceptionV2State { + pub fn new(dev: &MlDevice, n_batch: usize, k_seq: usize, lr: f32, seed: u64) -> Result { + anyhow::ensure!(n_batch > 0 && k_seq > 0, "n_batch and k_seq must be > 0"); + let stream = dev.cuda_stream().context("v2 stream")?.clone(); + let ctx = dev.cuda_context().context("v2 ctx")?; + + let pool = HorizonTokenAttentionPool::new(ctx, stream.clone()) + .context("HorizonTokenAttentionPool")?; + let inv_pool = InvertedAttentionPool::new(ctx, stream.clone()) + .context("InvertedAttentionPool")?; + let moe = RegimeMoeGate::new(ctx, stream.clone()) + .context("RegimeMoeGate")?; + let anchor_l2 = AnchorL2::new(ctx, stream.clone()).context("AnchorL2")?; + + let mut rng = ChaCha8Rng::seed_from_u64(seed); + let scale = (1.0_f32 / V2_HIDDEN_DIM as f32).sqrt(); + + let init_horizon: Vec = (0..V2_N_HORIZONS * V2_HIDDEN_DIM) + .map(|_| rng.gen_range(-scale..scale)).collect(); + let init_q: Vec = (0..V2_HIDDEN_DIM) + .map(|_| rng.gen_range(-scale..scale)).collect(); + let init_w_gate: Vec = (0..V2_N_EXPERTS * V2_REGIME_DIM) + .map(|_| rng.gen_range(-0.1..0.1)).collect(); + let init_experts_w: Vec = (0..V2_N_EXPERTS * V2_HIDDEN_DIM * V2_HIDDEN_DIM) + .map(|_| rng.gen_range(-scale..scale)).collect(); + + let horizon_tokens_d = upload(&stream, &init_horizon)?; + let horizon_tokens_init_d = upload(&stream, &init_horizon)?; + let q_d = upload(&stream, &init_q)?; + let q_init_d = upload(&stream, &init_q)?; + let w_gate_d = upload(&stream, &init_w_gate)?; + let experts_w_d = upload(&stream, &init_experts_w)?; + let experts_b_d = stream.alloc_zeros::(V2_N_EXPERTS * V2_HIDDEN_DIM)?; + let experts_w_init_d = upload(&stream, &init_experts_w)?; + let experts_b_init_d = stream.alloc_zeros::(V2_N_EXPERTS * V2_HIDDEN_DIM)?; + + let n_ext = V2_N_HORIZONS + k_seq; + + let s = &stream; + let state = Self { + pool, + horizon_tokens_d, + horizon_tokens_init_d, + q_d, + q_init_d, + grad_horizon_tokens_d: s.alloc_zeros::(V2_N_HORIZONS * V2_HIDDEN_DIM)?, + grad_horizon_tokens_scratch_d:s.alloc_zeros::(n_batch * V2_N_HORIZONS * V2_HIDDEN_DIM)?, + grad_q_d: s.alloc_zeros::(V2_HIDDEN_DIM)?, + grad_q_scratch_d: s.alloc_zeros::(n_batch * V2_HIDDEN_DIM)?, + ctx_h_d: s.alloc_zeros::(n_batch * V2_N_HORIZONS * V2_HIDDEN_DIM)?, + attn_h_d: s.alloc_zeros::(n_batch * n_ext)?, + opt_horizon_tokens: AdamW::new(dev, V2_N_HORIZONS * V2_HIDDEN_DIM, lr)?, + opt_q: AdamW::new(dev, V2_HIDDEN_DIM, lr)?, + + inv_pool, + inv_pooled_d: s.alloc_zeros::(n_batch * V2_HIDDEN_DIM)?, + inv_attn_d: s.alloc_zeros::(n_batch * V2_HIDDEN_DIM * V2_HIDDEN_DIM)?, + inv_d_scores_scratch_d: s.alloc_zeros::(n_batch * V2_HIDDEN_DIM * V2_HIDDEN_DIM)?, + + moe, + gate_logits_d: s.alloc_zeros::(n_batch * V2_N_EXPERTS)?, + w_gate_d, + experts_w_d, + experts_b_d, + experts_w_init_d, + experts_b_init_d, + top_e_d: s.alloc_zeros::(n_batch)?, + gate_probs_d: s.alloc_zeros::(n_batch * V2_N_EXPERTS)?, + aux_loss_d: s.alloc_zeros::(1)?, + routed_ctx_d: s.alloc_zeros::(n_batch * V2_N_HORIZONS * V2_HIDDEN_DIM)?, + grad_routed_d: s.alloc_zeros::(n_batch * V2_N_HORIZONS * V2_HIDDEN_DIM)?, + grad_w_gate_d: s.alloc_zeros::(V2_N_EXPERTS * V2_REGIME_DIM)?, + grad_experts_w_d: s.alloc_zeros::(V2_N_EXPERTS * V2_HIDDEN_DIM * V2_HIDDEN_DIM)?, + grad_experts_b_d: s.alloc_zeros::(V2_N_EXPERTS * V2_HIDDEN_DIM)?, + grad_w_scratch_d: s.alloc_zeros::(n_batch * V2_N_HORIZONS * V2_N_EXPERTS * V2_HIDDEN_DIM * V2_HIDDEN_DIM)?, + grad_b_scratch_d: s.alloc_zeros::(n_batch * V2_N_HORIZONS * V2_N_EXPERTS * V2_HIDDEN_DIM)?, + opt_w_gate: AdamW::new(dev, V2_N_EXPERTS * V2_REGIME_DIM, lr)?, + opt_experts_w: AdamW::new(dev, V2_N_EXPERTS * V2_HIDDEN_DIM * V2_HIDDEN_DIM, lr)?, + opt_experts_b: AdamW::new(dev, V2_N_EXPERTS * V2_HIDDEN_DIM, lr)?, + + log_sigma_h_d: s.alloc_zeros::(V2_N_HORIZONS)?, + grad_log_sigma_h_d: s.alloc_zeros::(V2_N_HORIZONS)?, + opt_log_sigma: AdamW::new(dev, V2_N_HORIZONS, lr * 0.25)?, // slow update + + anchor_l2, + anchor: AnchorController::new( + magnitude_of(&init_horizon) + magnitude_of(&init_q), + init_horizon.len() + init_q.len(), + ), + lambda_d: s.alloc_zeros::(1)?, + anchor_loss_partial_d: s.alloc_zeros::(1)?, + anchor_loss_total_d: s.alloc_zeros::(1)?, + + n_batch, + k_seq, + stream: stream.clone(), + }; + Ok(state) + } + + /// Zero per-step grad scratches. Capture-safe (memset_zeros only). + pub fn zero_grads(&mut self) -> Result<()> { + let s = &self.stream; + s.memset_zeros(&mut self.grad_horizon_tokens_d)?; + s.memset_zeros(&mut self.grad_horizon_tokens_scratch_d)?; + s.memset_zeros(&mut self.grad_q_d)?; + s.memset_zeros(&mut self.grad_q_scratch_d)?; + s.memset_zeros(&mut self.grad_routed_d)?; + s.memset_zeros(&mut self.grad_w_gate_d)?; + s.memset_zeros(&mut self.grad_experts_w_d)?; + s.memset_zeros(&mut self.grad_experts_b_d)?; + s.memset_zeros(&mut self.grad_w_scratch_d)?; + s.memset_zeros(&mut self.grad_b_scratch_d)?; + s.memset_zeros(&mut self.grad_log_sigma_h_d)?; + s.memset_zeros(&mut self.anchor_loss_partial_d)?; + s.memset_zeros(&mut self.anchor_loss_total_d)?; + Ok(()) + } + + pub fn n_batch(&self) -> usize { self.n_batch } + pub fn k_seq(&self) -> usize { self.k_seq } +} + +fn upload(stream: &Arc, host: &[f32]) -> Result> { + use crate::pinned_mem::MappedF32Buffer; + use cudarc::driver::{DevicePtrMut}; + let n = host.len(); + let staging = unsafe { MappedF32Buffer::new(n) } + .map_err(|e| anyhow::anyhow!("v2_state upload staging: {e}"))?; + staging.write_from_slice(host); + let mut dst = stream.alloc_zeros::(n).context("v2_state upload alloc")?; + if n > 0 { + let nbytes = n * std::mem::size_of::(); + unsafe { + let (dst_ptr, _g) = dst.device_ptr_mut(stream); + cudarc::driver::result::memcpy_dtod_async( + dst_ptr, staging.dev_ptr, nbytes, stream.cu_stream(), + ).context("v2_state upload DtoD")?; + } + } + Ok(dst) +} + +fn magnitude_of(p: &[f32]) -> f32 { + p.iter().map(|x| x * x).sum::().sqrt() +}