feat: per-branch liquid tau ODE replaces global spread_velocity

c51_grad_kernel reads liquid_mod[d] per branch instead of spread_velocity[0].
4 per-branch Q-gap EMAs with adaptive alpha drive continuous-time ODE.
Fast branch learning → large tau → slow adaptation → don't overshoot.
Stuck branch → small tau → fast adaptation → push harder.
spread_velocity infrastructure removed (-30 lines).

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-15 00:22:52 +02:00
parent 90f4e0531b
commit 70f1446965
4 changed files with 45 additions and 50 deletions

View File

@@ -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.

View File

@@ -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::<f32>(), 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),

View File

@@ -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() }

View File

@@ -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;