diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index a04f81885..172472aad 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -1504,11 +1504,12 @@ impl GpuDqnTrainer { let na = self.config.num_atoms; let param_sizes = compute_param_sizes(&self.config); let shmem = (na * 2) * std::mem::size_of::(); + let atom_positions_buf_ptr = self.atom_positions_buf.raw_ptr(); for branch in 0..4_usize { let spacing_ptr = self.ptrs.params_ptr + padded_byte_offset(¶m_sizes, 52 + branch); let out_offset = (branch * na * std::mem::size_of::()) as u64; - let out_ptr = self.atom_positions_buf.raw_ptr() + out_offset; + let out_ptr = atom_positions_buf_ptr + out_offset; unsafe { self.stream.launch_builder(&self.adaptive_atom_kernel) @@ -1616,17 +1617,19 @@ impl GpuDqnTrainer { let budget_max = lambda_base * 50.0; // budget scales with lambda // Zero total_penalty before reduction + let homeostatic_total_buf_ptr = self.homeostatic_total_buf.raw_ptr(); + let homeostatic_penalties_buf_ptr = self.homeostatic_penalties_buf.raw_ptr(); unsafe { cudarc::driver::sys::cuMemsetD8Async( - self.homeostatic_total_buf.raw_ptr(), 0, std::mem::size_of::(), self.stream.cu_stream(), + homeostatic_total_buf_ptr, 0, std::mem::size_of::(), self.stream.cu_stream(), ); } unsafe { self.stream.launch_builder(&self.homeostatic_kernel) .arg(&self.homeostatic_obs_dev_ptr) .arg(&self.homeostatic_targets_dev_ptr) - .arg(&self.homeostatic_penalties_buf) - .arg(&self.homeostatic_total_buf) + .arg(&homeostatic_penalties_buf_ptr) + .arg(&homeostatic_total_buf_ptr) .arg(&(HOMEOSTATIC_N_OBS as i32)) .arg(&lambda_base) .arg(&budget_max) @@ -2113,6 +2116,9 @@ impl GpuDqnTrainer { let w_gamma = self.ptrs.params_ptr + padded_byte_offset(¶m_sizes, 74); let b_gamma = self.ptrs.params_ptr + padded_byte_offset(¶m_sizes, 75); let k = ISV_K as i32; + let isv_embedding_buf_ptr = self.isv_embedding_buf.raw_ptr(); + let branch_gate_buf_ptr = self.branch_gate_buf.raw_ptr(); + let gamma_mod_buf_ptr = self.gamma_mod_buf.raw_ptr(); unsafe { self.stream.launch_builder(&self.isv_forward_kernel) @@ -2123,9 +2129,9 @@ impl GpuDqnTrainer { .arg(&w_fc2).arg(&b_fc2) .arg(&w_gate).arg(&b_gate) .arg(&w_gamma).arg(&b_gamma) - .arg(&self.isv_embedding_buf) - .arg(&self.branch_gate_buf) - .arg(&self.gamma_mod_buf) + .arg(&isv_embedding_buf_ptr) + .arg(&branch_gate_buf_ptr) + .arg(&gamma_mod_buf_ptr) .arg(&k) .launch(LaunchConfig { grid_dim: (1, 1, 1), @@ -2147,11 +2153,13 @@ impl GpuDqnTrainer { let blocks = ((batch_size as u32 + 255) / 256).max(1); let b_i32 = batch_size as i32; let sh2 = self.config.shared_h2 as i32; + let save_h_s2_ptr = self.save_h_s2.raw_ptr(); + let isv_embedding_buf_ptr = self.isv_embedding_buf.raw_ptr(); unsafe { self.stream.launch_builder(&self.isv_feature_gate_kernel) - .arg(&self.save_h_s2) // h_s2 in-place (WAW dep needed for Hopper graph replay) - .arg(&self.isv_embedding_buf) + .arg(&save_h_s2_ptr) // h_s2 in-place (WAW dep needed for Hopper graph replay) + .arg(&isv_embedding_buf_ptr) .arg(&w_gate) .arg(&b_gate) .arg(&b_i32) @@ -2174,13 +2182,15 @@ impl GpuDqnTrainer { let w_route = self.ptrs.params_ptr + padded_byte_offset(¶m_sizes, 80); let b_route = self.ptrs.params_ptr + padded_byte_offset(¶m_sizes, 81); let sh2 = self.config.shared_h2 as i32; + let isv_embedding_buf_ptr = self.isv_embedding_buf.raw_ptr(); + let temporal_weight_buf_ptr = self.temporal_weight_buf.raw_ptr(); unsafe { self.stream.launch_builder(&self.isv_temporal_route_kernel) - .arg(&self.isv_embedding_buf) + .arg(&isv_embedding_buf_ptr) .arg(&w_route) .arg(&b_route) - .arg(&self.temporal_weight_buf) + .arg(&temporal_weight_buf_ptr) .arg(&sh2) .launch(LaunchConfig { grid_dim: (1, 1, 1), @@ -2201,13 +2211,15 @@ impl GpuDqnTrainer { let blocks = ((batch_size as u32 + 255) / 256).max(1); let b_i32 = batch_size as i32; let sh2 = self.config.shared_h2 as i32; + let save_h_s2_ptr = self.save_h_s2.raw_ptr(); + let predicted_error_buf_ptr = self.predicted_error_buf.raw_ptr(); unsafe { self.stream.launch_builder(&self.recursive_conf_fwd_kernel) - .arg(&self.save_h_s2) + .arg(&save_h_s2_ptr) .arg(&w_conf) .arg(&b_conf) - .arg(&self.predicted_error_buf) + .arg(&predicted_error_buf_ptr) .arg(&b_i32) .arg(&sh2) .launch(LaunchConfig { @@ -2233,13 +2245,15 @@ impl GpuDqnTrainer { let b_i32 = batch_size as i32; let sh2 = self.config.shared_h2 as i32; let ah = self.config.adv_h as i32; + let save_h_s2_ptr = self.save_h_s2.raw_ptr(); + let plan_params_buf_ptr = self.plan_params_buf.raw_ptr(); unsafe { self.stream.launch_builder(&self.trade_plan_fwd_kernel) - .arg(&self.save_h_s2) + .arg(&save_h_s2_ptr) .arg(&w_fc).arg(&b_fc) .arg(&w_out).arg(&b_out) - .arg(&self.plan_params_buf) + .arg(&plan_params_buf_ptr) .arg(&b_i32).arg(&sh2).arg(&ah) .launch(LaunchConfig { grid_dim: (blocks, 1, 1), @@ -2259,10 +2273,11 @@ impl GpuDqnTrainer { let b_i32 = batch_size as i32; let step = self.adam_step; let noise_scale = 0.05_f32; // ±5% + let plan_params_buf_ptr = self.plan_params_buf.raw_ptr(); unsafe { self.stream.launch_builder(&self.plan_noise_kernel) - .arg(&self.plan_params_buf) + .arg(&plan_params_buf_ptr) .arg(&b_i32) .arg(&step) .arg(&noise_scale) @@ -2289,16 +2304,20 @@ impl GpuDqnTrainer { let b_i32 = batch_size as i32; let sh2 = self.config.shared_h2 as i32; let loss_weight = 0.01_f32; + let save_h_s2_ptr = self.save_h_s2.raw_ptr(); + let predicted_error_buf_ptr = self.predicted_error_buf.raw_ptr(); + let recursive_conf_partials_ptr = self.recursive_conf_partials.raw_ptr(); + let bw_d_h_s2_ptr = self.bw_d_h_s2.raw_ptr(); unsafe { // Stage 1: per-block partial reduction (no atomicAdd) self.stream.launch_builder(&self.recursive_conf_bwd_kernel) - .arg(&self.save_h_s2) - .arg(&self.predicted_error_buf) + .arg(&save_h_s2_ptr) + .arg(&predicted_error_buf_ptr) .arg(&self.lagged_td_error_dev_ptr) .arg(&w_conf) - .arg(&self.recursive_conf_partials) - .arg(&self.bw_d_h_s2) + .arg(&recursive_conf_partials_ptr) + .arg(&bw_d_h_s2_ptr) .arg(&b_i32) .arg(&sh2) .arg(&loss_weight) @@ -2314,7 +2333,7 @@ impl GpuDqnTrainer { let total_params = sh2 + 1; let reduce_blocks = ((total_params as u32 + 255) / 256).max(1); self.stream.launch_builder(&self.recursive_conf_reduce_kernel) - .arg(&self.recursive_conf_partials) + .arg(&recursive_conf_partials_ptr) .arg(&d_w_conf) .arg(&d_b_conf) .arg(&num_blocks_i32) @@ -2335,9 +2354,10 @@ impl GpuDqnTrainer { let blocks = ((self.config.batch_size as u32 + 255) / 256).max(1); let b = self.config.batch_size as i32; let gamma_mod_ptr = self.gamma_mod_buf.raw_ptr(); + let gamma_buf_ptr = self.gamma_buf.raw_ptr(); unsafe { self.stream.launch_builder(&self.fill_gamma_buf_kernel) - .arg(&self.gamma_buf) + .arg(&gamma_buf_ptr) .arg(&base_gamma) .arg(&gamma_mod_ptr) .arg(&b) @@ -2428,17 +2448,21 @@ impl GpuDqnTrainer { let w_out_ptr = self.ptrs.params_ptr + padded_byte_offset(¶m_sizes, 66); let b_out_ptr = self.ptrs.params_ptr + padded_byte_offset(¶m_sizes, 67); + let save_h_s2_ptr = self.save_h_s2.raw_ptr(); + let predicted_error_buf_ptr = self.predicted_error_buf.raw_ptr(); + let risk_hidden_buf_ptr = self.risk_hidden_buf.raw_ptr(); + let risk_budget_buf_ptr = self.risk_budget_buf.raw_ptr(); unsafe { self.stream.launch_builder(&self.risk_forward_kernel) - .arg(&self.save_h_s2) + .arg(&save_h_s2_ptr) .arg(&self.isv_signals_dev_ptr) // ISV raw [8] - .arg(&self.predicted_error_buf) // predicted_error [B] + .arg(&predicted_error_buf_ptr) // predicted_error [B] .arg(&w_fc_ptr) .arg(&b_fc_ptr) .arg(&w_out_ptr) .arg(&b_out_ptr) - .arg(&self.risk_hidden_buf) - .arg(&self.risk_budget_buf) + .arg(&risk_hidden_buf_ptr) + .arg(&risk_budget_buf_ptr) .arg(&(batch_size as i32)) .arg(&(self.config.shared_h2 as i32)) .arg(&(self.config.adv_h as i32)) @@ -2451,12 +2475,16 @@ impl GpuDqnTrainer { /// Apply risk budget: scales magnitude Q-values (Full×R, Half×sqrt(R)), /// produces per-sample CVaR alpha and commitment lambda. pub(crate) fn apply_risk_budget(&self, batch_size: usize) -> Result<(), MLError> { + let risk_budget_buf_ptr = self.risk_budget_buf.raw_ptr(); + let q_out_buf_ptr = self.q_out_buf.raw_ptr(); + let cvar_alpha_buf_ptr = self.cvar_alpha_buf.raw_ptr(); + let commit_lambda_buf_ptr = self.commit_lambda_buf.raw_ptr(); unsafe { self.stream.launch_builder(&self.risk_apply_kernel) - .arg(&self.risk_budget_buf) - .arg(&self.q_out_buf) - .arg(&self.cvar_alpha_buf) - .arg(&self.commit_lambda_buf) + .arg(&risk_budget_buf_ptr) + .arg(&q_out_buf_ptr) + .arg(&cvar_alpha_buf_ptr) + .arg(&commit_lambda_buf_ptr) .arg(&(batch_size as i32)) .arg(&(self.config.branch_0_size as i32)) .arg(&(self.config.branch_1_size as i32)) @@ -2530,9 +2558,10 @@ impl GpuDqnTrainer { pub(crate) fn apply_regime_dropout(&self, batch_size: usize, is_training: bool) -> Result<(), MLError> { let sh2 = self.config.shared_h2; let total = batch_size * sh2; + let save_h_s2_ptr = self.save_h_s2.raw_ptr(); unsafe { self.stream.launch_builder(&self.regime_dropout_kernel) - .arg(&self.save_h_s2) + .arg(&save_h_s2_ptr) .arg(&self.ptrs.states_buf) .arg(&(batch_size as i32)) .arg(&(sh2 as i32)) @@ -2557,10 +2586,12 @@ impl GpuDqnTrainer { /// Must be called AFTER apply_branch_confidence_routing and BEFORE launch_q_attention. pub(crate) fn apply_epistemic_gate(&self, batch_size: usize) -> Result<(), MLError> { let ta = self.total_actions() as i32; + let q_out_buf_ptr = self.q_out_buf.raw_ptr(); + let q_var_buf_trainer_ptr = self.q_var_buf_trainer.raw_ptr(); unsafe { self.stream.launch_builder(&self.epistemic_gate_kernel) - .arg(&self.q_out_buf) - .arg(&self.q_var_buf_trainer) + .arg(&q_out_buf_ptr) + .arg(&q_var_buf_trainer_ptr) .arg(&self.var_ema_dev_ptr) .arg(&(batch_size as i32)) .arg(&ta) @@ -2581,10 +2612,15 @@ impl GpuDqnTrainer { /// Informational — writes scalar to branch_indep_penalty_buf for logging. /// Gradient integration deferred (requires backward pass modification). pub(crate) fn compute_branch_independence(&self, batch_size: usize) -> Result<(), MLError> { + let branch_indep_penalty_buf_ptr = self.branch_indep_penalty_buf.raw_ptr(); + let save_h_b0_ptr = self.save_h_b0.raw_ptr(); + let save_h_b1_ptr = self.save_h_b1.raw_ptr(); + let save_h_b2_ptr = self.save_h_b2.raw_ptr(); + let save_h_b3_ptr = self.save_h_b3.raw_ptr(); // Use raw memset to avoid &mut self borrow conflict on the penalty buf unsafe { cudarc::driver::sys::cuMemsetD8Async( - self.branch_indep_penalty_buf.raw_ptr(), + branch_indep_penalty_buf_ptr, 0, std::mem::size_of::(), self.stream.cu_stream(), @@ -2593,11 +2629,11 @@ impl GpuDqnTrainer { let ah = self.config.adv_h as i32; unsafe { self.stream.launch_builder(&self.branch_indep_kernel) - .arg(&self.save_h_b0) - .arg(&self.save_h_b1) - .arg(&self.save_h_b2) - .arg(&self.save_h_b3) - .arg(&self.branch_indep_penalty_buf) + .arg(&save_h_b0_ptr) + .arg(&save_h_b1_ptr) + .arg(&save_h_b2_ptr) + .arg(&save_h_b3_ptr) + .arg(&branch_indep_penalty_buf_ptr) .arg(&(batch_size as i32)) .arg(&ah) .arg(&0.01_f32) // lambda_indep @@ -2614,10 +2650,13 @@ impl GpuDqnTrainer { /// G10: Compute temporal consistency penalty (Lipschitz on Q-diffs for similar states). /// Informational — writes scalar to temporal_penalty_buf for logging. pub(crate) fn compute_temporal_consistency(&self, batch_size: usize) -> Result<(), MLError> { + let temporal_penalty_buf_ptr = self.temporal_penalty_buf.raw_ptr(); + let save_h_s2_ptr = self.save_h_s2.raw_ptr(); + let q_out_buf_ptr = self.q_out_buf.raw_ptr(); // Use raw memset to avoid &mut self borrow conflict on the penalty buf unsafe { cudarc::driver::sys::cuMemsetD8Async( - self.temporal_penalty_buf.raw_ptr(), + temporal_penalty_buf_ptr, 0, std::mem::size_of::(), self.stream.cu_stream(), @@ -2627,9 +2666,9 @@ impl GpuDqnTrainer { let sh2 = self.config.shared_h2 as i32; unsafe { self.stream.launch_builder(&self.temporal_consistency_kernel) - .arg(&self.save_h_s2) - .arg(&self.q_out_buf) - .arg(&self.temporal_penalty_buf) + .arg(&save_h_s2_ptr) + .arg(&q_out_buf_ptr) + .arg(&temporal_penalty_buf_ptr) .arg(&(batch_size as i32)) .arg(&sh2) .arg(&ta) @@ -2649,11 +2688,13 @@ impl GpuDqnTrainer { /// reduces via c51_loss_reduce kernel. No cuMemsetD8Async or atomicAdd. pub(crate) fn compute_predictive_coding_loss(&self, batch_size: usize) -> Result<(), MLError> { let sh2 = self.config.shared_h2 as i32; + let save_h_s2_ptr = self.save_h_s2.raw_ptr(); + let predictive_per_sample_buf_ptr = self.predictive_per_sample_buf.raw_ptr(); // Step 1: per-sample MSE loss (graph-captured kernel, no memset needed) unsafe { self.stream.launch_builder(&self.predictive_coding_kernel) - .arg(&self.save_h_s2) - .arg(&self.predictive_per_sample_buf) + .arg(&save_h_s2_ptr) + .arg(&predictive_per_sample_buf_ptr) .arg(&(batch_size as i32)) .arg(&sh2) .arg(&0.1_f32) // lambda_pred @@ -2665,7 +2706,7 @@ impl GpuDqnTrainer { unsafe { self.stream .launch_builder(&self.c51_loss_reduce_kernel) - .arg(&self.predictive_per_sample_buf) + .arg(&predictive_per_sample_buf_ptr) .arg(&loss_ptr) .arg(&(batch_size as i32)) .launch(LaunchConfig { @@ -2838,12 +2879,16 @@ impl GpuDqnTrainer { let w_a_ptr = param_ptr; let w_b_ptr = param_ptr + (sh2 * MAMBA2_STATE_DIM * 4) as u64; let w_c_ptr = param_ptr + (2 * sh2 * MAMBA2_STATE_DIM * 4) as u64; + let mamba2_h_history_ptr = self.mamba2_h_history.raw_ptr(); + let save_h_s2_ptr = self.save_h_s2.raw_ptr(); + let mamba2_h_enriched_ptr = self.mamba2_h_enriched.raw_ptr(); + let temporal_weight_buf_ptr = self.temporal_weight_buf.raw_ptr(); unsafe { self.stream.launch_builder(&self.mamba2_scan_kernel) - .arg(&self.mamba2_h_history) - .arg(&self.save_h_s2) - .arg(&self.mamba2_h_enriched) + .arg(&mamba2_h_history_ptr) + .arg(&save_h_s2_ptr) + .arg(&mamba2_h_enriched_ptr) .arg(&w_a_ptr) .arg(&w_b_ptr) .arg(&w_c_ptr) @@ -2852,7 +2897,7 @@ impl GpuDqnTrainer { .arg(&(sh2 as i32)) .arg(&(MAMBA2_STATE_DIM as i32)) .arg(&self.isv_signals_dev_ptr) // regime-conditioned decay (ISV[11]) - .arg(&self.temporal_weight_buf) // per-feature temporal routing + .arg(&temporal_weight_buf_ptr) // per-feature temporal routing .launch(LaunchConfig::for_num_elems(batch_size as u32)) .map_err(|e| MLError::ModelError(format!("mamba2_temporal_scan: {e}")))?; } @@ -2863,11 +2908,13 @@ impl GpuDqnTrainer { pub(crate) fn mamba2_update_history(&self, batch_size: usize) -> Result<(), MLError> { let sh2 = self.config.shared_h2; let total = batch_size * sh2; + let mamba2_h_history_ptr = self.mamba2_h_history.raw_ptr(); + let save_h_s2_ptr = self.save_h_s2.raw_ptr(); unsafe { self.stream.launch_builder(&self.mamba2_update_kernel) - .arg(&self.mamba2_h_history) - .arg(&self.save_h_s2) + .arg(&mamba2_h_history_ptr) + .arg(&save_h_s2_ptr) .arg(&(batch_size as i32)) .arg(&(MAMBA2_HISTORY_K as i32)) .arg(&(sh2 as i32)) @@ -2884,10 +2931,12 @@ impl GpuDqnTrainer { // Graph-safe copy: enriched -> save_h_s2 so branch heads use enriched activations. // Uses a kernel instead of raw memcpy_dtod_async which is NOT captured in CUDA Graph. let n = (batch_size * self.config.shared_h2) as u32; + let mamba2_h_enriched_ptr = self.mamba2_h_enriched.raw_ptr(); + let save_h_s2_ptr = self.save_h_s2.raw_ptr(); unsafe { self.stream.launch_builder(&self.mamba2_copy_enriched_kernel) - .arg(&self.mamba2_h_enriched) - .arg(&self.save_h_s2) + .arg(&mamba2_h_enriched_ptr) + .arg(&save_h_s2_ptr) .arg(&(n as i32)) .launch(LaunchConfig::for_num_elems(n)) .map_err(|e| MLError::ModelError(format!("mamba2 enriched copy kernel: {e}")))?; @@ -2913,12 +2962,15 @@ impl GpuDqnTrainer { // Deterministic kernel: each thread writes exactly one element (plain write, // not atomicAdd), so no memset_zeros needed. let total_weight_params = (3 * sh2 * MAMBA2_STATE_DIM) as u32; + let mamba2_h_history_ptr = self.mamba2_h_history.raw_ptr(); + let save_h_s2_ptr = self.save_h_s2.raw_ptr(); + let bw_d_h_s2_ptr = self.bw_d_h_s2.raw_ptr(); unsafe { self.stream.launch_builder(&self.mamba2_backward_kernel) - .arg(&self.mamba2_h_history) - .arg(&self.save_h_s2) - .arg(&self.bw_d_h_s2) // d_h_enriched = trunk activation gradient [B, SH2] + .arg(&mamba2_h_history_ptr) + .arg(&save_h_s2_ptr) + .arg(&bw_d_h_s2_ptr) // d_h_enriched = trunk activation gradient [B, SH2] .arg(&w_a_ptr) .arg(&w_b_ptr) .arg(&w_c_ptr) @@ -2945,10 +2997,12 @@ impl GpuDqnTrainer { let lr = base_lr * self.new_component_lr_scale(); // SGD: params += (-lr) * grad + let mamba2_params_ptr = self.mamba2_params.raw_ptr(); + let mamba2_grad_ptr = self.mamba2_grad.raw_ptr(); unsafe { self.stream.launch_builder(&self.saxpy_f32_kernel) - .arg(&self.mamba2_params) - .arg(&self.mamba2_grad) + .arg(&mamba2_params_ptr) + .arg(&mamba2_grad_ptr) .arg(&(-lr)) .arg(&(n as i32)) .launch(LaunchConfig::for_num_elems(n as u32)) @@ -3840,10 +3894,11 @@ impl GpuDqnTrainer { // Single batched launch: 12 blocks × 256 threads, one matrix per block. // Descriptor buffer was built at construction with stable pointers. + let spectral_norm_descriptors_ptr = self.spectral_norm_descriptors.raw_ptr(); unsafe { self.stream .launch_builder(&self.spectral_norm_batched_kernel) - .arg(&self.spectral_norm_descriptors) + .arg(&spectral_norm_descriptors_ptr) .arg(&sigma_max) .arg(&num_matrices) .launch(LaunchConfig { @@ -6648,11 +6703,13 @@ impl GpuDqnTrainer { // causal_mean_reduce kernel fully overwrites out[0] via direct assignment // (single-block, thread 0 only, no atomicAdd) — no memset needed. let n_features_i32 = market_dim.min(14) as i32; + let causal_sensitivity_buf_ptr = self.causal_sensitivity_buf.raw_ptr(); + let causal_mean_scratch_ptr = self.causal_mean_scratch.raw_ptr(); unsafe { self.stream .launch_builder(&self.causal_mean_reduce_kernel) - .arg(&self.causal_sensitivity_buf) - .arg(&self.causal_mean_scratch) + .arg(&causal_sensitivity_buf_ptr) + .arg(&causal_mean_scratch_ptr) .arg(&n_features_i32) .launch(LaunchConfig { grid_dim: (1, 1, 1), @@ -6666,7 +6723,7 @@ impl GpuDqnTrainer { unsafe { cudarc::driver::sys::cuMemcpyDtoHAsync_v2( self.readback_pinned.add(10).cast(), - self.causal_mean_scratch.raw_ptr(), + causal_mean_scratch_ptr, std::mem::size_of::(), self.stream.cu_stream(), ); @@ -6764,11 +6821,13 @@ impl GpuDqnTrainer { // GPU-side mean reduction + async DtoH to pinned buffer let n_features_i32 = market_dim.min(14) as i32; + let causal_sensitivity_buf_ptr = self.causal_sensitivity_buf.raw_ptr(); + let causal_mean_scratch_ptr = self.causal_mean_scratch.raw_ptr(); unsafe { self.stream .launch_builder(&self.causal_mean_reduce_kernel) - .arg(&self.causal_sensitivity_buf) - .arg(&self.causal_mean_scratch) + .arg(&causal_sensitivity_buf_ptr) + .arg(&causal_mean_scratch_ptr) .arg(&n_features_i32) .launch(LaunchConfig { grid_dim: (1, 1, 1), @@ -6782,7 +6841,7 @@ impl GpuDqnTrainer { unsafe { cudarc::driver::sys::cuMemcpyDtoHAsync_v2( self.readback_pinned.add(10).cast(), - self.causal_mean_scratch.raw_ptr(), + causal_mean_scratch_ptr, std::mem::size_of::(), self.stream.cu_stream(), ); @@ -7227,6 +7286,7 @@ impl GpuDqnTrainer { let q_out_ptr = self.q_out_buf.raw_ptr(); let null_atom_stats = 0u64; let null_q_var = 0u64; + let atom_positions_buf_ptr = self.atom_positions_buf.raw_ptr(); unsafe { self.stream .launch_builder(&self.expected_q_kernel) @@ -7242,7 +7302,7 @@ impl GpuDqnTrainer { .arg(&self.per_sample_support_ptr) .arg(&null_atom_stats) .arg(&null_q_var) - .arg(&self.atom_positions_buf) + .arg(&atom_positions_buf_ptr) .launch(LaunchConfig { grid_dim: (grid_dim, 1, 1), block_dim: (block_dim, 1, 1), @@ -7348,6 +7408,7 @@ impl GpuDqnTrainer { let q_out_ptr = self.q_out_buf.raw_ptr(); let null_atom_stats = 0u64; let null_q_var = 0u64; + let atom_positions_buf_ptr = self.atom_positions_buf.raw_ptr(); unsafe { self.stream .launch_builder(&self.expected_q_kernel) @@ -7363,7 +7424,7 @@ impl GpuDqnTrainer { .arg(&self.per_sample_support_ptr) .arg(&null_atom_stats) .arg(&null_q_var) - .arg(&self.atom_positions_buf) + .arg(&atom_positions_buf_ptr) .launch(LaunchConfig { grid_dim: (grid_dim, 1, 1), block_dim: (block_dim, 1, 1), @@ -7405,13 +7466,16 @@ impl GpuDqnTrainer { let total_actions = self.total_actions() as i32; let n = batch_size as i32; let num_atoms = self.config.num_atoms as i32; + let q_out_buf_ptr = self.q_out_buf.raw_ptr(); + let atom_stats_buf_ptr = self.atom_stats_buf.raw_ptr(); + let q_stats_buf_ptr = self.q_stats_buf.raw_ptr(); unsafe { self.stream .launch_builder(&self.q_stats_kernel) - .arg(&self.q_out_buf) - .arg(&self.atom_stats_buf) - .arg(&self.q_stats_buf) + .arg(&q_out_buf_ptr) + .arg(&atom_stats_buf_ptr) + .arg(&q_stats_buf_ptr) .arg(&n) .arg(&total_actions) .arg(&num_atoms) @@ -7461,11 +7525,13 @@ impl GpuDqnTrainer { // q_stats_kernel writes 7 floats to q_readback_dev_ptr[0..7]. // CPU reads from q_readback_pinned — previous step's values. let stats_dev_ptr = self.q_readback_dev_ptr; + let q_out_buf_ptr = self.q_out_buf.raw_ptr(); + let atom_stats_buf_ptr = self.atom_stats_buf.raw_ptr(); unsafe { self.stream .launch_builder(&self.q_stats_kernel) - .arg(&self.q_out_buf) - .arg(&self.atom_stats_buf) + .arg(&q_out_buf_ptr) + .arg(&atom_stats_buf_ptr) .arg(&stats_dev_ptr) .arg(&n) .arg(&total_actions) @@ -7676,6 +7742,7 @@ impl GpuDqnTrainer { let q_out_ptr = self.denoise_target_q_buf.raw_ptr(); let null_atom_stats = 0u64; let null_q_var = 0u64; + let atom_positions_buf_ptr = self.atom_positions_buf.raw_ptr(); unsafe { self.stream .launch_builder(&self.expected_q_kernel) @@ -7691,7 +7758,7 @@ impl GpuDqnTrainer { .arg(&self.per_sample_support_ptr) .arg(&null_atom_stats) .arg(&null_q_var) - .arg(&self.atom_positions_buf) + .arg(&atom_positions_buf_ptr) .launch(LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), @@ -8535,6 +8602,10 @@ impl GpuDqnTrainer { let cur_input = self.config.market_dim + 3; let cur_hidden: usize = 128; let cur_output = self.config.market_dim; + let curiosity_input_buf_ptr = self.curiosity_input_buf.raw_ptr(); + let curiosity_hidden_buf_ptr = self.curiosity_hidden_buf.raw_ptr(); + let curiosity_pred_buf_ptr = self.curiosity_pred_buf.raw_ptr(); + let curiosity_error_buf_ptr = self.curiosity_error_buf.raw_ptr(); // Step 1: Prepare input — build [B, CUR_INPUT] from states + action one-hot let blocks_b = ((b + 255) / 256) as u32; @@ -8543,7 +8614,7 @@ impl GpuDqnTrainer { .launch_builder(&self.curiosity_prepare_input_func) .arg(&self.ptrs.states_buf) .arg(&self.ptrs.actions_buf) - .arg(&self.curiosity_input_buf) + .arg(&curiosity_input_buf_ptr) .arg(&n) .arg(&sd) .launch(cudarc::driver::LaunchConfig { @@ -8558,8 +8629,8 @@ impl GpuDqnTrainer { self.cublas_forward.sgemm_f32( &self.stream, self.curiosity_w1_ptr, - self.curiosity_input_buf.raw_ptr(), - self.curiosity_hidden_buf.raw_ptr(), + curiosity_input_buf_ptr, + curiosity_hidden_buf_ptr, cur_hidden, b, cur_input, "cur_gemm1", )?; @@ -8571,7 +8642,7 @@ impl GpuDqnTrainer { unsafe { self.stream .launch_builder(&self.curiosity_bias_leaky_relu_func) - .arg(&self.curiosity_hidden_buf) + .arg(&curiosity_hidden_buf_ptr) .arg(&self.curiosity_b1_ptr) .arg(&cur_hidden_i32) .arg(&total_hidden) @@ -8587,8 +8658,8 @@ impl GpuDqnTrainer { self.cublas_forward.sgemm_f32( &self.stream, self.curiosity_w2_ptr, - self.curiosity_hidden_buf.raw_ptr(), - self.curiosity_pred_buf.raw_ptr(), + curiosity_hidden_buf_ptr, + curiosity_pred_buf_ptr, cur_output, b, cur_hidden, "cur_gemm2", )?; @@ -8597,10 +8668,10 @@ impl GpuDqnTrainer { unsafe { self.stream .launch_builder(&self.curiosity_bias_mse_func) - .arg(&self.curiosity_pred_buf) + .arg(&curiosity_pred_buf_ptr) .arg(&self.curiosity_b2_ptr) .arg(&self.ptrs.next_states_buf) - .arg(&self.curiosity_error_buf) + .arg(&curiosity_error_buf_ptr) .arg(&n) .arg(&sd) .launch(cudarc::driver::LaunchConfig { @@ -8628,6 +8699,24 @@ impl GpuDqnTrainer { let b2 = self.config.branch_2_size; let b3 = self.config.branch_3_size; + // Extract raw u64 pointers for CudaSlice fields (eliminates cudarc overhead) + let on_v_logits_buf_ptr = self.on_v_logits_buf.raw_ptr(); + let tg_v_logits_buf_ptr = self.tg_v_logits_buf.raw_ptr(); + let on_next_v_logits_buf_ptr = self.on_next_v_logits_buf.raw_ptr(); + let actions_buf_ptr = self.actions_buf.raw_ptr(); + let rewards_buf_ptr = self.rewards_buf.raw_ptr(); + let dones_buf_ptr = self.dones_buf.raw_ptr(); + let is_weights_buf_ptr = self.is_weights_buf.raw_ptr(); + let per_sample_loss_buf_ptr = self.per_sample_loss_buf.raw_ptr(); + let td_errors_buf_ptr = self.td_errors_buf.raw_ptr(); + let save_current_lp_ptr = self.save_current_lp.raw_ptr(); + let save_projected_ptr = self.save_projected.raw_ptr(); + let curiosity_error_buf_ptr = self.curiosity_error_buf.raw_ptr(); + let drawdown_depths_buf_ptr = self.drawdown_depths_buf.raw_ptr(); + let ensemble_std_buf_ptr = self.ensemble_std_buf.raw_ptr(); + let atom_positions_buf_ptr = self.atom_positions_buf.raw_ptr(); + let cvar_alpha_buf_ptr = self.cvar_alpha_buf.raw_ptr(); + // Extract raw pointers for online branch logits (contiguous [B, (B0+B1+B2+B3)*NA], f32) let f32_sz = std::mem::size_of::(); let on_b_base = self.on_b_logits_buf.raw_ptr(); @@ -8670,37 +8759,37 @@ impl GpuDqnTrainer { self.stream .launch_builder(kernel) // ── Online on current states (5, f32 logits) ── - .arg(&self.on_v_logits_buf) + .arg(&on_v_logits_buf_ptr) .arg(&on_b0_ptr) .arg(&on_b1_ptr) .arg(&on_b2_ptr) .arg(&on_b3_ptr) // ── Target on next_states (5, f32 logits) ── - .arg(&self.tg_v_logits_buf) + .arg(&tg_v_logits_buf_ptr) .arg(&tg_b0_ptr) .arg(&tg_b1_ptr) .arg(&tg_b2_ptr) .arg(&tg_b3_ptr) // ── Online-next on next_states (Double DQN, 5, f32 logits) ── - .arg(&self.on_next_v_logits_buf) + .arg(&on_next_v_logits_buf_ptr) .arg(&on_next_b0_ptr) .arg(&on_next_b1_ptr) .arg(&on_next_b2_ptr) .arg(&on_next_b3_ptr) // ── Batch data (4) ── - .arg(&self.actions_buf) - .arg(&self.rewards_buf) - .arg(&self.dones_buf) - .arg(&self.is_weights_buf) + .arg(&actions_buf_ptr) + .arg(&rewards_buf_ptr) + .arg(&dones_buf_ptr) + .arg(&is_weights_buf_ptr) // ── Outputs (3) ── - .arg(&self.per_sample_loss_buf) - .arg(&self.td_errors_buf) + .arg(&per_sample_loss_buf_ptr) + .arg(&td_errors_buf_ptr) .arg(&self.total_loss_dev_ptr) // ── Saved for backward (2) ── - .arg(&self.save_current_lp) - .arg(&self.save_projected) + .arg(&save_current_lp_ptr) + .arg(&save_projected_ptr) // ── Curiosity Q-penalty (2) ── - .arg(&self.curiosity_error_buf) + .arg(&curiosity_error_buf_ptr) .arg(&self.config.curiosity_q_penalty_lambda) // ── Config (8 — per_sample_support replaces v_min+v_max) ── .arg(&gamma_buf_ptr) @@ -8712,10 +8801,10 @@ impl GpuDqnTrainer { .arg(&b2_i32) .arg(&b3_i32) // ── #18 Asymmetric DD loss (2) ── - .arg(&self.drawdown_depths_buf) + .arg(&drawdown_depths_buf_ptr) .arg(&self.asymmetric_dd_weight) // ── #27 Ensemble disagreement (2) ── - .arg(&self.ensemble_std_buf) + .arg(&ensemble_std_buf_ptr) .arg(&self.ensemble_disagreement_weight) // ── Spectral decoupling (1) ── .arg(&self.config.spectral_decoupling_lambda) @@ -8726,9 +8815,9 @@ impl GpuDqnTrainer { // ── Adam step counter for stochastic Expected SARSA ── .arg(&self.ptrs.t_buf) // ── Adaptive atom positions ── - .arg(&self.atom_positions_buf) + .arg(&atom_positions_buf_ptr) // ── Per-sample CVaR alpha from learned risk branch ── - .arg(&self.cvar_alpha_buf) + .arg(&cvar_alpha_buf_ptr) .launch(LaunchConfig { grid_dim: (b as u32, 1, 1), block_dim: (256, 1, 1), @@ -8744,10 +8833,11 @@ impl GpuDqnTrainer { /// Grid=(1,1,1), Block=(1,1,1). Must stay single-thread inside graph_mega on Hopper. fn launch_loss_reduce(&self, total_loss_dev_ptr: u64) -> Result<(), MLError> { let b = self.config.batch_size as i32; + let per_sample_loss_buf_ptr = self.per_sample_loss_buf.raw_ptr(); unsafe { self.stream .launch_builder(&self.c51_loss_reduce_kernel) - .arg(&self.per_sample_loss_buf) + .arg(&per_sample_loss_buf_ptr) .arg(&total_loss_dev_ptr) .arg(&b) .launch(LaunchConfig { @@ -8791,16 +8881,24 @@ impl GpuDqnTrainer { // Grid = B*NA (one thread per (sample, atom), loops over 4 branches). // Zero atomicAdd — fully deterministic gradient computation. let blocks = ((b * na + 255) / 256) as u32; + let save_current_lp_ptr = self.save_current_lp.raw_ptr(); + let save_projected_ptr = self.save_projected.raw_ptr(); + let is_weights_buf_ptr = self.is_weights_buf.raw_ptr(); + let actions_buf_ptr = self.actions_buf.raw_ptr(); + let d_value_logits_buf_ptr = self.d_value_logits_buf.raw_ptr(); + let d_adv_logits_buf_ptr = self.d_adv_logits_buf.raw_ptr(); + let liquid_mod_buf_ptr = self.liquid_mod_buf.raw_ptr(); + let atom_positions_buf_ptr = self.atom_positions_buf.raw_ptr(); unsafe { self.stream .launch_builder(&self.c51_grad_kernel) - .arg(&self.save_current_lp) - .arg(&self.save_projected) - .arg(&self.is_weights_buf) - .arg(&self.actions_buf) - .arg(&self.d_value_logits_buf) - .arg(&self.d_adv_logits_buf) + .arg(&save_current_lp_ptr) + .arg(&save_projected_ptr) + .arg(&is_weights_buf_ptr) + .arg(&actions_buf_ptr) + .arg(&d_value_logits_buf_ptr) + .arg(&d_adv_logits_buf_ptr) .arg(&batch_i32) .arg(&na_i32) .arg(&b0_i32) @@ -8811,8 +8909,8 @@ impl GpuDqnTrainer { .arg(&entropy_coeff) .arg(&self.branch_scales_ptr) .arg(&self.per_sample_support_ptr) - .arg(&self.liquid_mod_buf.raw_ptr()) - .arg(&self.atom_positions_buf) + .arg(&liquid_mod_buf_ptr) + .arg(&atom_positions_buf_ptr) .arg(&self.q_mean_ema_dev_ptr) // unused but preserves 19-param graph node structure on Hopper .launch(LaunchConfig { grid_dim: (blocks, 1, 1), @@ -8839,6 +8937,22 @@ impl GpuDqnTrainer { let b2 = self.config.branch_2_size; let b3 = self.config.branch_3_size; + // Extract raw u64 pointers for CudaSlice fields (eliminates cudarc overhead) + let on_v_logits_buf_ptr = self.on_v_logits_buf.raw_ptr(); + let tg_v_logits_buf_ptr = self.tg_v_logits_buf.raw_ptr(); + let on_next_v_logits_buf_ptr = self.on_next_v_logits_buf.raw_ptr(); + let actions_buf_ptr = self.actions_buf.raw_ptr(); + let rewards_buf_ptr = self.rewards_buf.raw_ptr(); + let dones_buf_ptr = self.dones_buf.raw_ptr(); + let is_weights_buf_ptr = self.is_weights_buf.raw_ptr(); + let per_sample_loss_buf_ptr = self.per_sample_loss_buf.raw_ptr(); + let td_errors_buf_ptr = self.td_errors_buf.raw_ptr(); + let save_current_lp_ptr = self.save_current_lp.raw_ptr(); + let save_projected_ptr = self.save_projected.raw_ptr(); + let curiosity_error_buf_ptr = self.curiosity_error_buf.raw_ptr(); + let drawdown_depths_buf_ptr = self.drawdown_depths_buf.raw_ptr(); + let ensemble_std_buf_ptr = self.ensemble_std_buf.raw_ptr(); + // Extract raw pointers for online branch logits (contiguous [B, (B0+B1+B2+B3)*NA], f32) let f32_sz = std::mem::size_of::(); let on_b_base = self.on_b_logits_buf.raw_ptr(); @@ -8880,38 +8994,38 @@ impl GpuDqnTrainer { self.stream .launch_builder(kernel) // ── Online on current states (5) ── - .arg(&self.on_v_logits_buf) + .arg(&on_v_logits_buf_ptr) .arg(&on_b0_ptr) .arg(&on_b1_ptr) .arg(&on_b2_ptr) .arg(&on_b3_ptr) // ── Target on next_states (5) ── - .arg(&self.tg_v_logits_buf) + .arg(&tg_v_logits_buf_ptr) .arg(&tg_b0_ptr) .arg(&tg_b1_ptr) .arg(&tg_b2_ptr) .arg(&tg_b3_ptr) // ── Online-next on next_states (Double DQN, 5) ── - .arg(&self.on_next_v_logits_buf) + .arg(&on_next_v_logits_buf_ptr) .arg(&on_next_b0_ptr) .arg(&on_next_b1_ptr) .arg(&on_next_b2_ptr) .arg(&on_next_b3_ptr) // ── Batch data (4) ── - .arg(&self.actions_buf) - .arg(&self.rewards_buf) - .arg(&self.dones_buf) - .arg(&self.is_weights_buf) + .arg(&actions_buf_ptr) + .arg(&rewards_buf_ptr) + .arg(&dones_buf_ptr) + .arg(&is_weights_buf_ptr) // ── Outputs (3) ── - .arg(&self.per_sample_loss_buf) - .arg(&self.td_errors_buf) + .arg(&per_sample_loss_buf_ptr) + .arg(&td_errors_buf_ptr) .arg(&self.mse_loss_dev_ptr) // MSE writes to separate accumulator (not total_loss) // ── Saved for backward (2) ── repurposed: save_current_lp = softmax probs, // save_projected = per-branch E[Q] values (4 floats per sample per branch) - .arg(&self.save_current_lp) - .arg(&self.save_projected) + .arg(&save_current_lp_ptr) + .arg(&save_projected_ptr) // ── Curiosity Q-penalty (2) ── - .arg(&self.curiosity_error_buf) + .arg(&curiosity_error_buf_ptr) .arg(&self.config.curiosity_q_penalty_lambda) // ── Config (8 — per_sample_support replaces v_min+v_max) ── .arg(&gamma) @@ -8923,10 +9037,10 @@ impl GpuDqnTrainer { .arg(&b2_i32) .arg(&b3_i32) // ── #18 Asymmetric DD loss (2) ── - .arg(&self.drawdown_depths_buf) + .arg(&drawdown_depths_buf_ptr) .arg(&self.asymmetric_dd_weight) // ── #27 Ensemble disagreement (2) ── - .arg(&self.ensemble_std_buf) + .arg(&ensemble_std_buf_ptr) .arg(&self.ensemble_disagreement_weight) // ── Spectral decoupling (1) ── .arg(&self.config.spectral_decoupling_lambda) @@ -8982,6 +9096,12 @@ impl GpuDqnTrainer { let b2_i32 = b2 as i32; let b3_i32 = b3 as i32; let total_branch_atoms_i32 = total_branch_atoms as i32; + let save_current_lp_ptr = self.save_current_lp.raw_ptr(); + let save_projected_ptr = self.save_projected.raw_ptr(); + let is_weights_buf_ptr = self.is_weights_buf.raw_ptr(); + let actions_buf_ptr = self.actions_buf.raw_ptr(); + let d_value_dst_ptr = d_value_dst.raw_ptr(); + let d_adv_dst_ptr = d_adv_dst.raw_ptr(); // Grid = B*NA (deterministic: one thread per (sample, atom), loops 4 branches) let blocks = ((b * na + 255) / 256) as u32; @@ -8989,12 +9109,12 @@ impl GpuDqnTrainer { unsafe { self.stream .launch_builder(&self.mse_grad_kernel) - .arg(&self.save_current_lp) - .arg(&self.save_projected) - .arg(&self.is_weights_buf) - .arg(&self.actions_buf) - .arg(d_value_dst) - .arg(d_adv_dst) + .arg(&save_current_lp_ptr) + .arg(&save_projected_ptr) + .arg(&is_weights_buf_ptr) + .arg(&actions_buf_ptr) + .arg(&d_value_dst_ptr) + .arg(&d_adv_dst_ptr) .arg(&batch_i32) .arg(&na_i32) .arg(&b0_i32)