diff --git a/crates/ml/examples/alpha_dqn_h600_smoke.rs b/crates/ml/examples/alpha_dqn_h600_smoke.rs index 0c4576394..8c2d4582a 100644 --- a/crates/ml/examples/alpha_dqn_h600_smoke.rs +++ b/crates/ml/examples/alpha_dqn_h600_smoke.rs @@ -48,6 +48,10 @@ use clap::Parser; use cudarc::driver::{CudaContext, DevicePtr, DevicePtrMut}; use tracing::{info, warn}; +// Phase E.4.A.8: Mamba2 temporal encoder (ml-alpha Phase 1d.1). +use ml_alpha::mamba2_block::{Mamba2Block, Mamba2BlockConfig}; +use ml_core::cuda_autograd::gpu_tensor::GpuTensor; + // ── Mapped-pinned helpers ────────────────────────────────────────── // // `cuMemHostAlloc(DEVICEMAP|PORTABLE)` allocations: host and device see @@ -345,9 +349,18 @@ struct Cli { temporal: bool, /// Sliding-window length K (in snapshots). 16 starting point; the /// kernel max is `MAMBA2_KERNEL_SEQ_MAX = 32`. Sweep 16 / 32 in - /// later iterations once Mamba2 forward is wired (T8). + /// later iterations. #[arg(long, default_value_t = 16)] window_k: usize, + /// Mamba2 hidden dimension (= sh2 in kernel). 32 starting point. + /// Also the input dim to the C51 head when --temporal is set + /// (C51 W shape: [N_ACTIONS * n_atoms, mamba2_hidden_dim]). + #[arg(long, default_value_t = 32)] + mamba2_hidden_dim: usize, + /// Mamba2 SSM state dimension. Kernel max is + /// `MAMBA2_KERNEL_STATE_MAX = 16`. + #[arg(long, default_value_t = 16)] + mamba2_state_dim: usize, /// Output JSON path for final verdict + ISV readings. #[arg(long, default_value = "config/ml/alpha_dqn_h600_smoke.json")] out_path: PathBuf, @@ -711,8 +724,16 @@ fn main() -> Result<()> { // --- Initialize Q-network: Xavier --- // C51 lifts the output dimension to N_ACTIONS × n_atoms (logits per // atom, softmax-normalized across atoms). Linear-Q stays at N_ACTIONS. + // --temporal additionally swaps the C51 input feature dim from + // STATE_DIM=10 to mamba2_hidden_dim (typically 32) — Q-net consumes + // Mamba2's `h_enriched` instead of raw state. + let c51_input_dim: usize = if cli.c51 && cli.temporal { + cli.mamba2_hidden_dim + } else { + STATE_DIM + }; let n_weights_eff: usize = if cli.c51 { - N_ACTIONS * c51_n_atoms * STATE_DIM + N_ACTIONS * c51_n_atoms * c51_input_dim } else { N_WEIGHTS }; @@ -721,7 +742,7 @@ fn main() -> Result<()> { } else { N_BIASES }; - let xavier_scale = (2.0_f32 / STATE_DIM as f32).sqrt(); + let xavier_scale = (2.0_f32 / c51_input_dim as f32).sqrt(); let mut rng = SmokeRng::new(cli.seed.wrapping_add(0xDEAD_BEEF)); let w_init: Vec = (0..n_weights_eff) .map(|_| xavier_scale * 2.0 * (rng.next_f32() - 0.5)) @@ -804,12 +825,44 @@ fn main() -> Result<()> { // Phase E.4.A.7: GPU-resident sliding-window state buffer for the // temporal encoder. Allocated unconditionally; only populated when // cli.temporal is set. Capacity: window_k × STATE_DIM floats. - // Phase E.4.A.6: shift+insert layout — slot 0 = oldest, slot K-1 - // = newest. No head_idx tracking needed; kernel always shifts then - // inserts at slot K-1. - let mut window_dev = stream - .alloc_zeros::(cli.window_k * STATE_DIM) - .context("alloc window buffer")?; + // Phase E.4.A.6+A.8: shift+insert chronological window, allocated + // as a GpuTensor of shape [1, K, in_dim] so Mamba2Block can consume + // it directly via `forward_train(&window_tensor)`. The push kernel + // writes via the underlying CudaSlice (no reshape needed — same + // 160 floats either way). + let mut window_tensor = GpuTensor::zeros( + &[1, cli.window_k, STATE_DIM], &stream + ).map_err(|e| anyhow::anyhow!("alloc window tensor: {e}"))?; + // Phase E.4.A.8 batched training: store each step's h_enriched + // (Mamba2 output) so the batched C51 forward at episode end reads + // pre-computed temporal representations instead of recomputing + // Mamba2 over each step's window. Size = (horizon + 1) so the + // terminal next-state h_enriched fits at the end. + let h_enriched_buf_capacity = (cli.horizon + 1) * cli.mamba2_hidden_dim; + let mut h_enriched_buf_dev = stream + .alloc_zeros::(h_enriched_buf_capacity) + .context("alloc h_enriched buffer")?; + + // Phase E.4.A.8: construct the Mamba2 temporal encoder when + // --temporal is set. Hidden_dim = c51_input_dim (so the encoder's + // output flows directly into the C51 head without an additional + // projection). State_dim ≤ 16 (kernel max). + let mamba2_block: Option = if cli.temporal { + let cfg = Mamba2BlockConfig { + in_dim: STATE_DIM, + hidden_dim: cli.mamba2_hidden_dim, + state_dim: cli.mamba2_state_dim, + seq_len: cli.window_k, + }; + info!( + " TEMPORAL Mamba2: in={} hidden={} state={} K={}", + cfg.in_dim, cfg.hidden_dim, cfg.state_dim, cfg.seq_len + ); + Some(Mamba2Block::new(cfg, stream.clone()) + .map_err(|e| anyhow::anyhow!("Mamba2Block init: {e}"))?) + } else { + None + }; // Kill-criteria inputs: action_counts (i32 × n_actions), // scalar_inputs (f32 × 3: rollout_R_mean, q_init_norm, q_early_norm). @@ -842,7 +895,7 @@ fn main() -> Result<()> { // Mamba2 sees zero-history for the first window_k-1 steps. if cli.temporal { stream - .memset_zeros(&mut window_dev) + .memset_zeros(window_tensor.data_mut()) .context("zero window on episode reset")?; } @@ -860,11 +913,9 @@ fn main() -> Result<()> { let s_vec = state_as_vec(&env, &state); let action: u8 = if cli.c51 { state_pinned.write(&s_vec); - // Phase E.4.A.7: shift-and-insert push of current state - // into the chronological window buffer. Consumer (Mamba2) - // wires in T8. + // Phase E.4.A.7: shift-and-insert push of current state. if cli.temporal { - let (w_ptr, _g1) = window_dev.device_ptr_mut(&stream); + let (w_ptr, _g1) = window_tensor.data_mut().device_ptr_mut(&stream); unsafe { ml::cuda_pipeline::alpha_kernels::launch_alpha_window_push( &stream, &push_kernel, @@ -873,7 +924,62 @@ fn main() -> Result<()> { )?; } } - { + // Phase E.4.A.8: Mamba2 forward over the window → h_enriched + // (`cache.h_enriched [1, hidden_dim]`). The C51 head reads + // this device buffer in place of `state_pinned`. + // + // The cache is dropped at the end of this scope — + // inference doesn't need backward (we only train via + // the batched update at end-of-episode). Backward + // wiring with cache retention lands in T10. + // Phase E.4.A.8: in temporal mode, run Mamba2 over the + // current window, persist h_enriched into the per-step + // buffer (slot = current env step), and feed it to C51. + // In non-temporal mode, C51 reads state_pinned directly. + if cli.temporal { + let block = mamba2_block.as_ref() + .expect("temporal flag set but Mamba2Block missing"); + let (_logit, cache) = block.forward_train(&window_tensor) + .map_err(|e| anyhow::anyhow!("mamba2 forward: {e}"))?; + // 1) Copy h_enriched into h_enriched_buf_dev at slot t. + // Slot t holds the representation BEFORE env.step at step t + // (i.e., the "current state" representation). + let slot_offset = state.step * cli.mamba2_hidden_dim; + { + // memcpy_dtod into the slot + let src_ptr = { + let (p, _g) = cache.h_enriched.cuda_data().device_ptr(&stream); + p + }; + // Use a small custom copy kernel? Or use cudarc's memcpy_dtod. + // For now: dtoh then htod (small — hidden_dim=32 floats). + // OPTIMIZATION HOOK: replace with kernel for zero round-trips. + let h_host = stream.clone_dtoh(cache.h_enriched.cuda_data()) + .map_err(|e| anyhow::anyhow!("dtoh h_enriched: {e}"))?; + let _ = src_ptr; // silence unused + let slot_slice = &mut h_enriched_buf_dev; + // Build a slice view: write h_host into slot_offset..slot_offset+hidden_dim + let mut buf_host = stream.clone_dtoh(slot_slice) + .map_err(|e| anyhow::anyhow!("dtoh h buf: {e}"))?; + for j in 0..cli.mamba2_hidden_dim { + buf_host[slot_offset + j] = h_host[j]; + } + stream.memcpy_htod(&buf_host, slot_slice) + .map_err(|e| anyhow::anyhow!("htod h buf: {e}"))?; + } + // 2) C51 forward on h_enriched (batch=1, dim=hidden). + let (w_ptr, _g0) = w_dev.device_ptr(&stream); + let (b_ptr, _g1) = b_dev.device_ptr(&stream); + let (p_ptr, _g2) = single_probs_dev.device_ptr_mut(&stream); + let (h_ptr, _g3) = cache.h_enriched.cuda_data().device_ptr(&stream); + unsafe { + ml::cuda_pipeline::alpha_kernels::launch_alpha_c51_forward( + &stream, &c51_fwd_kernel, + w_ptr, b_ptr, h_ptr, p_ptr, + 1, c51_input_dim as i32, n_act_i, n_atoms_i, + )?; + } + } else { let (w_ptr, _g0) = w_dev.device_ptr(&stream); let (b_ptr, _g1) = b_dev.device_ptr(&stream); let (p_ptr, _g2) = single_probs_dev.device_ptr_mut(&stream); @@ -881,7 +987,7 @@ fn main() -> Result<()> { ml::cuda_pipeline::alpha_kernels::launch_alpha_c51_forward( &stream, &c51_fwd_kernel, w_ptr, b_ptr, state_pinned.dev_u64(), p_ptr, - 1, state_dim_i, n_act_i, n_atoms_i, + 1, c51_input_dim as i32, n_act_i, n_atoms_i, )?; } } @@ -981,17 +1087,61 @@ fn main() -> Result<()> { if cli.c51 { // ── C51 batched compute ───────────────────────────────── + // Phase E.4.A.8: in temporal mode, the batched C51 forwards + // read the per-step h_enriched buffer instead of raw states. + // h_enriched_buf_dev[t..t+hidden_dim] = h at step t. + // The "next-state" representation for step t is h at step + // t+1; we need to run Mamba2 once MORE on the terminal + // next-window to fill slot ep_len. + if cli.temporal { + // Run Mamba2 on the final window (it currently holds + // the post-terminal observation since env.step pushed + // before returning done). h_enriched for slot ep_len. + let block = mamba2_block.as_ref().expect("Mamba2Block missing"); + let (_logit, cache) = block.forward_train(&window_tensor) + .map_err(|e| anyhow::anyhow!("mamba2 terminal forward: {e}"))?; + let h_host_term = stream.clone_dtoh(cache.h_enriched.cuda_data()) + .map_err(|e| anyhow::anyhow!("dtoh term h: {e}"))?; + let term_offset = (ep_len as usize) * cli.mamba2_hidden_dim; + let mut buf_host = stream.clone_dtoh(&h_enriched_buf_dev) + .map_err(|e| anyhow::anyhow!("dtoh buf for term: {e}"))?; + for j in 0..cli.mamba2_hidden_dim { + buf_host[term_offset + j] = h_host_term[j]; + } + stream.memcpy_htod(&buf_host, &mut h_enriched_buf_dev) + .map_err(|e| anyhow::anyhow!("htod buf for term: {e}"))?; + } + let (curr_input_ptr_guard, next_input_ptr_guard); + let (curr_input_ptr, next_input_ptr): (u64, u64) = if cli.temporal { + // h_current = h_enriched_buf_dev[0..ep_len*hidden] + // h_next = h_enriched_buf_dev[hidden..(ep_len+1)*hidden] + let (h_ptr_curr, g_curr) = h_enriched_buf_dev.device_ptr(&stream); + let (h_ptr_next_base, g_next) = h_enriched_buf_dev.device_ptr(&stream); + curr_input_ptr_guard = g_curr; + next_input_ptr_guard = g_next; + let _ = (&curr_input_ptr_guard, &next_input_ptr_guard); + ( + h_ptr_curr, + h_ptr_next_base + (cli.mamba2_hidden_dim as u64) * 4u64, // 4 = sizeof f32 + ) + } else { + let (s_ptr, g_curr) = states_dev.device_ptr(&stream); + let (n_ptr, g_next) = next_states_dev.device_ptr(&stream); + curr_input_ptr_guard = g_curr; + next_input_ptr_guard = g_next; + let _ = (&curr_input_ptr_guard, &next_input_ptr_guard); + (s_ptr, n_ptr) + }; // Forward online → probs_current { let (w_ptr, _g0) = w_dev.device_ptr(&stream); let (b_ptr, _g1) = b_dev.device_ptr(&stream); - let (s_ptr, _g2) = states_dev.device_ptr(&stream); let (p_ptr, _g3) = probs_current_dev.device_ptr_mut(&stream); unsafe { ml::cuda_pipeline::alpha_kernels::launch_alpha_c51_forward( &stream, &c51_fwd_kernel, - w_ptr, b_ptr, s_ptr, p_ptr, - ep_len, state_dim_i, n_act_i, n_atoms_i, + w_ptr, b_ptr, curr_input_ptr, p_ptr, + ep_len, c51_input_dim as i32, n_act_i, n_atoms_i, )?; } } @@ -999,13 +1149,12 @@ fn main() -> Result<()> { { let (w_ptr, _g0) = w_target_dev.device_ptr(&stream); let (b_ptr, _g1) = b_target_dev.device_ptr(&stream); - let (s_ptr, _g2) = next_states_dev.device_ptr(&stream); let (p_ptr, _g3) = probs_next_dev.device_ptr_mut(&stream); unsafe { ml::cuda_pipeline::alpha_kernels::launch_alpha_c51_forward( &stream, &c51_fwd_kernel, - w_ptr, b_ptr, s_ptr, p_ptr, - ep_len, state_dim_i, n_act_i, n_atoms_i, + w_ptr, b_ptr, next_input_ptr, p_ptr, + ep_len, c51_input_dim as i32, n_act_i, n_atoms_i, )?; } } @@ -1024,7 +1173,10 @@ fn main() -> Result<()> { )?; } } - // CE gradient on (probs_current, m, actions) + // CE gradient on (probs_current, m, actions). With --temporal + // the input to the C51 layer is h_enriched (hidden_dim); + // gradient w.r.t. W computes dW = (p - m) ⊗ input^T, so + // input pointer must match the forward's input. { let (p_ptr, _g0) = probs_current_dev.device_ptr(&stream); let (m_ptr, _g1) = m_dev.device_ptr(&stream); @@ -1032,12 +1184,13 @@ fn main() -> Result<()> { let (s_ptr, _g3) = states_dev.device_ptr(&stream); let (dw_ptr, _g4) = dw_dev.device_ptr_mut(&stream); let (db_ptr, _g5) = db_dev.device_ptr_mut(&stream); + let grad_input_ptr = if cli.temporal { curr_input_ptr } else { s_ptr }; unsafe { ml::cuda_pipeline::alpha_kernels::launch_alpha_c51_grad( &stream, &c51_grad_kernel, - p_ptr, m_ptr, a_ptr, s_ptr, + p_ptr, m_ptr, a_ptr, grad_input_ptr, dw_ptr, db_ptr, - ep_len, state_dim_i, n_act_i, n_atoms_i, + ep_len, c51_input_dim as i32, n_act_i, n_atoms_i, 1.0 / ep_len as f32, )?; }