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:
@@ -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.");
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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];
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user