feat(phase-e-4-a): batched parallel-env TRAINING (greenfield)

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 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-05-15 23:11:01 +02:00
parent 90c9d54454
commit 0ab54dc8ef
4 changed files with 296 additions and 286 deletions

View File

@@ -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::<f32>(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::<f32>(train_batch * STATE_DIM)?;
let mut next_states_train_dev = stream.alloc_zeros::<f32>(train_batch * STATE_DIM)?;
let mut actions_train_dev = stream.alloc_zeros::<i32>(train_batch)?;
let mut rewards_train_dev = stream.alloc_zeros::<f32>(train_batch)?;
let mut dones_train_dev = stream.alloc_zeros::<f32>(train_batch)?;
let mut probs_curr_train_dev = stream.alloc_zeros::<f32>(train_batch * N_ACTIONS * c51_n_atoms)?;
let mut probs_next_train_dev = stream.alloc_zeros::<f32>(train_batch * N_ACTIONS * c51_n_atoms)?;
let mut m_train_dev = stream.alloc_zeros::<f32>(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::<f32>(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::<f32>(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<f32> = Vec::with_capacity(cli.horizon * STATE_DIM);
let mut next_states_host: Vec<f32> = Vec::with_capacity(cli.horizon * STATE_DIM);
let mut actions_host: Vec<i32> = Vec::with_capacity(cli.horizon);
let mut rewards_host: Vec<f32> = Vec::with_capacity(cli.horizon);
let mut dones_host: Vec<f32> = 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<ExecutionEnv> = (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<EpisodeState> = 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<f32> = vec![0.0; train_batch * STATE_DIM];
let mut next_states_host_train: Vec<f32> = vec![0.0; train_batch * STATE_DIM];
let mut actions_host_train: Vec<i32> = vec![0; train_batch];
let mut rewards_host_train: Vec<f32> = vec![0.0; train_batch];
let mut dones_host_train: Vec<f32> = 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<f32> = 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.");

View File

@@ -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

View File

@@ -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];
}