@@ -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 . w rapping_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 + t rain_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 < i 32> = 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 f 32 / 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 (
& s tream , & 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 ) * 4 u64 ,
)
} 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 _pt r( & 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_pt r , 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 _pa r {
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_ar r , 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 ) * 4 u64 ,
)
} 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_episode s) ;
if ( epoch + 1 ) % 5 = = 0 | | epoch + 1 = = train_n_epochs {
info! ( " train epoch {}/{} " , epoch + 1 , train_n_ epoch s ) ;
}
}
info! ( " Training complete. Frozen policy ready for eval. " ) ;