perf: convert all CudaSlice kernel args to raw u64 — eliminate cudarc overhead in graph

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-18 23:01:51 +02:00
parent 4b6d9d2e82
commit c7185d2483

View File

@@ -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::<f32>();
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(&param_sizes, 52 + branch);
let out_offset = (branch * na * std::mem::size_of::<f32>()) 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::<f32>(), self.stream.cu_stream(),
homeostatic_total_buf_ptr, 0, std::mem::size_of::<f32>(), 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(&param_sizes, 74);
let b_gamma = self.ptrs.params_ptr + padded_byte_offset(&param_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(&param_sizes, 80);
let b_route = self.ptrs.params_ptr + padded_byte_offset(&param_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(&param_sizes, 66);
let b_out_ptr = self.ptrs.params_ptr + padded_byte_offset(&param_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::<f32>(),
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::<f32>(),
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::<f32>(),
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::<f32>(),
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::<f32>();
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::<f32>();
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)