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:
@@ -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(¶m_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(¶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::<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)
|
||||
|
||||
Reference in New Issue
Block a user