fix(4branch): add branch 3 args to all 6 loss/grad kernel launches — fixes CUDA_ERROR_INVALID_VALUE
The CUDA kernels (.cu files) were updated to accept 4 branch pointers and b3_size, but the Rust launch functions still only passed 3 branches. This caused CUDA_ERROR_INVALID_VALUE on H100 due to argument count mismatch. Fixed 7 launch sites in gpu_dqn_trainer.rs: - launch_c51_loss: added on_b3/tg_b3/on_next_b3 pointers + b3_i32 - launch_c51_mixup: added branch_3_size arg - launch_c51_grad: added b3_i32, fixed thread count 3*na → 4*na - launch_mse_loss: added on_b3/tg_b3/on_next_b3 pointers + b3_i32 - launch_mse_grad_inner: added b3_i32, fixed thread count 3*na → 4*na - apply_cql_gradient: added b3_i32 - compute_expected_q: added b3 Also updated max_branch calculations to include b3. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -1717,6 +1717,7 @@ impl GpuDqnTrainer {
|
||||
let b0_i32 = b0 as i32;
|
||||
let b1_i32 = b1 as i32;
|
||||
let b2_i32 = b2 as i32;
|
||||
let b3_i32 = b3 as i32;
|
||||
let v_min = self.config.v_min;
|
||||
let v_max = self.config.v_max;
|
||||
|
||||
@@ -1735,6 +1736,7 @@ impl GpuDqnTrainer {
|
||||
.arg(&b0_i32)
|
||||
.arg(&b1_i32)
|
||||
.arg(&b2_i32)
|
||||
.arg(&b3_i32)
|
||||
.arg(&v_min)
|
||||
.arg(&v_max)
|
||||
.launch(LaunchConfig {
|
||||
@@ -4058,6 +4060,7 @@ impl GpuDqnTrainer {
|
||||
let b0 = self.config.branch_0_size as i32;
|
||||
let b1 = self.config.branch_1_size as i32;
|
||||
let b2 = self.config.branch_2_size as i32;
|
||||
let b3 = self.config.branch_3_size as i32;
|
||||
let v_min = self.config.v_min;
|
||||
let v_max = self.config.v_max;
|
||||
|
||||
@@ -4080,6 +4083,7 @@ impl GpuDqnTrainer {
|
||||
.arg(&b0)
|
||||
.arg(&b1)
|
||||
.arg(&b2)
|
||||
.arg(&b3)
|
||||
.arg(&v_min)
|
||||
.arg(&v_max)
|
||||
.launch(LaunchConfig {
|
||||
@@ -4999,6 +5003,7 @@ impl GpuDqnTrainer {
|
||||
let b0 = self.config.branch_0_size;
|
||||
let b1 = self.config.branch_1_size;
|
||||
let b2 = self.config.branch_2_size;
|
||||
let b3 = self.config.branch_3_size;
|
||||
|
||||
// Extract raw pointers for online branch logits (contiguous [B, (B0+B1+B2+B3)*NA], f32)
|
||||
let f32_sz = std::mem::size_of::<f32>();
|
||||
@@ -5006,18 +5011,21 @@ impl GpuDqnTrainer {
|
||||
let on_b0_ptr = on_b_base;
|
||||
let on_b1_ptr = on_b_base + (b * b0 * na * f32_sz) as u64;
|
||||
let on_b2_ptr = on_b1_ptr + (b * b1 * na * f32_sz) as u64;
|
||||
let on_b3_ptr = on_b2_ptr + (b * b2 * na * f32_sz) as u64;
|
||||
|
||||
// Target branch logits (f32)
|
||||
let tg_b_base = self.tg_b_logits_buf.raw_ptr();
|
||||
let tg_b0_ptr = tg_b_base;
|
||||
let tg_b1_ptr = tg_b_base + (b * b0 * na * f32_sz) as u64;
|
||||
let tg_b2_ptr = tg_b1_ptr + (b * b1 * na * f32_sz) as u64;
|
||||
let tg_b3_ptr = tg_b2_ptr + (b * b2 * na * f32_sz) as u64;
|
||||
|
||||
// Online-next branch logits (Double DQN action selector, f32)
|
||||
let on_next_b_base = self.on_next_b_logits_buf.raw_ptr();
|
||||
let on_next_b0_ptr = on_next_b_base;
|
||||
let on_next_b1_ptr = on_next_b_base + (b * b0 * na * f32_sz) as u64;
|
||||
let on_next_b2_ptr = on_next_b1_ptr + (b * b1 * na * f32_sz) as u64;
|
||||
let on_next_b3_ptr = on_next_b2_ptr + (b * b2 * na * f32_sz) as u64;
|
||||
|
||||
// N-step returns: use gamma^n for the Bellman projection.
|
||||
// The experience collector pre-computes R_n = sum(gamma^i * r_i).
|
||||
@@ -5029,10 +5037,11 @@ impl GpuDqnTrainer {
|
||||
let b0_i32 = b0 as i32;
|
||||
let b1_i32 = b1 as i32;
|
||||
let b2_i32 = b2 as i32;
|
||||
let b3_i32 = b3 as i32;
|
||||
|
||||
// Shared memory: float arrays (support + val + adv + lp + proj + current_lp + reduce)
|
||||
// Float (4 bytes/elem) for numerically stable softmax/projection
|
||||
let max_branch = b0.max(b1).max(b2);
|
||||
let max_branch = b0.max(b1).max(b2).max(b3);
|
||||
let shmem_floats = na + na + max_branch * na + na + na + na + 8;
|
||||
let shmem_bytes = (shmem_floats * std::mem::size_of::<f32>()) as u32;
|
||||
|
||||
@@ -5042,21 +5051,24 @@ impl GpuDqnTrainer {
|
||||
unsafe {
|
||||
self.stream
|
||||
.launch_builder(kernel)
|
||||
// ── Online on current states (4, f32 logits) ──
|
||||
// ── Online on current states (5, f32 logits) ──
|
||||
.arg(&self.on_v_logits_buf)
|
||||
.arg(&on_b0_ptr)
|
||||
.arg(&on_b1_ptr)
|
||||
.arg(&on_b2_ptr)
|
||||
// ── Target on next_states (4, f32 logits) ──
|
||||
.arg(&on_b3_ptr)
|
||||
// ── Target on next_states (5, f32 logits) ──
|
||||
.arg(&self.tg_v_logits_buf)
|
||||
.arg(&tg_b0_ptr)
|
||||
.arg(&tg_b1_ptr)
|
||||
.arg(&tg_b2_ptr)
|
||||
// ── Online-next on next_states (Double DQN, 4, f32 logits) ──
|
||||
.arg(&tg_b3_ptr)
|
||||
// ── Online-next on next_states (Double DQN, 5, f32 logits) ──
|
||||
.arg(&self.on_next_v_logits_buf)
|
||||
.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)
|
||||
@@ -5072,7 +5084,7 @@ impl GpuDqnTrainer {
|
||||
// ── Curiosity Q-penalty (2) ──
|
||||
.arg(&self.curiosity_error_buf)
|
||||
.arg(&self.config.curiosity_q_penalty_lambda)
|
||||
// ── Config (8) ──
|
||||
// ── Config (9) ──
|
||||
.arg(&gamma)
|
||||
.arg(&batch_i32)
|
||||
.arg(&na_i32)
|
||||
@@ -5081,6 +5093,7 @@ impl GpuDqnTrainer {
|
||||
.arg(&b0_i32)
|
||||
.arg(&b1_i32)
|
||||
.arg(&b2_i32)
|
||||
.arg(&b3_i32)
|
||||
// ── #18 Asymmetric DD loss (2) ──
|
||||
.arg(&self.drawdown_depths_buf)
|
||||
.arg(&self.asymmetric_dd_weight)
|
||||
@@ -5140,6 +5153,7 @@ impl GpuDqnTrainer {
|
||||
.arg(&(self.config.branch_0_size as i32))
|
||||
.arg(&(self.config.branch_1_size as i32))
|
||||
.arg(&(self.config.branch_2_size as i32))
|
||||
.arg(&(self.config.branch_3_size as i32))
|
||||
.launch(LaunchConfig {
|
||||
grid_dim: (b as u32, 1, 1),
|
||||
block_dim: (256, 1, 1),
|
||||
@@ -5168,10 +5182,11 @@ impl GpuDqnTrainer {
|
||||
let b0_i32 = b0 as i32;
|
||||
let b1_i32 = b1 as i32;
|
||||
let b2_i32 = b2 as i32;
|
||||
let b3_i32 = b3 as i32;
|
||||
let total_branch_atoms_i32 = total_branch_atoms as i32;
|
||||
let entropy_coeff = self.config.entropy_coefficient;
|
||||
|
||||
let blocks = ((b * 3 * na + 255) / 256) as u32;
|
||||
let blocks = ((b * 4 * na + 255) / 256) as u32;
|
||||
|
||||
unsafe {
|
||||
self.stream
|
||||
@@ -5187,6 +5202,7 @@ impl GpuDqnTrainer {
|
||||
.arg(&b0_i32)
|
||||
.arg(&b1_i32)
|
||||
.arg(&b2_i32)
|
||||
.arg(&b3_i32)
|
||||
.arg(&total_branch_atoms_i32)
|
||||
.arg(&entropy_coeff)
|
||||
.launch(LaunchConfig {
|
||||
@@ -5212,6 +5228,7 @@ impl GpuDqnTrainer {
|
||||
let b0 = self.config.branch_0_size;
|
||||
let b1 = self.config.branch_1_size;
|
||||
let b2 = self.config.branch_2_size;
|
||||
let b3 = self.config.branch_3_size;
|
||||
|
||||
// Extract raw pointers for online branch logits (contiguous [B, (B0+B1+B2+B3)*NA], f32)
|
||||
let f32_sz = std::mem::size_of::<f32>();
|
||||
@@ -5219,18 +5236,21 @@ impl GpuDqnTrainer {
|
||||
let on_b0_ptr = on_b_base;
|
||||
let on_b1_ptr = on_b_base + (b * b0 * na * f32_sz) as u64;
|
||||
let on_b2_ptr = on_b1_ptr + (b * b1 * na * f32_sz) as u64;
|
||||
let on_b3_ptr = on_b2_ptr + (b * b2 * na * f32_sz) as u64;
|
||||
|
||||
// Target branch logits (f32)
|
||||
let tg_b_base = self.tg_b_logits_buf.raw_ptr();
|
||||
let tg_b0_ptr = tg_b_base;
|
||||
let tg_b1_ptr = tg_b_base + (b * b0 * na * f32_sz) as u64;
|
||||
let tg_b2_ptr = tg_b1_ptr + (b * b1 * na * f32_sz) as u64;
|
||||
let tg_b3_ptr = tg_b2_ptr + (b * b2 * na * f32_sz) as u64;
|
||||
|
||||
// Online-next branch logits (Double DQN action selector, f32)
|
||||
let on_next_b_base = self.on_next_b_logits_buf.raw_ptr();
|
||||
let on_next_b0_ptr = on_next_b_base;
|
||||
let on_next_b1_ptr = on_next_b_base + (b * b0 * na * f32_sz) as u64;
|
||||
let on_next_b2_ptr = on_next_b1_ptr + (b * b1 * na * f32_sz) as u64;
|
||||
let on_next_b3_ptr = on_next_b2_ptr + (b * b2 * na * f32_sz) as u64;
|
||||
|
||||
// N-step returns: use gamma^n for the Bellman projection.
|
||||
let gamma = self.config.gamma.powi(self.config.n_steps as i32);
|
||||
@@ -5241,31 +5261,35 @@ impl GpuDqnTrainer {
|
||||
let b0_i32 = b0 as i32;
|
||||
let b1_i32 = b1 as i32;
|
||||
let b2_i32 = b2 as i32;
|
||||
let b3_i32 = b3 as i32;
|
||||
|
||||
// Shared memory: float arrays (support + val + adv + lp + proj + current_lp + reduce)
|
||||
// Float (4 bytes/elem) for numerically stable softmax — BF16 exp() overflows at logit > 11.1
|
||||
let max_branch = b0.max(b1).max(b2);
|
||||
let max_branch = b0.max(b1).max(b2).max(b3);
|
||||
let shmem_floats = na + na + max_branch * na + na + na + na + 8;
|
||||
let shmem_bytes = (shmem_floats * std::mem::size_of::<f32>()) as u32;
|
||||
|
||||
unsafe {
|
||||
self.stream
|
||||
.launch_builder(kernel)
|
||||
// ── Online on current states (4) ──
|
||||
// ── Online on current states (5) ──
|
||||
.arg(&self.on_v_logits_buf)
|
||||
.arg(&on_b0_ptr)
|
||||
.arg(&on_b1_ptr)
|
||||
.arg(&on_b2_ptr)
|
||||
// ── Target on next_states (4) ──
|
||||
.arg(&on_b3_ptr)
|
||||
// ── Target on next_states (5) ──
|
||||
.arg(&self.tg_v_logits_buf)
|
||||
.arg(&tg_b0_ptr)
|
||||
.arg(&tg_b1_ptr)
|
||||
.arg(&tg_b2_ptr)
|
||||
// ── Online-next on next_states (Double DQN, 4) ──
|
||||
.arg(&tg_b3_ptr)
|
||||
// ── Online-next on next_states (Double DQN, 5) ──
|
||||
.arg(&self.on_next_v_logits_buf)
|
||||
.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)
|
||||
@@ -5276,13 +5300,13 @@ impl GpuDqnTrainer {
|
||||
.arg(&self.td_errors_buf)
|
||||
.arg(&self.mse_loss_buf) // MSE writes to separate accumulator (not total_loss_buf)
|
||||
// ── Saved for backward (2) ── repurposed: save_current_lp = softmax probs,
|
||||
// save_projected = per-branch E[Q] values (3 floats per sample per branch)
|
||||
// save_projected = per-branch E[Q] values (4 floats per sample per branch)
|
||||
.arg(&self.save_current_lp)
|
||||
.arg(&self.save_projected)
|
||||
// ── Curiosity Q-penalty (2) ──
|
||||
.arg(&self.curiosity_error_buf)
|
||||
.arg(&self.config.curiosity_q_penalty_lambda)
|
||||
// ── Config (8) ──
|
||||
// ── Config (9) ──
|
||||
.arg(&gamma)
|
||||
.arg(&batch_i32)
|
||||
.arg(&na_i32)
|
||||
@@ -5291,6 +5315,7 @@ impl GpuDqnTrainer {
|
||||
.arg(&b0_i32)
|
||||
.arg(&b1_i32)
|
||||
.arg(&b2_i32)
|
||||
.arg(&b3_i32)
|
||||
// ── #18 Asymmetric DD loss (2) ──
|
||||
.arg(&self.drawdown_depths_buf)
|
||||
.arg(&self.asymmetric_dd_weight)
|
||||
@@ -5349,11 +5374,12 @@ impl GpuDqnTrainer {
|
||||
let b0_i32 = b0 as i32;
|
||||
let b1_i32 = b1 as i32;
|
||||
let b2_i32 = b2 as i32;
|
||||
let b3_i32 = b3 as i32;
|
||||
let total_branch_atoms_i32 = total_branch_atoms as i32;
|
||||
let v_min = self.config.v_min;
|
||||
let v_max = self.config.v_max;
|
||||
|
||||
let blocks = ((b * 3 * na + 255) / 256) as u32;
|
||||
let blocks = ((b * 4 * na + 255) / 256) as u32;
|
||||
|
||||
unsafe {
|
||||
self.stream
|
||||
@@ -5369,6 +5395,7 @@ impl GpuDqnTrainer {
|
||||
.arg(&b0_i32)
|
||||
.arg(&b1_i32)
|
||||
.arg(&b2_i32)
|
||||
.arg(&b3_i32)
|
||||
.arg(&total_branch_atoms_i32)
|
||||
.arg(&v_min)
|
||||
.arg(&v_max)
|
||||
|
||||
Reference in New Issue
Block a user