diff --git a/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu b/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu index f6a6cebfe..c5187d2cc 100644 --- a/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu +++ b/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu @@ -25,7 +25,7 @@ extern "C" __global__ void c51_grad_kernel( float entropy_coeff, const float* __restrict__ branch_scales, /* [B, 4] per-sample per-branch gradient scale */ const float* __restrict__ per_sample_support, /* [B, 3] per-sample [v_min, v_max, delta_z] */ - const float* __restrict__ spread_velocity) /* [1] pinned device-mapped modulator */ + const float* __restrict__ liquid_mod) /* [4] pinned device-mapped per-branch modulators */ { int tid = blockIdx.x * blockDim.x + threadIdx.x; int total_elems = batch_size * num_atoms; @@ -87,7 +87,7 @@ extern "C" __global__ void c51_grad_kernel( /* Scale proportional to delta_z — tighter atoms = smaller spread needed. * No hardcoded constants: spread = inv_batch * delta_z (same order as CE grad). */ float delta_z = per_sample_support[b * 3 + 2]; - float velocity_mod = spread_velocity[0]; + float velocity_mod = liquid_mod[d]; float spread_scale = inv_batch * delta_z * velocity_mod; /* Asymmetric spread: challenger (adjacent to taken) gets pushed UP. diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index f2aeb5b46..d76dc3b8d 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -711,13 +711,6 @@ pub struct GpuDqnTrainer { /// EMA of Q-divergence for adaptive tau computation. q_div_ema: f32, - /// Q-gap EMA for momentum-modulated spread gradient. - q_gap_ema: f32, - /// Spread velocity scalar — written by CPU, read by GPU via pinned device-mapped memory. - spread_velocity_pinned: *mut f32, - /// Device pointer for spread_velocity_pinned (from cuMemHostGetDevicePointer). - spread_velocity_dev_ptr: u64, - // ── Training state ────────────────────────────────────────────── pub(crate) adam_step: i32, total_params: usize, @@ -1063,9 +1056,6 @@ impl Drop for GpuDqnTrainer { if !self.adaptive_clip_pinned.is_null() { let _ = unsafe { cudarc::driver::result::free_host(self.adaptive_clip_pinned.cast()) }; } - if !self.spread_velocity_pinned.is_null() { - let _ = unsafe { cudarc::driver::result::free_host(self.spread_velocity_pinned.cast()) }; - } if !self.liquid_mod_pinned.is_null() { let _ = unsafe { cudarc::driver::result::free_host(self.liquid_mod_pinned.cast()) }; } @@ -1086,9 +1076,42 @@ impl GpuDqnTrainer { pub fn eval_v_range_ptr(&self) -> u64 { self.eval_v_range_ptr } + /// Update per-branch liquid tau modulation from per-branch Q-gaps. + /// + /// Liquid ODE per branch d: + /// velocity_d = q_gap_d - q_gap_ema_d + /// tau_d = tau_min + (tau_max - tau_min) * sigmoid(velocity_d / delta_z) + /// liquid_mod_d += (1/tau_d) * (1.0 - liquid_mod_d) + pub fn update_liquid_tau(&mut self, per_branch_q_gaps: [f32; 4]) { + let tau_min = 0.01_f32; + let tau_max = 1.0_f32; + let delta_z_approx = self.eval_q_std_ema.max(0.01); + + for d in 0..4 { + let q_gap_d = per_branch_q_gaps[d]; + + if self.per_branch_q_gap_ema[d] < 1e-12 { + self.per_branch_q_gap_ema[d] = q_gap_d; + } else { + let err = (q_gap_d - self.per_branch_q_gap_ema[d]).abs(); + let alpha = (err / (err + self.per_branch_q_gap_ema[d].max(0.001))).clamp(0.01, 0.3); + self.per_branch_q_gap_ema[d] = (1.0 - alpha) * self.per_branch_q_gap_ema[d] + alpha * q_gap_d; + } + + let velocity_d = q_gap_d - self.per_branch_q_gap_ema[d]; + let sigmoid_arg = velocity_d / delta_z_approx; + let sig = 1.0 / (1.0 + (-sigmoid_arg).exp()); + let tau_d = tau_min + (tau_max - tau_min) * sig; + + let current = unsafe { *self.liquid_mod_pinned.add(d) }; + let updated = current + (1.0 / tau_d) * (1.0 - current); + unsafe { *self.liquid_mod_pinned.add(d) = updated.clamp(0.1, 2.0); } + } + } + /// Update eval v_range from observed Q-value statistics. /// Called at epoch boundary after compute_q_stats. - pub fn update_eval_v_range(&mut self, q_mean: f32, q_std: f32, q_gap: f32) { + pub fn update_eval_v_range(&mut self, q_mean: f32, q_std: f32, q_gap: f32, per_branch_q_gaps: [f32; 4]) { // Adaptive-rate EMA with baseline proportional to Q-std. // Baseline = max(q_std_ema, 0.001) so normal Q-oscillation doesn't // trigger aggressive tracking — only abnormal shifts do. @@ -1107,19 +1130,8 @@ impl GpuDqnTrainer { self.eval_q_std_ema = (1.0 - alpha_std) * self.eval_q_std_ema + alpha_std * q_std; } - // Q-gap momentum: modulate spread gradient based on Q-gap velocity - let q_gap_val = q_gap; - if self.q_gap_ema < 1e-12 { - self.q_gap_ema = q_gap_val; - } else { - let err = (q_gap_val - self.q_gap_ema).abs(); - let alpha = (err / (err + self.q_gap_ema.max(0.001))).clamp(0.01, 0.3); - self.q_gap_ema = (1.0 - alpha) * self.q_gap_ema + alpha * q_gap_val; - } - let velocity = q_gap_val - self.q_gap_ema; - let delta_z_approx = self.eval_q_std_ema.max(0.01); - let modulator = (1.0 - velocity / delta_z_approx).clamp(0.1, 2.0); - unsafe { *self.spread_velocity_pinned = modulator; } + // Replace old spread_velocity with per-branch liquid tau ODE + self.update_liquid_tau(per_branch_q_gaps); // Width from Q-gap (action differentiation range), not Q-std. // Floor = max(10 * q_gap, 3 * q_std) — enough atoms to resolve actions. @@ -2582,23 +2594,6 @@ impl GpuDqnTrainer { dev_ptr }; - let spread_velocity_pinned: *mut f32 = unsafe { - let flags = cudarc::driver::sys::CU_MEMHOSTALLOC_DEVICEMAP; - cudarc::driver::result::malloc_host(std::mem::size_of::(), flags) - .map_err(|e| MLError::ModelError(format!("pinned spread_velocity alloc: {e}")))? - as *mut f32 - }; - unsafe { *spread_velocity_pinned = 1.0; } // Full spread initially - let spread_velocity_dev_ptr = unsafe { - let mut dev_ptr: u64 = 0; - cudarc::driver::sys::cuMemHostGetDevicePointer_v2( - &mut dev_ptr as *mut u64, - spread_velocity_pinned.cast(), - 0, - ); - dev_ptr - }; - // ── Allocate consolidated transfer buffers ───────────────── // Upload staging: states + next_states (f32) // Rewards/dones uploaded separately as f32 @@ -3330,9 +3325,6 @@ impl GpuDqnTrainer { adaptive_clip_dev_ptr, grad_norm_ema: 0.0, q_div_ema: 0.0, - q_gap_ema: 0.0, - spread_velocity_pinned, - spread_velocity_dev_ptr, adam_step: 0, total_params, params_initialized: false, @@ -5651,7 +5643,7 @@ impl GpuDqnTrainer { .arg(&entropy_coeff) .arg(&self.branch_scales_ptr) .arg(&self.per_sample_support_ptr) - .arg(&self.spread_velocity_dev_ptr) + .arg(&self.liquid_mod_dev_ptr) .launch(LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index ab77e79b1..f75788904 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -2051,8 +2051,8 @@ impl FusedTrainingCtx { vr } /// Update eval v_range from observed Q-value statistics. - pub(crate) fn update_eval_v_range(&mut self, q_mean: f32, q_std: f32, q_gap: f32) { - self.trainer.update_eval_v_range(q_mean, q_std, q_gap); + pub(crate) fn update_eval_v_range(&mut self, q_mean: f32, q_std: f32, q_gap: f32, per_branch_q_gaps: [f32; 4]) { + self.trainer.update_eval_v_range(q_mean, q_std, q_gap, per_branch_q_gaps); } /// Per-sample epsilon from IQL expectile gap. pub(crate) fn per_sample_epsilon_ptr(&self) -> u64 { self.gpu_iql.per_sample_epsilon_ptr() } diff --git a/crates/ml/src/trainers/dqn/trainer/training_loop.rs b/crates/ml/src/trainers/dqn/trainer/training_loop.rs index 33f5fc66f..5bcb4636e 100644 --- a/crates/ml/src/trainers/dqn/trainer/training_loop.rs +++ b/crates/ml/src/trainers/dqn/trainer/training_loop.rs @@ -1391,7 +1391,10 @@ impl DQNTrainer { // Update eval v_range from observed Q-stats + Q-gap let q_std = stats.q_variance.max(0.0).sqrt(); let q_gap = (stats.avg_max_q as f32 - stats.q_mean).max(0.0); - fused.update_eval_v_range(stats.q_mean, q_std, q_gap); + // Simple approximation: use global q_gap for all branches initially. + // Per-branch extraction requires reading 12 Q-values from GPU. + let per_branch_q_gaps = [q_gap; 4]; // uniform until per-branch extraction is added + fused.update_eval_v_range(stats.q_mean, q_std, q_gap, per_branch_q_gaps); } } train_step_count += 1;