diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 677c8967f..aa2299f9f 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -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::(); @@ -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::()) 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::(); @@ -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::()) 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)