From 0ab54dc8efc632484cd8a8f64a0279c47f2d53b6 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Fri, 15 May 2026 23:11:01 +0200 Subject: [PATCH] feat(phase-e-4-a): batched parallel-env TRAINING (greenfield) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit T15: training rewrite mirroring T14 eval. N_par parallel envs lockstep H steps per epoch; ONE batched C51 update at B = N_par * H. Expected ~15× speedup vs sequential. - NEW kernel alpha_h_enriched_store_batched_kernel for batched h_enriched slot writes - Training section greenfielded: legacy sequential loop deleted - CLI flag --n-train-par (default 50) - Terminal next-state slot zeroed; done=1 at horizon masks Q_next contribution in Bellman projection — no terminal Mamba2 forward - docs/isv-slots.md updated per kernel-audit-doc hook Co-Authored-By: Claude Opus 4.7 --- crates/ml/examples/alpha_compose_backtest.rs | 509 ++++++++---------- crates/ml/src/cuda_pipeline/alpha_kernels.rs | 38 ++ .../ml/src/cuda_pipeline/alpha_window_push.cu | 22 + docs/isv-slots.md | 13 + 4 files changed, 296 insertions(+), 286 deletions(-) diff --git a/crates/ml/examples/alpha_compose_backtest.rs b/crates/ml/examples/alpha_compose_backtest.rs index 80b58bced..233066f83 100644 --- a/crates/ml/examples/alpha_compose_backtest.rs +++ b/crates/ml/examples/alpha_compose_backtest.rs @@ -265,6 +265,14 @@ struct Cli { /// rate / Sharpe. #[arg(long, default_value_t = false)] isv_continual: bool, + /// Phase E.4.A.T15.batched (2026-05-15): parallel envs per training + /// "epoch". n_train_episodes / n_train_par epochs total. Each epoch + /// runs N_par envs in lockstep H steps + ONE batched C51 update at + /// B = N_par * H. Default 50 keeps SGD semantics (W is updated 20 + /// times over 1000 training eps); raise to 1000 for one-epoch full- + /// batch training. + #[arg(long, default_value_t = 50)] + n_train_par: usize, #[arg(long, default_value = "config/ml/alpha_compose_backtest.json")] out_path: PathBuf, } @@ -391,6 +399,7 @@ fn main() -> Result<()> { let push_kernel = push_module.load_function("alpha_window_push_kernel")?; let h_store_kernel = push_module.load_function("alpha_h_enriched_store_kernel")?; let push_batched_kernel = push_module.load_function("alpha_window_push_batched_kernel")?; + let h_store_batched_kernel = push_module.load_function("alpha_h_enriched_store_batched_kernel")?; // Phase E.4.A.7 controller (used in --isv-continual eval path). let ctl_module = ctx .load_cubin(ml::cuda_pipeline::alpha_kernels::STACKER_THRESHOLD_CONTROLLER_CUBIN.to_vec()) @@ -535,316 +544,245 @@ fn main() -> Result<()> { let mut isv_dev = stream.clone_htod(&isv_host).context("upload isv")?; let mut ctl_wiener_dev = stream.alloc_zeros::(3).context("alloc ctl wiener")?; - // --------------------------------------------------------------- - // Phase 1: TRAIN DQN on train segment. - // --------------------------------------------------------------- - info!("=== Training phase: {} episodes on train segment ===", cli.n_train_episodes); - let mut episode_rng = SmokeRng::new(cli.seed.wrapping_add(0xFEED)); + // ============================================================ + // Phase 1: TRAIN DQN — BATCHED PARALLEL ENVS (Phase E.4.A.T15) + // ============================================================ + let train_n_par = cli.n_train_par; + let train_n_epochs = (cli.n_train_episodes + train_n_par - 1) / train_n_par; + info!("=== Training phase: {} epochs × N_par={} parallel envs × H={} steps (BATCHED) ===", + train_n_epochs, train_n_par, cli.horizon); let train_max_start = (n_train.saturating_sub(cli.horizon + 1)).max(1); - for ep in 0..cli.n_train_episodes { - let eps = cli.eps_start - + (cli.eps_end - cli.eps_start) - * (ep as f32 / cli.n_train_episodes.max(1) as f32); - let start_cursor = (episode_rng.next_u64() as usize) % train_max_start; - let env_seed = episode_rng.next_u64(); - env.reset_at(env_seed, start_cursor); - let mut state = EpisodeState::new(); + let train_h = cli.horizon; + let train_batch = train_n_par * train_h; + let mut states_train_dev = stream.alloc_zeros::(train_batch * STATE_DIM)?; + let mut next_states_train_dev = stream.alloc_zeros::(train_batch * STATE_DIM)?; + let mut actions_train_dev = stream.alloc_zeros::(train_batch)?; + let mut rewards_train_dev = stream.alloc_zeros::(train_batch)?; + let mut dones_train_dev = stream.alloc_zeros::(train_batch)?; + let mut probs_curr_train_dev = stream.alloc_zeros::(train_batch * N_ACTIONS * c51_n_atoms)?; + let mut probs_next_train_dev = stream.alloc_zeros::(train_batch * N_ACTIONS * c51_n_atoms)?; + let mut m_train_dev = stream.alloc_zeros::(train_batch * c51_n_atoms)?; + let h_enriched_train_capacity = (train_h + 1) * train_n_par * cli.mamba2_hidden_dim; + let mut h_enriched_train_dev = stream.alloc_zeros::(h_enriched_train_capacity)?; + let batched_state_pinned_train = unsafe { MappedF32::new(train_n_par * STATE_DIM)? }; + let batched_action_pinned_train = unsafe { MappedI32::new(train_n_par)? }; + let mut batched_window_tensor_train = GpuTensor::zeros( + &[train_n_par, cli.window_k, STATE_DIM], &stream + ).map_err(|e| anyhow::anyhow!("alloc batched window train: {e}"))?; + let mut batched_probs_inference_train = stream + .alloc_zeros::(train_n_par * N_ACTIONS * c51_n_atoms)?; + let snapshots_arc_train = env.snapshots_arc(); + let fill_model_for_train = env.fill_model.clone(); + let env_config_template_train = env.config.clone(); + let mut episode_rng = SmokeRng::new(cli.seed.wrapping_add(0xFEED)); - let mut states_host: Vec = Vec::with_capacity(cli.horizon * STATE_DIM); - let mut next_states_host: Vec = Vec::with_capacity(cli.horizon * STATE_DIM); - let mut actions_host: Vec = Vec::with_capacity(cli.horizon); - let mut rewards_host: Vec = Vec::with_capacity(cli.horizon); - let mut dones_host: Vec = Vec::with_capacity(cli.horizon); - // Phase E.4.A.7 (T12): zero window at episode start. + for epoch in 0..train_n_epochs { + let _eps = cli.eps_start + (cli.eps_end - cli.eps_start) + * (epoch as f32 / train_n_epochs.max(1) as f32); + let mut par_envs: Vec = (0..train_n_par).map(|_| { + let mut cfg = env_config_template_train.clone(); + cfg.cost_per_contract = cli.train_cost; + ExecutionEnv::new_arc( + cfg, fill_model_for_train.clone(), + std::sync::Arc::clone(&snapshots_arc_train), + 0, + ) + }).collect(); + for env_i in par_envs.iter_mut() { + let start = (episode_rng.next_u64() as usize) % train_max_start; + let seed = episode_rng.next_u64(); + env_i.reset_at(seed, start); + } + let mut par_states: Vec = vec![EpisodeState::new(); train_n_par]; + let mut done_flags = vec![false; train_n_par]; if cli.temporal { - stream.memset_zeros(window_tensor.data_mut()) - .context("zero window on episode reset (train)")?; + stream.memset_zeros(batched_window_tensor_train.data_mut())?; } - loop { - let s_vec = env.state(&state).to_vec(); - let action: u8 = if cli.c51 { - state_pinned.write(&s_vec); - // Phase E.4.A.7 (T12): push and run Mamba2. - if cli.temporal { - 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, - state_pinned.dev_u64(), w_ptr, - cli.window_k as i32, state_dim_i, - )?; - } + stream.memset_zeros(&mut h_enriched_train_dev)?; + let mut states_host_train: Vec = vec![0.0; train_batch * STATE_DIM]; + let mut next_states_host_train: Vec = vec![0.0; train_batch * STATE_DIM]; + let mut actions_host_train: Vec = vec![0; train_batch]; + let mut rewards_host_train: Vec = vec![0.0; train_batch]; + let mut dones_host_train: Vec = vec![0.0; train_batch]; + + for step in 0..train_h { + let mut batched_states_host = vec![0.0_f32; train_n_par * STATE_DIM]; + for i in 0..train_n_par { + if !done_flags[i] { + let s = par_envs[i].state(&par_states[i]); + batched_states_host[i * STATE_DIM..(i + 1) * STATE_DIM].copy_from_slice(&s); } - if cli.temporal { - let block = mamba2_block.as_ref().expect("Mamba2Block missing"); - let (_logit, cache) = block.forward_train(&window_tensor) - .map_err(|e| anyhow::anyhow!("mamba2 forward: {e}"))?; - // Phase E.4.A.8.fix: GPU-side slot copy. - { - let slot_offset = (state.step * cli.mamba2_hidden_dim) as i32; - let (src_ptr, _g_src) = cache.h_enriched.cuda_data().device_ptr(&stream); - let (buf_ptr, _g_buf) = h_enriched_buf_dev.device_ptr_mut(&stream); - unsafe { - ml::cuda_pipeline::alpha_kernels::launch_alpha_h_enriched_store( - &stream, &h_store_kernel, - src_ptr, buf_ptr, slot_offset, - cli.mamba2_hidden_dim as i32, - )?; - } - } - // C51 forward on h_enriched. - 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); - unsafe { - ml::cuda_pipeline::alpha_kernels::launch_alpha_c51_forward( - &stream, &c51_fwd_kernel, - w_ptr, b_ptr, state_pinned.dev_u64(), p_ptr, - 1, c51_input_dim as i32, n_act_i, n_atoms_i, - )?; - } - } - let step_seed = episode_rng.next_u64() as u32; - { - let (p_ptr, _g0) = single_probs_dev.device_ptr(&stream); - unsafe { - ml::cuda_pipeline::alpha_kernels::launch_alpha_c51_thompson_select( - &stream, &c51_thompson_kernel, - p_ptr, - state_pinned.dev_u64(), - cli.train_threshold, - 1, - state_dim_i, - c51_v_min, c51_delta_z, - step_seed, - action_pinned.dev_u64(), - 1, n_act_i, n_atoms_i, - )?; - } - } - stream.synchronize()?; - action_pinned.read() as u8 - } else { - stream.memcpy_htod(&s_vec, &mut single_state_dev)?; - { - let (w_ptr, _g0) = w_dev.device_ptr(&stream); - let (b_ptr, _g1) = b_dev.device_ptr(&stream); - let (s_ptr, _g2) = single_state_dev.device_ptr(&stream); - let (q_ptr, _g3) = single_q_dev.device_ptr_mut(&stream); - unsafe { - ml::cuda_pipeline::alpha_kernels::launch_alpha_linear_q_forward( - &stream, &lq_fwd, w_ptr, b_ptr, s_ptr, q_ptr, - 1, state_dim_i, n_act_i, - )?; - } - } - stream.synchronize()?; - let q_host = stream.clone_dtoh(&single_q_dev)?; - epsilon_greedy_gated( - &q_host, - s_vec[1], - cli.train_threshold, - eps, - allowed_actions, - &mut episode_rng, - ) - }; - let (_s_next, reward, done) = env - .step(action, &mut state) - .ok_or_else(|| anyhow::anyhow!("step returned None"))?; - let s_next_vec = env.state(&state).to_vec(); - states_host.extend_from_slice(&s_vec); - next_states_host.extend_from_slice(&s_next_vec); - actions_host.push(action as i32); - rewards_host.push(reward); - dones_host.push(if done { 1.0 } else { 0.0 }); - if done { - break; } - } - let ep_len = actions_host.len() as i32; - if ep_len < 2 { - continue; - } - // Batched train update (same logic as alpha_dqn_h600_smoke). - let rewards_norm: Vec = rewards_host - .iter() - .map(|r| r / cli.reward_scale) - .collect(); - stream.memcpy_htod(&states_host, &mut states_dev)?; - stream.memcpy_htod(&next_states_host, &mut next_states_dev)?; - stream.memcpy_htod(&actions_host, &mut actions_dev)?; - stream.memcpy_htod(&rewards_norm, &mut rewards_dev)?; - stream.memcpy_htod(&dones_host, &mut dones_dev)?; - if cli.c51 { - // Phase E.4.A.8 (T12): in temporal mode, the batched C51 - // forwards read h_enriched_buf_dev populated during the - // per-step inference. Run one extra Mamba2 forward on the - // terminal window to fill slot ep_len. + batched_state_pinned_train.write(&batched_states_host); + if cli.temporal { + { + let (w_ptr, _g) = batched_window_tensor_train.data_mut().device_ptr_mut(&stream); + unsafe { + ml::cuda_pipeline::alpha_kernels::launch_alpha_window_push_batched( + &stream, &push_batched_kernel, + batched_state_pinned_train.dev_u64(), w_ptr, + train_n_par as i32, cli.window_k as i32, state_dim_i, + )?; + } + } 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}"))?; - // Phase E.4.A.8.fix: GPU-side slot copy for terminal h_enriched. - let term_offset = ((ep_len as usize) * cli.mamba2_hidden_dim) as i32; - let (src_ptr, _g_src) = cache.h_enriched.cuda_data().device_ptr(&stream); - let (buf_ptr, _g_buf) = h_enriched_buf_dev.device_ptr_mut(&stream); + let (_logit, cache) = block.forward_train(&batched_window_tensor_train) + .map_err(|e| anyhow::anyhow!("train mamba2: {e}"))?; + { + let (src_p, _g_src) = cache.h_enriched.cuda_data().device_ptr(&stream); + let (buf_p, _g_buf) = h_enriched_train_dev.device_ptr_mut(&stream); + let step_row_offset = (step * train_n_par) as i32; + unsafe { + ml::cuda_pipeline::alpha_kernels::launch_alpha_h_enriched_store_batched( + &stream, &h_store_batched_kernel, + src_p, buf_p, step_row_offset, + train_n_par as i32, cli.mamba2_hidden_dim as i32, + )?; + } + } + let (h_ptr, _g_h) = cache.h_enriched.cuda_data().device_ptr(&stream); + let (w_ptr, _g0) = w_dev.device_ptr(&stream); + let (b_ptr, _g1) = b_dev.device_ptr(&stream); + let (p_ptr, _g3) = batched_probs_inference_train.device_ptr_mut(&stream); unsafe { - ml::cuda_pipeline::alpha_kernels::launch_alpha_h_enriched_store( - &stream, &h_store_kernel, - src_ptr, buf_ptr, term_offset, - cli.mamba2_hidden_dim as i32, + ml::cuda_pipeline::alpha_kernels::launch_alpha_c51_forward( + &stream, &c51_fwd_kernel, + w_ptr, b_ptr, h_ptr, p_ptr, + train_n_par as i32, c51_input_dim as i32, n_act_i, n_atoms_i, )?; } - } - let (curr_input_ptr_guard, next_input_ptr_guard); - let (curr_input_ptr, next_input_ptr): (u64, u64) = if cli.temporal { - 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, - ) } 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) - }; - // C51 forward online → probs_current - { let (w_ptr, _g0) = w_dev.device_ptr(&stream); let (b_ptr, _g1) = b_dev.device_ptr(&stream); - let (p_ptr, _g3) = probs_current_dev.device_ptr_mut(&stream); + let (p_ptr, _g3) = batched_probs_inference_train.device_ptr_mut(&stream); unsafe { ml::cuda_pipeline::alpha_kernels::launch_alpha_c51_forward( &stream, &c51_fwd_kernel, - w_ptr, b_ptr, curr_input_ptr, p_ptr, - ep_len, c51_input_dim as i32, n_act_i, n_atoms_i, + w_ptr, b_ptr, batched_state_pinned_train.dev_u64(), p_ptr, + train_n_par as i32, c51_input_dim as i32, n_act_i, n_atoms_i, )?; } } - // C51 forward target → probs_next + let step_seed = episode_rng.next_u64() as u32; { - let (w_ptr, _g0) = w_target_dev.device_ptr(&stream); - let (b_ptr, _g1) = b_target_dev.device_ptr(&stream); - let (p_ptr, _g3) = probs_next_dev.device_ptr_mut(&stream); + let (p_ptr, _g0) = batched_probs_inference_train.device_ptr(&stream); unsafe { - ml::cuda_pipeline::alpha_kernels::launch_alpha_c51_forward( - &stream, &c51_fwd_kernel, - w_ptr, b_ptr, next_input_ptr, p_ptr, - ep_len, c51_input_dim as i32, n_act_i, n_atoms_i, + ml::cuda_pipeline::alpha_kernels::launch_alpha_c51_thompson_select( + &stream, &c51_thompson_kernel, + p_ptr, + batched_state_pinned_train.dev_u64(), + cli.train_threshold, + 1, state_dim_i, c51_v_min, c51_delta_z, step_seed, + batched_action_pinned_train.dev_u64(), + train_n_par as i32, n_act_i, n_atoms_i, )?; } } - { - let (pn_ptr, _g0) = probs_next_dev.device_ptr(&stream); - let (r_ptr, _g1) = rewards_dev.device_ptr(&stream); - let (d_ptr, _g2) = dones_dev.device_ptr(&stream); - let (m_ptr, _g3) = m_dev.device_ptr_mut(&stream); - unsafe { - ml::cuda_pipeline::alpha_kernels::launch_alpha_c51_project( - &stream, &c51_project_kernel, - pn_ptr, r_ptr, d_ptr, - c51_v_min, c51_v_max, cli.gamma, c51_delta_z, - m_ptr, ep_len, n_act_i, n_atoms_i, - )?; - } - } - { - let (p_ptr, _g0) = probs_current_dev.device_ptr(&stream); - let (m_ptr, _g1) = m_dev.device_ptr(&stream); - let (a_ptr, _g2) = actions_dev.device_ptr(&stream); - 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, grad_input_ptr, dw_ptr, db_ptr, - ep_len, c51_input_dim as i32, n_act_i, n_atoms_i, - 1.0 / ep_len as f32, - )?; - } - } - } else { - // Linear-Q forward/Munchausen/grad - { - 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 (q_ptr, _g3) = q_current_dev.device_ptr_mut(&stream); - unsafe { - ml::cuda_pipeline::alpha_kernels::launch_alpha_linear_q_forward( - &stream, &lq_fwd, w_ptr, b_ptr, s_ptr, q_ptr, - ep_len, state_dim_i, n_act_i, - )?; - } - } - { - 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 (q_ptr, _g3) = q_next_dev.device_ptr_mut(&stream); - unsafe { - ml::cuda_pipeline::alpha_kernels::launch_alpha_linear_q_forward( - &stream, &lq_fwd, w_ptr, b_ptr, s_ptr, q_ptr, - ep_len, state_dim_i, n_act_i, - )?; - } - } - { - let (qn_ptr, _g0) = q_next_dev.device_ptr(&stream); - let (qc_ptr, _g1) = q_current_dev.device_ptr(&stream); - let (a_ptr, _g2) = actions_dev.device_ptr(&stream); - let (r_ptr, _g3) = rewards_dev.device_ptr(&stream); - let (d_ptr, _g4) = dones_dev.device_ptr(&stream); - let (t_ptr, _g5) = target_dev.device_ptr_mut(&stream); - unsafe { - ml::cuda_pipeline::alpha_kernels::launch_alpha_munchausen_target( - &stream, &munch_kernel, - qn_ptr, qc_ptr, a_ptr, r_ptr, d_ptr, - cli.gamma, cli.alpha_m, cli.tau, cli.log_clip_min, - t_ptr, ep_len, n_act_i, - )?; - } - } - { - let (qc_ptr, _g0) = q_current_dev.device_ptr(&stream); - let (t_ptr, _g1) = target_dev.device_ptr(&stream); - let (a_ptr, _g2) = actions_dev.device_ptr(&stream); - 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); - unsafe { - ml::cuda_pipeline::alpha_kernels::launch_alpha_linear_q_grad( - &stream, &lq_grad, - qc_ptr, t_ptr, a_ptr, s_ptr, dw_ptr, db_ptr, - ep_len, state_dim_i, n_act_i, - 1.0 / ep_len as f32, - )?; + stream.synchronize()?; + let actions = batched_action_pinned_train.read_all(); + + for i in 0..train_n_par { + if done_flags[i] { continue; } + let s = &batched_states_host[i * STATE_DIM..(i + 1) * STATE_DIM]; + let row = step * train_n_par + i; + states_host_train[row * STATE_DIM..(row + 1) * STATE_DIM].copy_from_slice(s); + actions_host_train[row] = actions[i]; + let action = actions[i] as u8; + let (next_state_arr, reward, done) = par_envs[i].step(action, &mut par_states[i]) + .ok_or_else(|| anyhow::anyhow!("env step None"))?; + next_states_host_train[row * STATE_DIM..(row + 1) * STATE_DIM] + .copy_from_slice(&next_state_arr); + rewards_host_train[row] = reward / cli.reward_scale; + dones_host_train[row] = if done { 1.0 } else { 0.0 }; + if done { + done_flags[i] = true; } } } - // Clip + SGD (shared kernels; sizes from n_weights_eff / n_biases_eff) + + stream.memcpy_htod(&states_host_train, &mut states_train_dev)?; + stream.memcpy_htod(&next_states_host_train, &mut next_states_train_dev)?; + stream.memcpy_htod(&actions_host_train, &mut actions_train_dev)?; + stream.memcpy_htod(&rewards_host_train, &mut rewards_train_dev)?; + stream.memcpy_htod(&dones_host_train, &mut dones_train_dev)?; + + let bt = train_batch as i32; + let (curr_input_ptr_guard, next_input_ptr_guard); + let (curr_input_ptr, next_input_ptr): (u64, u64) = if cli.temporal { + let (h_base, g) = h_enriched_train_dev.device_ptr(&stream); + let (h_base_next, g_next) = h_enriched_train_dev.device_ptr(&stream); + curr_input_ptr_guard = g; + next_input_ptr_guard = g_next; + let _ = (&curr_input_ptr_guard, &next_input_ptr_guard); + ( + h_base, + h_base_next + (train_n_par as u64) * (cli.mamba2_hidden_dim as u64) * 4u64, + ) + } else { + let (s_ptr, g_c) = states_train_dev.device_ptr(&stream); + let (n_ptr, g_n) = next_states_train_dev.device_ptr(&stream); + curr_input_ptr_guard = g_c; + next_input_ptr_guard = g_n; + let _ = (&curr_input_ptr_guard, &next_input_ptr_guard); + (s_ptr, n_ptr) + }; + + { + let (w_ptr, _g0) = w_dev.device_ptr(&stream); + let (b_ptr, _g1) = b_dev.device_ptr(&stream); + let (p_ptr, _g3) = probs_curr_train_dev.device_ptr_mut(&stream); + unsafe { + ml::cuda_pipeline::alpha_kernels::launch_alpha_c51_forward( + &stream, &c51_fwd_kernel, + w_ptr, b_ptr, curr_input_ptr, p_ptr, + bt, c51_input_dim as i32, n_act_i, n_atoms_i, + )?; + } + } + { + let (w_ptr, _g0) = w_target_dev.device_ptr(&stream); + let (b_ptr, _g1) = b_target_dev.device_ptr(&stream); + let (p_ptr, _g3) = probs_next_train_dev.device_ptr_mut(&stream); + unsafe { + ml::cuda_pipeline::alpha_kernels::launch_alpha_c51_forward( + &stream, &c51_fwd_kernel, + w_ptr, b_ptr, next_input_ptr, p_ptr, + bt, c51_input_dim as i32, n_act_i, n_atoms_i, + )?; + } + } + { + let (pn_ptr, _g0) = probs_next_train_dev.device_ptr(&stream); + let (r_ptr, _g1) = rewards_train_dev.device_ptr(&stream); + let (d_ptr, _g2) = dones_train_dev.device_ptr(&stream); + let (m_ptr, _g3) = m_train_dev.device_ptr_mut(&stream); + unsafe { + ml::cuda_pipeline::alpha_kernels::launch_alpha_c51_project( + &stream, &c51_project_kernel, + pn_ptr, r_ptr, d_ptr, + c51_v_min, c51_v_max, cli.gamma, c51_delta_z, + m_ptr, bt, n_act_i, n_atoms_i, + )?; + } + } + { + let (p_ptr, _g0) = probs_curr_train_dev.device_ptr(&stream); + let (m_ptr, _g1) = m_train_dev.device_ptr(&stream); + let (a_ptr, _g2) = actions_train_dev.device_ptr(&stream); + let (s_ptr, _g3) = states_train_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, grad_input_ptr, dw_ptr, db_ptr, + bt, c51_input_dim as i32, n_act_i, n_atoms_i, + 1.0 / bt as f32, + )?; + } + } { let (dw_ptr, _g0) = dw_dev.device_ptr_mut(&stream); unsafe { @@ -879,16 +817,15 @@ fn main() -> Result<()> { )?; } } - // Target net hard update - if (ep + 1) % cli.target_update_every == 0 { + if (epoch + 1) % cli.target_update_every == 0 { stream.synchronize()?; let w_now = stream.clone_dtoh(&w_dev)?; let b_now = stream.clone_dtoh(&b_dev)?; stream.memcpy_htod(&w_now, &mut w_target_dev)?; stream.memcpy_htod(&b_now, &mut b_target_dev)?; } - if (ep + 1) % 200 == 0 { - info!(" train ep {}/{}", ep + 1, cli.n_train_episodes); + if (epoch + 1) % 5 == 0 || epoch + 1 == train_n_epochs { + info!(" train epoch {}/{}", epoch + 1, train_n_epochs); } } info!("Training complete. Frozen policy ready for eval."); diff --git a/crates/ml/src/cuda_pipeline/alpha_kernels.rs b/crates/ml/src/cuda_pipeline/alpha_kernels.rs index 642fe28b8..64fb6de9e 100644 --- a/crates/ml/src/cuda_pipeline/alpha_kernels.rs +++ b/crates/ml/src/cuda_pipeline/alpha_kernels.rs @@ -646,6 +646,44 @@ pub unsafe fn launch_alpha_c51_expected_q( Ok(()) } +/// Launch `alpha_h_enriched_store_batched_kernel`. Copies a batch of +/// `B` h_enriched vectors of width `hidden_dim` into a destination +/// buffer at a step-aligned row offset. Used to build the training- +/// time h_enriched buffer of shape `[H+1, B, hidden_dim]` (step-outer). +/// +/// # Safety +/// `src_dev` and `buf_dev` must be valid device pointers. +pub unsafe fn launch_alpha_h_enriched_store_batched( + stream: &cudarc::driver::CudaStream, + kernel: &cudarc::driver::CudaFunction, + src_dev: u64, + buf_dev: u64, + step_row_offset: i32, + b: i32, + hidden_dim: i32, +) -> Result<(), MLError> { + use cudarc::driver::{LaunchConfig, PushKernelArg}; + debug_assert!(b > 0 && hidden_dim > 0); + debug_assert!(step_row_offset >= 0); + const BLOCK: u32 = 32; + let grid_x = ((hidden_dim as u32) + BLOCK - 1) / BLOCK; + let cfg = LaunchConfig { + grid_dim: (grid_x.max(1), b as u32, 1), + block_dim: (BLOCK, 1, 1), + shared_mem_bytes: 0, + }; + stream + .launch_builder(kernel) + .arg(&src_dev) + .arg(&buf_dev) + .arg(&step_row_offset) + .arg(&b) + .arg(&hidden_dim) + .launch(cfg) + .map_err(|e| MLError::ModelError(format!("alpha_h_enriched_store_batched launch: {e}")))?; + Ok(()) +} + /// Launch `alpha_window_push_batched_kernel`. Same chronological /// shift+insert as `alpha_window_push_kernel` but across B parallel /// windows. One thread per (batch, feature). Used by the batched-eval diff --git a/crates/ml/src/cuda_pipeline/alpha_window_push.cu b/crates/ml/src/cuda_pipeline/alpha_window_push.cu index b60988861..9725ec5a7 100644 --- a/crates/ml/src/cuda_pipeline/alpha_window_push.cu +++ b/crates/ml/src/cuda_pipeline/alpha_window_push.cu @@ -72,3 +72,25 @@ extern "C" __global__ void alpha_window_push_batched_kernel( } windows[base + (K - 1) * state_dim + j] = states_in[b * state_dim + j]; } + +// ---------------------------------------------------------------------- +// Phase E.4.A.T15.batched (2026-05-15): batched h_enriched store for +// training-time transition collection. Copies cache.h_enriched +// [B, hidden_dim] into buf[step_offset_in_rows .. step_offset+B] of +// a [H+1, B, hidden_dim] buffer (step-outer layout, contiguous block +// per step). One thread per (env, feature). +// ---------------------------------------------------------------------- +extern "C" __global__ void alpha_h_enriched_store_batched_kernel( + const float* __restrict__ src, // [B, hidden_dim] + float* __restrict__ buf, // [..., B, hidden_dim] + int step_row_offset, // step * B (in rows of width hidden_dim) + int B, + int hidden_dim +) { + const int j = blockIdx.x * blockDim.x + threadIdx.x; + const int b = blockIdx.y; + if (b >= B || j >= hidden_dim) return; + const int dst_idx = (step_row_offset + b) * hidden_dim + j; + const int src_idx = b * hidden_dim + j; + buf[dst_idx] = src[src_idx]; +} diff --git a/docs/isv-slots.md b/docs/isv-slots.md index 7fdae29c3..cd8c6cb26 100644 --- a/docs/isv-slots.md +++ b/docs/isv-slots.md @@ -797,6 +797,19 @@ Deferred pending backtest validation — if frozen-Mamba2 already lifts backtest Sharpe, T10 becomes optimisation rather than prerequisite. +### Phase E.4.A.T15.batched — parallel-env TRAINING (2026-05-15) + +Training mirrors the T14 eval batching: N_par parallel envs run +lockstep H steps per epoch; ONE batched C51 update at B = N_par * H +per epoch. NEW kernel `alpha_h_enriched_store_batched_kernel` +copies cache.h_enriched [B, hidden] into a [H+1, B, hidden] +step-outer training buffer. The terminal-step slot stays zero — +done=1 at step H-1 masks Q_next contribution in the Bellman +projection so no terminal Mamba2 forward is needed. + +Pure-GPU per-step. No CPU compute. No per-step CPU↔GPU roundtrips +on the hot path beyond the unavoidable env.step calls. + ### Phase E.4.A.T14.batched — parallel-env eval (2026-05-15) The backtest binary's 2D-sweep eval loop was rewritten from