diff --git a/crates/ml-alpha/src/cfc/trunk.rs b/crates/ml-alpha/src/cfc/trunk.rs index c8bf4f10d..4423a14bb 100644 --- a/crates/ml-alpha/src/cfc/trunk.rs +++ b/crates/ml-alpha/src/cfc/trunk.rs @@ -19,7 +19,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, PROJ_DIM}; +use crate::heads::{HEAD_MID_DIM, HIDDEN_DIM, N_HORIZONS, PROJ_DIM}; use serde::{Deserialize, Serialize}; /// On-disk checkpoint envelope for CfcTrunk weights. Bumped version on @@ -51,11 +51,15 @@ const PROJ_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/projection.c pub struct CfcConfig { pub n_in: usize, pub n_hid: usize, + /// Mamba2 state dimension — used by the v2 trunk weight skeleton + /// (X1) and the per-stack Mamba2Block constructions added in X3/X5. + /// Default matches `PerceptionTrainerConfig::default().mamba2_state_dim`. + pub mamba2_state_dim: usize, } impl Default for CfcConfig { fn default() -> Self { - Self { n_in: FEATURE_DIM, n_hid: HIDDEN_DIM } + Self { n_in: FEATURE_DIM, n_hid: HIDDEN_DIM, mamba2_state_dim: 16 } } } @@ -73,7 +77,9 @@ pub struct CfcTrunk { heads_fn: CudaFunction, proj_fn: CudaFunction, - // Device-resident weights + // Device-resident weights — V1 (CfC + simple heads + projection). + // Preserved unchanged until X8/X9 consolidate against the v2 fields + // added below. w_in_d: CudaSlice, w_rec_d: CudaSlice, b_d: CudaSlice, @@ -85,6 +91,29 @@ pub struct CfcTrunk { proj_g_d: CudaSlice, proj_n_d: CudaSlice, + // === V2 weight skeleton (X1) — zero-initialised slabs + None for + // Mamba2 stacks. Each is wired in a subsequent commit (X2..X9): + // X2 VSN, X3 Mamba2_l1, X4 LN_a, X5 Mamba2_l2, X6 LN_b, + // X7 attn_pool, X8 CfC consolidation, X9 GRN heads. + // Public so PerceptionTrainer can read them during migration — + // visibility tightens to crate-private after X10 hoists the + // forward kernels into CfcTrunk methods. === + pub vsn_w_d: CudaSlice, + pub vsn_b_d: CudaSlice, + pub mamba2_stack_1: Option, + pub mamba2_stack_2: Option, + pub ln_a_gain_d: CudaSlice, + pub ln_a_bias_d: CudaSlice, + pub ln_b_gain_d: CudaSlice, + pub ln_b_bias_d: CudaSlice, + pub attn_q_d: CudaSlice, + pub heads_w_gate_d: CudaSlice, + pub heads_b_gate_d: CudaSlice, + pub heads_w_main_d: CudaSlice, + pub heads_b_main_d: CudaSlice, + pub heads_w_skip_d: CudaSlice, + pub heads_b_skip_d: CudaSlice, + // Ping-pong hidden state h_ping: CudaSlice, h_pong: CudaSlice, @@ -191,6 +220,36 @@ impl CfcTrunk { snap_feat_d: stream.alloc_zeros::(FEATURE_DIM).context("snap_feat alloc")?, probs_d: stream.alloc_zeros::(N_HORIZONS).context("probs alloc")?, proj_out_d: stream.alloc_zeros::(PROJ_DIM).context("proj_out alloc")?, + // === V2 skeleton fields (X1) — allocated BEFORE `stream` is + // moved into the struct below. Populated by X2..X9. === + vsn_w_d: stream.alloc_zeros::(cfg.n_in * cfg.n_in) + .context("v2 vsn_w alloc")?, + vsn_b_d: stream.alloc_zeros::(cfg.n_in) + .context("v2 vsn_b alloc")?, + mamba2_stack_1: None, + mamba2_stack_2: None, + ln_a_gain_d: stream.alloc_zeros::(cfg.n_hid) + .context("v2 ln_a_gain alloc")?, + ln_a_bias_d: stream.alloc_zeros::(cfg.n_hid) + .context("v2 ln_a_bias alloc")?, + ln_b_gain_d: stream.alloc_zeros::(cfg.n_hid) + .context("v2 ln_b_gain alloc")?, + ln_b_bias_d: stream.alloc_zeros::(cfg.n_hid) + .context("v2 ln_b_bias alloc")?, + attn_q_d: stream.alloc_zeros::(cfg.n_hid) + .context("v2 attn_q alloc")?, + heads_w_gate_d: stream.alloc_zeros::(N_HORIZONS * HEAD_MID_DIM) + .context("v2 heads_w_gate alloc")?, + heads_b_gate_d: stream.alloc_zeros::(N_HORIZONS) + .context("v2 heads_b_gate alloc")?, + heads_w_main_d: stream.alloc_zeros::(N_HORIZONS * HEAD_MID_DIM) + .context("v2 heads_w_main alloc")?, + heads_b_main_d: stream.alloc_zeros::(N_HORIZONS) + .context("v2 heads_b_main alloc")?, + heads_w_skip_d: stream.alloc_zeros::(N_HORIZONS * cfg.n_hid) + .context("v2 heads_w_skip alloc")?, + heads_b_skip_d: stream.alloc_zeros::(N_HORIZONS) + .context("v2 heads_b_skip alloc")?, stream, _snap_module: snap_module, _step_module: step_module,