diff --git a/crates/ml-alpha/src/trainer/integrated.rs b/crates/ml-alpha/src/trainer/integrated.rs index cd741df50..1118550e7 100644 --- a/crates/ml-alpha/src/trainer/integrated.rs +++ b/crates/ml-alpha/src/trainer/integrated.rs @@ -102,6 +102,7 @@ use crate::rl::ppo::{PolicyHead, PpoHeadsConfig, ValueHead}; use crate::rl::reward::RlLobBackend; use crate::trainer::optim::AdamW; use crate::trainer::perception::{PerceptionTrainer, PerceptionTrainerConfig}; +use crate::trainer::raw_launch::{RawArgs, raw_launch}; const GRAD_H_ACCUMULATE_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/grad_h_accumulate.cubin")); @@ -2855,24 +2856,22 @@ impl IntegratedTrainer { &self, b_size: usize, ) -> Result<()> { - let cfg = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (1, 1, 1), - shared_mem_bytes: 0, - }; let alpha = RL_LR_CONTROLLER_ALPHA; let b_size_i = b_size as i32; - let mut launch = self.stream.launch_builder(&self.rl_fused_controllers_fn); - launch - .arg(&self.isv_d) - .arg(&self.dones_d) - .arg(&alpha) - .arg(&b_size_i) - .arg(&self.fused_ctrl_input_slots_d); + let mut args = RawArgs::new(); + args.push_ptr(self.isv_d.raw_ptr()); + args.push_ptr(self.dones_d.raw_ptr()); + args.push_f32(alpha); + args.push_i32(b_size_i); + args.push_ptr(self.fused_ctrl_input_slots_d.raw_ptr()); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch - .launch(cfg) - .context("rl_fused_controllers launch")?; + raw_launch( + self.rl_fused_controllers_fn.cu_function(), + (1, 1, 1), (1, 1, 1), 0, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_fused_controllers: {:?}", e))?; } Ok(()) } @@ -2899,23 +2898,24 @@ impl IntegratedTrainer { debug_assert_eq!(obs_d.len(), b_size); debug_assert_eq!(dones_d.len(), b_size); debug_assert!(b_size <= 1024, "ema_update_on_done: b_size > 1024 needs multi-block reduce"); - let cfg = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (b_size as u32, 1, 1), - shared_mem_bytes: (2 * b_size * std::mem::size_of::()) as u32, - }; + let smem = (2 * b_size * std::mem::size_of::()) as u32; let slot_i = slot_index as i32; let b_size_i = b_size as i32; - let mut launch = self.stream.launch_builder(&self.ema_update_on_done_fn); - launch - .arg(&self.isv_d) - .arg(&slot_i) - .arg(&alpha) - .arg(obs_d) - .arg(dones_d) - .arg(&b_size_i); + let mut args = RawArgs::new(); + args.push_ptr(self.isv_d.raw_ptr()); + args.push_i32(slot_i); + args.push_f32(alpha); + args.push_ptr(obs_d.raw_ptr()); + args.push_ptr(dones_d.raw_ptr()); + args.push_i32(b_size_i); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch.launch(cfg).context("ema_update_on_done launch")?; + raw_launch( + self.ema_update_on_done_fn.cu_function(), + (1, 1, 1), (b_size as u32, 1, 1), smem, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("ema_update_on_done: {:?}", e))?; } Ok(()) } @@ -2934,22 +2934,23 @@ impl IntegratedTrainer { ) -> Result<()> { debug_assert_eq!(obs_d.len(), b_size); debug_assert!(b_size <= 1024, "ema_update_per_step: b_size > 1024 needs multi-block reduce"); - let cfg = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (b_size as u32, 1, 1), - shared_mem_bytes: (b_size * std::mem::size_of::()) as u32, - }; + let smem = (b_size * std::mem::size_of::()) as u32; let slot_i = slot_index as i32; let b_size_i = b_size as i32; - let mut launch = self.stream.launch_builder(&self.ema_update_per_step_fn); - launch - .arg(&self.isv_d) - .arg(&slot_i) - .arg(&alpha) - .arg(obs_d) - .arg(&b_size_i); + let mut args = RawArgs::new(); + args.push_ptr(self.isv_d.raw_ptr()); + args.push_i32(slot_i); + args.push_f32(alpha); + args.push_ptr(obs_d.raw_ptr()); + args.push_i32(b_size_i); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch.launch(cfg).context("ema_update_per_step launch")?; + raw_launch( + self.ema_update_per_step_fn.cu_function(), + (1, 1, 1), (b_size as u32, 1, 1), smem, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("ema_update_per_step: {:?}", e))?; } Ok(()) } @@ -4658,21 +4659,19 @@ impl IntegratedTrainer { { let b_i = b_size as i32; let n_tau_i = n_tau as i32; - let tau_cfg = LaunchConfig { - grid_dim: (b_size as u32, 1, 1), - block_dim: (n_tau as u32, 1, 1), - shared_mem_bytes: 0, - }; - let mut tau_launch = self.stream.launch_builder(&self.rl_sample_tau_fn); - tau_launch - .arg(&mut self.iqn_prng_state_d) - .arg(&mut self.iqn_tau_d) - .arg(&b_i) - .arg(&n_tau_i); + let mut args = RawArgs::new(); + args.push_ptr(self.iqn_prng_state_d.raw_ptr()); + args.push_ptr(self.iqn_tau_d.raw_ptr()); + args.push_i32(b_i); + args.push_i32(n_tau_i); + let mut ptrs = args.build_arg_ptrs(); unsafe { - tau_launch - .launch(tau_cfg) - .context("dqn_replay_step: rl_sample_tau online")?; + raw_launch( + self.rl_sample_tau_fn.cu_function(), + (b_size as u32, 1, 1), (n_tau as u32, 1, 1), 0, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("dqn_replay_step: rl_sample_tau online: {:?}", e))?; } } @@ -4725,21 +4724,19 @@ impl IntegratedTrainer { { let b_i = b_size as i32; let n_tau_i = n_tau as i32; - let tau_cfg = LaunchConfig { - grid_dim: (b_size as u32, 1, 1), - block_dim: (n_tau as u32, 1, 1), - shared_mem_bytes: 0, - }; - let mut tau_launch = self.stream.launch_builder(&self.rl_sample_tau_fn); - tau_launch - .arg(&mut self.iqn_prng_state_d) - .arg(&mut self.ss_iqn_tau_target_d) - .arg(&b_i) - .arg(&n_tau_i); + let mut args = RawArgs::new(); + args.push_ptr(self.iqn_prng_state_d.raw_ptr()); + args.push_ptr(self.ss_iqn_tau_target_d.raw_ptr()); + args.push_i32(b_i); + args.push_i32(n_tau_i); + let mut ptrs = args.build_arg_ptrs(); unsafe { - tau_launch - .launch(tau_cfg) - .context("dqn_replay_step: rl_sample_tau target")?; + raw_launch( + self.rl_sample_tau_fn.cu_function(), + (b_size as u32, 1, 1), (n_tau as u32, 1, 1), 0, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("dqn_replay_step: rl_sample_tau target: {:?}", e))?; } } @@ -4810,21 +4807,19 @@ impl IntegratedTrainer { // ── 2. Double-DQN argmax on online Q at h_tp1 ─────────────── { - let cfg_argmax = LaunchConfig { - grid_dim: (b_size as u32, 1, 1), - block_dim: (N_ACTIONS as u32, 1, 1), - shared_mem_bytes: 0, - }; - let mut launch = self.stream.launch_builder(&self.argmax_expected_q_fn); - launch - .arg(&self.q_logits_tp1_d) - .arg(&self.atom_supports_d) - .arg(&mut self.sampled_next_actions_d) - .arg(&b_size_i); + let mut args = RawArgs::new(); + args.push_ptr(self.q_logits_tp1_d.raw_ptr()); + args.push_ptr(self.atom_supports_d.raw_ptr()); + args.push_ptr(self.sampled_next_actions_d.raw_ptr()); + args.push_i32(b_size_i); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch - .launch(cfg_argmax) - .context("dqn_replay_step: argmax_expected_q launch")?; + raw_launch( + self.argmax_expected_q_fn.cu_function(), + (b_size as u32, 1, 1), (N_ACTIONS as u32, 1, 1), 0, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("dqn_replay_step: argmax_expected_q: {:?}", e))?; } } @@ -4904,22 +4899,21 @@ impl IntegratedTrainer { let dqn_rows = (N_ACTIONS * Q_N_ATOMS) as i32; let dqn_cols = HIDDEN_DIM as i32; let block_x = (dqn_rows as u32).min(256); - // Shared: cols + block_x floats for power iteration. let smem = ((HIDDEN_DIM as u32) + block_x) * (std::mem::size_of::() as u32); - let cfg = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (block_x, 1, 1), - shared_mem_bytes: smem, - }; - let mut launch = self.stream.launch_builder(&self.rl_spectral_norm_fn); - launch - .arg(&mut self.dqn_head.w_d) - .arg(&mut self.spectral_v_buffer_dqn) - .arg(&self.isv_d) - .arg(&dqn_rows) - .arg(&dqn_cols); + let mut args = RawArgs::new(); + args.push_ptr(self.dqn_head.w_d.raw_ptr()); + args.push_ptr(self.spectral_v_buffer_dqn.raw_ptr()); + args.push_ptr(self.isv_d.raw_ptr()); + args.push_i32(dqn_rows); + args.push_i32(dqn_cols); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch.launch(cfg).context("rl_spectral_norm(dqn_w) launch")?; + raw_launch( + self.rl_spectral_norm_fn.cu_function(), + (1, 1, 1), (block_x, 1, 1), smem, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_spectral_norm(dqn_w): {:?}", e))?; } } { @@ -4927,20 +4921,20 @@ impl IntegratedTrainer { let iqn_cols = N_ACTIONS as i32; let block_x = (iqn_rows as u32).min(256); let smem = ((N_ACTIONS as u32) + block_x) * (std::mem::size_of::() as u32); - let cfg = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (block_x, 1, 1), - shared_mem_bytes: smem, - }; - let mut launch = self.stream.launch_builder(&self.rl_spectral_norm_fn); - launch - .arg(&mut self.iqn_head.w_out_d) - .arg(&mut self.spectral_v_buffer_iqn) - .arg(&self.isv_d) - .arg(&iqn_rows) - .arg(&iqn_cols); + let mut args = RawArgs::new(); + args.push_ptr(self.iqn_head.w_out_d.raw_ptr()); + args.push_ptr(self.spectral_v_buffer_iqn.raw_ptr()); + args.push_ptr(self.isv_d.raw_ptr()); + args.push_i32(iqn_rows); + args.push_i32(iqn_cols); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch.launch(cfg).context("rl_spectral_norm(iqn_w_out) launch")?; + raw_launch( + self.rl_spectral_norm_fn.cu_function(), + (1, 1, 1), (block_x, 1, 1), smem, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_spectral_norm(iqn_w_out): {:?}", e))?; } } { @@ -4948,20 +4942,20 @@ impl IntegratedTrainer { let pi_cols = HIDDEN_DIM as i32; let block_x = (pi_rows as u32).min(256); let smem = ((HIDDEN_DIM as u32) + block_x) * (std::mem::size_of::() as u32); - let cfg = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (block_x, 1, 1), - shared_mem_bytes: smem, - }; - let mut launch = self.stream.launch_builder(&self.rl_spectral_norm_fn); - launch - .arg(&mut self.policy_head.w_d) - .arg(&mut self.spectral_v_buffer_pi) - .arg(&self.isv_d) - .arg(&pi_rows) - .arg(&pi_cols); + let mut args = RawArgs::new(); + args.push_ptr(self.policy_head.w_d.raw_ptr()); + args.push_ptr(self.spectral_v_buffer_pi.raw_ptr()); + args.push_ptr(self.isv_d.raw_ptr()); + args.push_i32(pi_rows); + args.push_i32(pi_cols); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch.launch(cfg).context("rl_spectral_norm(policy_w) launch")?; + raw_launch( + self.rl_spectral_norm_fn.cu_function(), + (1, 1, 1), (block_x, 1, 1), smem, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_spectral_norm(policy_w): {:?}", e))?; } } @@ -4969,21 +4963,20 @@ impl IntegratedTrainer { // Track mean(Q_predicted - actual_return), emit correction to ISV. { let block_x = (b_size as u32).min(256).max(1); - let cfg = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (block_x, 1, 1), - shared_mem_bytes: block_x * (std::mem::size_of::() as u32), - }; - // q_predicted: use ensemble_q_d (expected Q at sampled action). - // actual_return: use sampled_rewards_d (replay reward signal). - let mut launch = self.stream.launch_builder(&self.rl_q_bias_correction_fn); - launch - .arg(&self.ensemble_q_d) - .arg(&self.sampled_rewards_d) - .arg(&mut self.isv_d) - .arg(&b_size_i); + let smem = block_x * (std::mem::size_of::() as u32); + let mut args = RawArgs::new(); + args.push_ptr(self.ensemble_q_d.raw_ptr()); + args.push_ptr(self.sampled_rewards_d.raw_ptr()); + args.push_ptr(self.isv_d.raw_ptr()); + args.push_i32(b_size_i); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch.launch(cfg).context("rl_q_bias_correction launch")?; + raw_launch( + self.rl_q_bias_correction_fn.cu_function(), + (1, 1, 1), (block_x, 1, 1), smem, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_q_bias_correction: {:?}", e))?; } } @@ -5017,20 +5010,20 @@ impl IntegratedTrainer { // ── 9. Per-branch LR controller ────────────────────────────── // Single-thread kernel adjusts LR scale factors from loss deltas. { - let cfg = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (1, 1, 1), - shared_mem_bytes: 0, - }; - let mut launch = self.stream.launch_builder(&self.rl_per_branch_lr_fn); - launch - .arg(&mut self.isv_d) - .arg(&self.last_q_loss) - .arg(&self.last_pi_loss) - .arg(&self.last_v_loss) - .arg(&l_q); // IQN loss proxied by this step's Q loss + let mut args = RawArgs::new(); + args.push_ptr(self.isv_d.raw_ptr()); + args.push_f32(self.last_q_loss); + args.push_f32(self.last_pi_loss); + args.push_f32(self.last_v_loss); + args.push_f32(l_q); // IQN loss proxied by this step's Q loss + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch.launch(cfg).context("rl_per_branch_lr launch")?; + raw_launch( + self.rl_per_branch_lr_fn.cu_function(), + (1, 1, 1), (1, 1, 1), 0, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_per_branch_lr: {:?}", e))?; } } @@ -5131,17 +5124,16 @@ impl IntegratedTrainer { // Must run BEFORE any kernel that reads current_step from ISV. // Single thread, single block — graph-safe (no scalar args change). { - let cfg = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (1, 1, 1), - shared_mem_bytes: 0, - }; - let mut launch = self.stream.launch_builder(&self.rl_increment_step_fn); - launch.arg(&self.isv_d); + let mut args = RawArgs::new(); + args.push_ptr(self.isv_d.raw_ptr()); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch - .launch(cfg) - .context("rl_increment_step launch")?; + raw_launch( + self.rl_increment_step_fn.cu_function(), + (1, 1, 1), (1, 1, 1), 0, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_increment_step: {:?}", e))?; } } @@ -5572,17 +5564,17 @@ impl IntegratedTrainer { // Launched OUTSIDE graph so the scalar write doesn't break capture. { let ts_ns_val: u64 = last_snap.ts_ns; - let cfg = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (1, 1, 1), - shared_mem_bytes: 0, - }; - let mut launch = self.stream.launch_builder(&self.rl_write_u64_fn); - launch - .arg(&mut self.ts_ns_d) - .arg(&ts_ns_val); + let mut args = RawArgs::new(); + args.push_ptr(self.ts_ns_d.raw_ptr()); + args.push_u64(ts_ns_val); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch.launch(cfg).context("rl_write_u64 (ts_ns) launch")?; + raw_launch( + self.rl_write_u64_fn.cu_function(), + (1, 1, 1), (1, 1, 1), 0, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_write_u64 (ts_ns): {:?}", e))?; } } @@ -6243,29 +6235,28 @@ impl IntegratedTrainer { let pos_d_ref = lobsim.pos_d(); let bid_px_d = lobsim.bid_px_d(); let ask_px_d = lobsim.ask_px_d(); - let cfg = LaunchConfig { - grid_dim: (((b_size as u32) + 31) / 32, 1, 1), - block_dim: (32, 1, 1), - shared_mem_bytes: 0, - }; + let grid_x = ((b_size as u32) + 31) / 32; let b_size_i = b_size as i32; let pos_bytes_i = lobsim.pos_bytes() as i32; - let mut launch = self.stream.launch_builder(&self.rl_hindsight_track_fn); - launch - .arg(pos_d_ref) - .arg(bid_px_d) - .arg(ask_px_d) - .arg(&mut self.hindsight.mid_ring_d) - .arg(&mut self.hindsight.ring_write_idx_d) - .arg(&mut self.hindsight.peak_mid_d) - .arg(&mut self.hindsight.entry_mid_d) - .arg(&mut self.hindsight.position_dir_d) - .arg(&b_size_i) - .arg(&pos_bytes_i); + let mut args = RawArgs::new(); + args.push_ptr(pos_d_ref.raw_ptr()); + args.push_ptr(bid_px_d.raw_ptr()); + args.push_ptr(ask_px_d.raw_ptr()); + args.push_ptr(self.hindsight.mid_ring_d.raw_ptr()); + args.push_ptr(self.hindsight.ring_write_idx_d.raw_ptr()); + args.push_ptr(self.hindsight.peak_mid_d.raw_ptr()); + args.push_ptr(self.hindsight.entry_mid_d.raw_ptr()); + args.push_ptr(self.hindsight.position_dir_d.raw_ptr()); + args.push_i32(b_size_i); + args.push_i32(pos_bytes_i); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch - .launch(cfg) - .context("step_with_lobsim: rl_hindsight_track")?; + raw_launch( + self.rl_hindsight_track_fn.cu_function(), + (grid_x, 1, 1), (32, 1, 1), 0, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_hindsight_track: {:?}", e))?; } } @@ -6304,30 +6295,28 @@ impl IntegratedTrainer { // Phase 1: ring write + flush decision { let b_size_i = b_size as i32; - let cfg = LaunchConfig { - grid_dim: (b_size as u32, 1, 1), - block_dim: (128, 1, 1), - shared_mem_bytes: 0, - }; - let mut launch = self.stream.launch_builder(&self.rl_per_push_ring_fn); - launch - .arg(self.perception.h_t_view()) - .arg(&self.rewards_d) - .arg(&self.raw_rewards_d) - .arg(&self.dones_d) - .arg(&self.log_pi_old_d) - .arg(&self.actions_d) - .arg(&self.isv_d) - .arg(&mut self.gpu_replay.nstep_scalars_d) - .arg(&mut self.gpu_replay.nstep_h_t_d) - .arg(&mut self.gpu_replay.nstep_write_idx_d) - .arg(&mut self.gpu_replay.nstep_count_d) - .arg(&mut self.gpu_replay.flush_flags_d) - .arg(&b_size_i); + let mut args = RawArgs::new(); + args.push_ptr(self.perception.h_t_view().raw_ptr()); + args.push_ptr(self.rewards_d.raw_ptr()); + args.push_ptr(self.raw_rewards_d.raw_ptr()); + args.push_ptr(self.dones_d.raw_ptr()); + args.push_ptr(self.log_pi_old_d.raw_ptr()); + args.push_ptr(self.actions_d.raw_ptr()); + args.push_ptr(self.isv_d.raw_ptr()); + args.push_ptr(self.gpu_replay.nstep_scalars_d.raw_ptr()); + args.push_ptr(self.gpu_replay.nstep_h_t_d.raw_ptr()); + args.push_ptr(self.gpu_replay.nstep_write_idx_d.raw_ptr()); + args.push_ptr(self.gpu_replay.nstep_count_d.raw_ptr()); + args.push_ptr(self.gpu_replay.flush_flags_d.raw_ptr()); + args.push_i32(b_size_i); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch - .launch(cfg) - .context("step_with_lobsim: rl_per_push_ring")?; + raw_launch( + self.rl_per_push_ring_fn.cu_function(), + (b_size as u32, 1, 1), (128, 1, 1), 0, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_per_push_ring: {:?}", e))?; } } // Phase 2: coordinated flush (coalesced replay write). @@ -6339,41 +6328,39 @@ impl IntegratedTrainer { { let b_size_i = b_size as i32; let cap_i = self.gpu_replay.capacity as i32; - let cfg = LaunchConfig { - grid_dim: (b_size as u32, 1, 1), - block_dim: (128, 1, 1), - shared_mem_bytes: 0, - }; // nstep_count_d appears twice in the kernel signature: once as // const read (arg 5) and once as mutable reset (arg 17). Same // device pointer — extract raw to satisfy the borrow checker. let nstep_count_ptr = self.gpu_replay.nstep_count_d.raw_ptr(); - let mut launch = self.stream.launch_builder(&self.rl_per_push_flush_fn); - launch - .arg(&self.gpu_replay.flush_flags_d) - .arg(&mut self.gpu_replay.write_offsets_d) - .arg(&self.gpu_replay.nstep_scalars_d) - .arg(&self.gpu_replay.nstep_h_t_d) - .arg(&self.gpu_replay.nstep_write_idx_d) - .arg(&nstep_count_ptr) - .arg(&self.h_tp1_d) - .arg(&self.log_pi_old_d) - .arg(&self.actions_d) - .arg(&self.isv_d) - .arg(&mut self.gpu_replay.h_t_d) - .arg(&mut self.gpu_replay.h_tp1_d) - .arg(&mut self.gpu_replay.scalars_d) - .arg(&mut self.gpu_replay.priority_tree_d) - .arg(&mut self.gpu_replay.write_head_d) - .arg(&mut self.gpu_replay.replay_len_d) - .arg(&mut self.gpu_replay.max_priority_d) - .arg(&nstep_count_ptr) - .arg(&b_size_i) - .arg(&cap_i); + let mut args = RawArgs::new(); + args.push_ptr(self.gpu_replay.flush_flags_d.raw_ptr()); + args.push_ptr(self.gpu_replay.write_offsets_d.raw_ptr()); + args.push_ptr(self.gpu_replay.nstep_scalars_d.raw_ptr()); + args.push_ptr(self.gpu_replay.nstep_h_t_d.raw_ptr()); + args.push_ptr(self.gpu_replay.nstep_write_idx_d.raw_ptr()); + args.push_ptr(nstep_count_ptr); + args.push_ptr(self.h_tp1_d.raw_ptr()); + args.push_ptr(self.log_pi_old_d.raw_ptr()); + args.push_ptr(self.actions_d.raw_ptr()); + args.push_ptr(self.isv_d.raw_ptr()); + args.push_ptr(self.gpu_replay.h_t_d.raw_ptr()); + args.push_ptr(self.gpu_replay.h_tp1_d.raw_ptr()); + args.push_ptr(self.gpu_replay.scalars_d.raw_ptr()); + args.push_ptr(self.gpu_replay.priority_tree_d.raw_ptr()); + args.push_ptr(self.gpu_replay.write_head_d.raw_ptr()); + args.push_ptr(self.gpu_replay.replay_len_d.raw_ptr()); + args.push_ptr(self.gpu_replay.max_priority_d.raw_ptr()); + args.push_ptr(nstep_count_ptr); + args.push_i32(b_size_i); + args.push_i32(cap_i); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch - .launch(cfg) - .context("step_with_lobsim: rl_per_push_flush")?; + raw_launch( + self.rl_per_push_flush_fn.cu_function(), + (b_size as u32, 1, 1), (128, 1, 1), 0, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_per_push_flush: {:?}", e))?; } } @@ -6383,40 +6370,39 @@ impl IntegratedTrainer { { let b_size_i = b_size as i32; let cap_i = self.gpu_replay.capacity as i32; - let cfg = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (b_size as u32, 1, 1), - shared_mem_bytes: (2 * b_size * std::mem::size_of::()) as u32, - }; + let smem = (2 * b_size * std::mem::size_of::()) as u32; // peak_mid_d appears twice: const read (arg 2) and mutable // reset (arg 15). Same device pointer — extract raw to // satisfy the borrow checker (same pattern as nstep_count_d // in rl_per_push_flush). let peak_mid_ptr = self.hindsight.peak_mid_d.raw_ptr(); - let mut launch = self.stream.launch_builder(&self.rl_hindsight_inject_fn); - launch - .arg(&self.dones_d) - .arg(&self.raw_rewards_d) - .arg(&peak_mid_ptr) - .arg(&self.hindsight.entry_mid_d) - .arg(&self.hindsight.position_dir_d) - .arg(self.perception.h_t_view()) - .arg(&self.isv_d) - .arg(&mut self.gpu_replay.h_t_d) - .arg(&mut self.gpu_replay.h_tp1_d) - .arg(&mut self.gpu_replay.scalars_d) - .arg(&mut self.gpu_replay.priority_tree_d) - .arg(&mut self.gpu_replay.write_head_d) - .arg(&mut self.gpu_replay.replay_len_d) - .arg(&mut self.gpu_replay.max_priority_d) - .arg(&mut self.hindsight.ring_write_idx_d) - .arg(&peak_mid_ptr) - .arg(&b_size_i) - .arg(&cap_i); + let mut args = RawArgs::new(); + args.push_ptr(self.dones_d.raw_ptr()); + args.push_ptr(self.raw_rewards_d.raw_ptr()); + args.push_ptr(peak_mid_ptr); + args.push_ptr(self.hindsight.entry_mid_d.raw_ptr()); + args.push_ptr(self.hindsight.position_dir_d.raw_ptr()); + args.push_ptr(self.perception.h_t_view().raw_ptr()); + args.push_ptr(self.isv_d.raw_ptr()); + args.push_ptr(self.gpu_replay.h_t_d.raw_ptr()); + args.push_ptr(self.gpu_replay.h_tp1_d.raw_ptr()); + args.push_ptr(self.gpu_replay.scalars_d.raw_ptr()); + args.push_ptr(self.gpu_replay.priority_tree_d.raw_ptr()); + args.push_ptr(self.gpu_replay.write_head_d.raw_ptr()); + args.push_ptr(self.gpu_replay.replay_len_d.raw_ptr()); + args.push_ptr(self.gpu_replay.max_priority_d.raw_ptr()); + args.push_ptr(self.hindsight.ring_write_idx_d.raw_ptr()); + args.push_ptr(peak_mid_ptr); + args.push_i32(b_size_i); + args.push_i32(cap_i); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch - .launch(cfg) - .context("step_with_lobsim: rl_hindsight_inject")?; + raw_launch( + self.rl_hindsight_inject_fn.cu_function(), + (1, 1, 1), (b_size as u32, 1, 1), smem, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_hindsight_inject: {:?}", e))?; } } @@ -6426,32 +6412,31 @@ impl IntegratedTrainer { { let bid_px_d = lobsim.bid_px_d(); let ask_px_d = lobsim.ask_px_d(); - let cfg = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (256, 1, 1), // CLOSED_RING_SIZE - shared_mem_bytes: (2 * 256 * std::mem::size_of::()) as u32, - }; + let smem = (2 * 256 * std::mem::size_of::()) as u32; let cap_i = self.gpu_replay.capacity as i32; - let mut launch = self.stream.launch_builder(&self.rl_hindsight_forward_fn); - launch - .arg(bid_px_d) - .arg(ask_px_d) - .arg(&self.isv_d) - .arg(&mut self.hindsight.closed_ring_d) - .arg(&self.hindsight.closed_h_t_d) - .arg(&self.hindsight.closed_step_d) - .arg(&mut self.gpu_replay.h_t_d) - .arg(&mut self.gpu_replay.h_tp1_d) - .arg(&mut self.gpu_replay.scalars_d) - .arg(&mut self.gpu_replay.priority_tree_d) - .arg(&mut self.gpu_replay.write_head_d) - .arg(&mut self.gpu_replay.replay_len_d) - .arg(&mut self.gpu_replay.max_priority_d) - .arg(&cap_i); + let mut args = RawArgs::new(); + args.push_ptr(bid_px_d.raw_ptr()); + args.push_ptr(ask_px_d.raw_ptr()); + args.push_ptr(self.isv_d.raw_ptr()); + args.push_ptr(self.hindsight.closed_ring_d.raw_ptr()); + args.push_ptr(self.hindsight.closed_h_t_d.raw_ptr()); + args.push_ptr(self.hindsight.closed_step_d.raw_ptr()); + args.push_ptr(self.gpu_replay.h_t_d.raw_ptr()); + args.push_ptr(self.gpu_replay.h_tp1_d.raw_ptr()); + args.push_ptr(self.gpu_replay.scalars_d.raw_ptr()); + args.push_ptr(self.gpu_replay.priority_tree_d.raw_ptr()); + args.push_ptr(self.gpu_replay.write_head_d.raw_ptr()); + args.push_ptr(self.gpu_replay.replay_len_d.raw_ptr()); + args.push_ptr(self.gpu_replay.max_priority_d.raw_ptr()); + args.push_i32(cap_i); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch - .launch(cfg) - .context("step_with_lobsim: rl_hindsight_forward")?; + raw_launch( + self.rl_hindsight_forward_fn.cu_function(), + (1, 1, 1), (256, 1, 1), smem, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_hindsight_forward: {:?}", e))?; } } @@ -6533,34 +6518,32 @@ impl IntegratedTrainer { { let b_size_i = b_size as i32; let cap_i = self.gpu_replay.capacity as i32; - let cfg = LaunchConfig { - grid_dim: (b_size as u32, 1, 1), - block_dim: (1, 1, 1), - shared_mem_bytes: 0, - }; - let mut launch = self.train_stream.launch_builder(&self.rl_per_sample_fn); - launch - .arg(&self.gpu_replay.priority_tree_d) - .arg(&self.gpu_replay.h_t_d) - .arg(&self.gpu_replay.h_tp1_d) - .arg(&self.gpu_replay.scalars_d) - .arg(&self.gpu_replay.replay_len_d) - .arg(&self.isv_d) - .arg(&mut self.gpu_replay.sample_prng_d) - .arg(&mut self.sampled_h_t_d) - .arg(&mut self.sampled_h_tp1_d) - .arg(&mut self.sampled_rewards_d) - .arg(&mut self.sampled_dones_d) - .arg(&mut self.sampled_log_pi_old_d) - .arg(&mut self.sampled_n_step_gammas_d) - .arg(&mut self.sampled_actions_d) - .arg(&mut self.gpu_replay.sample_indices_d) - .arg(&b_size_i) - .arg(&cap_i); + let mut args = RawArgs::new(); + args.push_ptr(self.gpu_replay.priority_tree_d.raw_ptr()); + args.push_ptr(self.gpu_replay.h_t_d.raw_ptr()); + args.push_ptr(self.gpu_replay.h_tp1_d.raw_ptr()); + args.push_ptr(self.gpu_replay.scalars_d.raw_ptr()); + args.push_ptr(self.gpu_replay.replay_len_d.raw_ptr()); + args.push_ptr(self.isv_d.raw_ptr()); + args.push_ptr(self.gpu_replay.sample_prng_d.raw_ptr()); + args.push_ptr(self.sampled_h_t_d.raw_ptr()); + args.push_ptr(self.sampled_h_tp1_d.raw_ptr()); + args.push_ptr(self.sampled_rewards_d.raw_ptr()); + args.push_ptr(self.sampled_dones_d.raw_ptr()); + args.push_ptr(self.sampled_log_pi_old_d.raw_ptr()); + args.push_ptr(self.sampled_n_step_gammas_d.raw_ptr()); + args.push_ptr(self.sampled_actions_d.raw_ptr()); + args.push_ptr(self.gpu_replay.sample_indices_d.raw_ptr()); + args.push_i32(b_size_i); + args.push_i32(cap_i); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch - .launch(cfg) - .context("step_with_lobsim: rl_per_sample")?; + raw_launch( + self.rl_per_sample_fn.cu_function(), + (b_size as u32, 1, 1), (1, 1, 1), 0, + self.train_stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_per_sample: {:?}", e))?; } } @@ -6606,25 +6589,23 @@ impl IntegratedTrainer { { let b_size_i = b_size as i32; let cap_i = self.gpu_replay.capacity as i32; - let cfg = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (b_size as u32, 1, 1), - shared_mem_bytes: (b_size * std::mem::size_of::()) as u32, - }; - let mut launch = - self.train_stream.launch_builder(&self.rl_per_update_priority_fn); - launch - .arg(&self.gpu_replay.sample_indices_d) - .arg(&self.td_per_sample_d) - .arg(&mut self.gpu_replay.priority_tree_d) - .arg(&mut self.gpu_replay.max_priority_d) - .arg(&self.isv_d) - .arg(&b_size_i) - .arg(&cap_i); + let smem = (b_size * std::mem::size_of::()) as u32; + let mut args = RawArgs::new(); + args.push_ptr(self.gpu_replay.sample_indices_d.raw_ptr()); + args.push_ptr(self.td_per_sample_d.raw_ptr()); + args.push_ptr(self.gpu_replay.priority_tree_d.raw_ptr()); + args.push_ptr(self.gpu_replay.max_priority_d.raw_ptr()); + args.push_ptr(self.isv_d.raw_ptr()); + args.push_i32(b_size_i); + args.push_i32(cap_i); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch - .launch(cfg) - .context("step_with_lobsim: rl_per_update_priority")?; + raw_launch( + self.rl_per_update_priority_fn.cu_function(), + (1, 1, 1), (b_size as u32, 1, 1), smem, + self.train_stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_per_update_priority: {:?}", e))?; } } @@ -6633,20 +6614,17 @@ impl IntegratedTrainer { // Each internal node written by exactly one thread (no atomics). { let cap_i = self.gpu_replay.capacity as i32; - let cfg = LaunchConfig { - grid_dim: (128, 1, 1), - block_dim: (256, 1, 1), - shared_mem_bytes: 0, - }; - let mut launch = - self.train_stream.launch_builder(&self.rl_per_tree_rebuild_fn); - launch - .arg(&mut self.gpu_replay.priority_tree_d) - .arg(&cap_i); + let mut args = RawArgs::new(); + args.push_ptr(self.gpu_replay.priority_tree_d.raw_ptr()); + args.push_i32(cap_i); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch - .launch(cfg) - .context("step_with_lobsim: rl_per_tree_rebuild")?; + raw_launch( + self.rl_per_tree_rebuild_fn.cu_function(), + (128, 1, 1), (256, 1, 1), 0, + self.train_stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_per_tree_rebuild: {:?}", e))?; } } } @@ -6666,21 +6644,20 @@ impl IntegratedTrainer { { let n_total = self.dqn_head.w_d.len() as i32; let block_dim: u32 = 256; - let cfg = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (block_dim, 1, 1), - shared_mem_bytes: (block_dim as usize * std::mem::size_of::()) as u32, - }; - let mut launch = self.stream.launch_builder(&self.rl_l2_diff_norm_fn); - launch - .arg(&self.dqn_head.w_d) - .arg(&self.dqn_head.w_target_d) - .arg(&mut self.ema_input_scratch_d) - .arg(&n_total); + let smem = (block_dim as usize * std::mem::size_of::()) as u32; + let mut args = RawArgs::new(); + args.push_ptr(self.dqn_head.w_d.raw_ptr()); + args.push_ptr(self.dqn_head.w_target_d.raw_ptr()); + args.push_ptr(self.ema_input_scratch_d.raw_ptr()); + args.push_i32(n_total); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch - .launch(cfg) - .context("rl_l2_diff_norm(W, W_target) launch")?; + raw_launch( + self.rl_l2_diff_norm_fn.cu_function(), + (1, 1, 1), (block_dim, 1, 1), smem, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_l2_diff_norm: {:?}", e))?; } } self.launch_ema_update_per_step( @@ -6712,59 +6689,33 @@ impl IntegratedTrainer { observed_loss_pi: f32, observed_loss_v: f32, ) -> Result<()> { - let cfg = LaunchConfig { - grid_dim: (1, 1, 1), - block_dim: (1, 1, 1), - shared_mem_bytes: 0, - }; - let observed_loss_bce: f32 = 0.0; - let observed_loss_aux: f32 = 0.0; - let q_loss_ema_slot: i32 = - crate::rl::isv_slots::RL_LR_Q_LOSS_EMA_INDEX as i32; - let q_best_slot: i32 = - crate::rl::isv_slots::RL_LR_Q_BEST_LOSS_INDEX as i32; - let q_counter_slot: i32 = - crate::rl::isv_slots::RL_LR_Q_STEPS_SINCE_BEST_INDEX as i32; - let q_warmup_slot: i32 = - crate::rl::isv_slots::RL_LR_Q_WARMUP_COUNTER_INDEX as i32; - let pi_loss_ema_slot: i32 = - crate::rl::isv_slots::RL_LR_PI_LOSS_EMA_INDEX as i32; - let pi_best_slot: i32 = - crate::rl::isv_slots::RL_LR_PI_BEST_LOSS_INDEX as i32; - let pi_counter_slot: i32 = - crate::rl::isv_slots::RL_LR_PI_STEPS_SINCE_BEST_INDEX as i32; - let pi_warmup_slot: i32 = - crate::rl::isv_slots::RL_LR_PI_WARMUP_COUNTER_INDEX as i32; - let v_loss_ema_slot: i32 = - crate::rl::isv_slots::RL_LR_V_LOSS_EMA_INDEX as i32; - let v_best_slot: i32 = - crate::rl::isv_slots::RL_LR_V_BEST_LOSS_INDEX as i32; - let v_counter_slot: i32 = - crate::rl::isv_slots::RL_LR_V_STEPS_SINCE_BEST_INDEX as i32; - let v_warmup_slot: i32 = - crate::rl::isv_slots::RL_LR_V_WARMUP_COUNTER_INDEX as i32; - let mut launch = self.stream.launch_builder(&self.rl_lr_controller_fn); - launch - .arg(&self.isv_d) - .arg(&observed_loss_bce) - .arg(&observed_loss_q) - .arg(&observed_loss_pi) - .arg(&observed_loss_v) - .arg(&observed_loss_aux) - .arg(&q_loss_ema_slot) - .arg(&q_best_slot) - .arg(&q_counter_slot) - .arg(&q_warmup_slot) - .arg(&pi_loss_ema_slot) - .arg(&pi_best_slot) - .arg(&pi_counter_slot) - .arg(&pi_warmup_slot) - .arg(&v_loss_ema_slot) - .arg(&v_best_slot) - .arg(&v_counter_slot) - .arg(&v_warmup_slot); + let mut args = RawArgs::new(); + args.push_ptr(self.isv_d.raw_ptr()); + args.push_f32(0.0_f32); // observed_loss_bce + args.push_f32(observed_loss_q); + args.push_f32(observed_loss_pi); + args.push_f32(observed_loss_v); + args.push_f32(0.0_f32); // observed_loss_aux + args.push_i32(crate::rl::isv_slots::RL_LR_Q_LOSS_EMA_INDEX as i32); + args.push_i32(crate::rl::isv_slots::RL_LR_Q_BEST_LOSS_INDEX as i32); + args.push_i32(crate::rl::isv_slots::RL_LR_Q_STEPS_SINCE_BEST_INDEX as i32); + args.push_i32(crate::rl::isv_slots::RL_LR_Q_WARMUP_COUNTER_INDEX as i32); + args.push_i32(crate::rl::isv_slots::RL_LR_PI_LOSS_EMA_INDEX as i32); + args.push_i32(crate::rl::isv_slots::RL_LR_PI_BEST_LOSS_INDEX as i32); + args.push_i32(crate::rl::isv_slots::RL_LR_PI_STEPS_SINCE_BEST_INDEX as i32); + args.push_i32(crate::rl::isv_slots::RL_LR_PI_WARMUP_COUNTER_INDEX as i32); + args.push_i32(crate::rl::isv_slots::RL_LR_V_LOSS_EMA_INDEX as i32); + args.push_i32(crate::rl::isv_slots::RL_LR_V_BEST_LOSS_INDEX as i32); + args.push_i32(crate::rl::isv_slots::RL_LR_V_STEPS_SINCE_BEST_INDEX as i32); + args.push_i32(crate::rl::isv_slots::RL_LR_V_WARMUP_COUNTER_INDEX as i32); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch.launch(cfg).context("rl_lr_controller launch")?; + raw_launch( + self.rl_lr_controller_fn.cu_function(), + (1, 1, 1), (1, 1, 1), 0, + self.stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("rl_lr_controller: {:?}", e))?; } Ok(()) } @@ -6914,19 +6865,20 @@ fn accumulate_grad_h( let n = grad_h_head.len(); debug_assert_eq!(grad_h_encoder.len(), n); let n_i = n as i32; - let cfg = LaunchConfig { - grid_dim: ((n as u32).div_ceil(256), 1, 1), - block_dim: (256, 1, 1), - shared_mem_bytes: 0, - }; - let mut launch = stream.launch_builder(grad_h_accumulate_fn); - launch - .arg(grad_h_head) - .arg(&lambda) - .arg(&n_i) - .arg(grad_h_encoder); + let grid_x = (n as u32).div_ceil(256); + let mut args = RawArgs::new(); + args.push_ptr(grad_h_head.raw_ptr()); + args.push_f32(lambda); + args.push_i32(n_i); + args.push_ptr(grad_h_encoder.raw_ptr()); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch.launch(cfg).context("grad_h_accumulate launch")?; + raw_launch( + grad_h_accumulate_fn.cu_function(), + (grid_x, 1, 1), (256, 1, 1), 0, + stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("grad_h_accumulate: {:?}", e))?; } Ok(()) } @@ -6948,19 +6900,20 @@ fn reduce_axis0_free( debug_assert_eq!(out.len(), n_tail); let n_batch_i = b_size as i32; let n_tail_i = n_tail as i32; - let cfg = LaunchConfig { - grid_dim: ((n_tail as u32).div_ceil(32), 1, 1), - block_dim: (32, 8, 1), - shared_mem_bytes: 0, - }; - let mut launch = stream.launch_builder(reduce_fn); - launch - .arg(per_batch) - .arg(&n_batch_i) - .arg(&n_tail_i) - .arg(out); + let grid_x = (n_tail as u32).div_ceil(32); + let mut args = RawArgs::new(); + args.push_ptr(per_batch.raw_ptr()); + args.push_i32(n_batch_i); + args.push_i32(n_tail_i); + args.push_ptr(out.raw_ptr()); + let mut ptrs = args.build_arg_ptrs(); unsafe { - launch.launch(cfg).context("reduce_axis0 launch")?; + raw_launch( + reduce_fn.cu_function(), + (grid_x, 1, 1), (32, 8, 1), 0, + stream.cu_stream(), + &mut ptrs[..args.len()], + ).map_err(|e| anyhow::anyhow!("reduce_axis0: {:?}", e))?; } Ok(()) }