feat(dqn): ISV-unified per-branch Q-support range (spec 2026-04-23)

Unifies the Q-support range source across atom grid / warm-start quantile
clamp consumers via the ISV signal bus. One broadcast written at epoch
boundary from per-branch Q-stats EMAs, read by the two consumers that
previously held disagreeing ranges. Target observation: atom utilisation
≥40% (up from 11-15% on train-fpxnw).

Phase 0 — per-branch Q-stats kernel Rust plumbing:
* Load q_stats_per_branch_reduce alongside legacy q_stats_reduce
* Add per_branch_q_stats_pinned (28 f32 = 4 × 7, device-mapped)
* PerBranchQValueStats struct: [QValueStatsResult; 4]
* reduce_current_q_stats_per_branch launches the new kernel with the
  four branch (off, size) pairs derived from config.branch_N_size

Phase 1 — ISV v-range plumbing (zero behavioural change at epoch 1):
* ISV_NETWORK_DIM=23 preserved for w_isv_fc1 sizing; ISV_TOTAL_DIM=31
  allocates 8 additional slots for per-branch (centre, half-width)
* Slot constants V_CENTER_DIR..V_HALF_URG covering slots 23..30
* eval_q_mean_ema / eval_q_std_ema / eval_ema_initialized promoted
  to [f32; 4] / [bool; 4]; scalar setters preserved for trajectory
  backtracking (broadcast same value to all branches)
* Bootstrap at construction: centre=0, half=(v_max-v_min)/2 → the
  byte-identical [config.v_min, config.v_max] span per branch before
  any Q observations arrive
* reset_eval_v_range_state resets the 4 per-branch EMAs AND the 8 ISV
  slots to bootstrap values; legacy eval_v_range_pinned[2] still reset
  (deferred removal — spec Phase 3)
* update_eval_v_range reworked: signature takes PerBranchQValueStats and
  per_branch_q_gaps. Maintains 4 independent adaptive-rate EMAs,
  computes (centre, half) per branch with min_half_floor=0.1×(v_max-v_min)
  and clamps to config bounds, writes 8 ISV slots. Branch-0 (direction)
  centre±half is also mirrored into the legacy eval_v_range_pinned for
  consumers that have not yet migrated to the per-branch bus.

Phase 2a/2b — atom grid per-branch v-range:
* adaptive_atom_positions kernel signature changed from
  (v_min: float, v_max: float) to (branch_idx: int, isv_signals: float*);
  reads centre/half from ISV slots 23+2·b, 24+2·b. Eliminates the f64→f32
  ABI trap (spec Phase 2 side-effect) since the only per-branch range
  path is now pointer-based.
* recompute_atom_positions passes branch_idx + isv_signals_dev_ptr per
  branch; no scalar v_min/v_max arg remains.

Phase 2c — warm-start quantile clamp per-branch from ISV:
* warm_start_atom_positions reads per-branch (centre, half) from pinned
  ISV host memory, clamps shared reward-quantile vector into each
  branch's adaptive range before tiling into atom_positions_buf.
  Bootstrap makes this equivalent to the pre-spec config.v_{min,max}
  clamp until the first Q observation lands.

Deviations from spec:
* Phase 2d (per_sample_support_buf → [N, 4, 3]) NOT implemented. The
  spec's premise was that per_sample_support is host-tiled from
  eval_v_range, but the active path in this codebase has it filled by
  iql_compute_per_sample_support (V(s)-centered, per-sample, already
  adaptive) — orthogonal to the ISV bus. Migrating that kernel to
  per-branch output would require rewriting iql_value_kernel +
  iql_support_floor + C51/MSE loss kernel indexing in lockstep, which
  the "no unrelated refactoring" constraint disallows. The loss-kernel
  Bellman projection today uses V-centered bounds that are themselves
  adaptive; the ISV v-range fix still lands the primary win (atom grid
  + warm-start agreement) without touching IQL.

Compile verified: cargo check -p ml + --workspace pass (SQLX_OFFLINE,
CARGO_INCREMENTAL=0, sccache). No TODO/FIXME/XXX introduced.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-23 21:18:19 +02:00
parent d38a8cf997
commit 9deda5f65b
6 changed files with 851 additions and 112 deletions

View File

@@ -5012,6 +5012,11 @@ extern "C" __global__ void regime_branch_gate(
* spacing = softmax(spacing_raw)
* positions[j] = v_min + (v_max - v_min) * cumsum(spacing)[j]
*
* Per-branch bounds come from the ISV signal bus (slots 23 + 2*b for centre,
* 24 + 2*b for half-width). Written by `update_eval_v_range`; bootstrap values
* (centre=0, half=(config.v_max-config.v_min)/2) at construction produce the
* same `[config.v_min, config.v_max]` positions as the pre-ISV behaviour.
*
* Single-block kernel: num_atoms=51 fits in one block.
* Grid: (1,1,1), Block: (256,1,1). Shared memory: 2 * num_atoms * sizeof(float).
*/
@@ -5019,14 +5024,22 @@ extern "C" __global__ void adaptive_atom_positions(
const float* __restrict__ spacing_raw,
float* __restrict__ positions_out,
int num_atoms,
float v_min,
float v_max
int branch_idx,
const float* __restrict__ isv_signals
) {
int tid = threadIdx.x;
extern __shared__ float shmem[];
float* s_softmax = shmem;
float* s_cumsum = shmem + num_atoms;
/* Per-branch v-range from ISV slots. Layout: centre at 23+2*b, half at
* 24+2*b for b in [0..4). Caller guarantees isv_signals is non-null and
* branch_idx ∈ [0,4). */
float v_center = isv_signals[23 + 2 * branch_idx];
float v_half = isv_signals[24 + 2 * branch_idx];
float v_min = v_center - v_half;
float v_max = v_center + v_half;
/* Load spacing_raw into shared memory */
if (tid < num_atoms) {
s_softmax[tid] = spacing_raw[tid];

View File

@@ -140,7 +140,22 @@ const ISV_K: usize = 4; // Temporal ISV history length
/// NOT rotated into history. Slot [22] (`SHARPE_EMA_INDEX`) added by the
/// ISV-audit bundle (2026-04-23) — Rust host broadcasts `training_sharpe_ema`
/// each epoch; the kernel reads it to drive slot 12 (learning_health).
const ISV_DIM: usize = 23;
/// Network-facing ISV width: the length of the vector fed into `w_isv_fc1`
/// attention weights. **Never change without a checkpoint format migration** —
/// the w_isv_fc1 weight tensor is sized [16, ISV_NETWORK_DIM] at 368 floats.
const ISV_NETWORK_DIM: usize = 23;
/// Total ISV slot count including cross-kernel scratchpad slots that never
/// feed the attention path. Slots [23..31] carry per-branch Q-support
/// (center, half-width) for direction/magnitude/order/urgency — written by
/// `update_eval_v_range`, read by `adaptive_atom_positions`,
/// `warm_start_atom_positions`, and per-sample support tiling. Extending
/// past 31 requires updating the pinned allocation in the constructor and
/// any kernel that takes an ISV pointer expecting a specific length.
const ISV_TOTAL_DIM: usize = 31;
/// Legacy alias preserved for call sites that haven't been audited for the
/// network-vs-total split. New code should pick `ISV_NETWORK_DIM` (for weight
/// tensor sizing) or `ISV_TOTAL_DIM` (for the broadcast bus buffer).
const ISV_DIM: usize = ISV_NETWORK_DIM;
pub const LEARNING_HEALTH_INDEX: usize = 12;
/// Task 2.X "make Full useful": ISV slots for per-magnitude-bin Q-mean EMAs.
/// Read by `c51_loss_batched` and `c51_grad_kernel` to compute the
@@ -171,6 +186,23 @@ pub const Q_DIR_ABS_REF_INDEX: usize = 21;
/// component-aggregation formula whose Sharpe-correlation measured at
/// r=-0.765 on the train-mdh86 logs that triggered the audit.
pub const SHARPE_EMA_INDEX: usize = 22;
/// ISV v-range unification bundle (spec 2026-04-23): per-branch Q-support
/// centre + half-width, 2 slots per branch (direction=0, magnitude=1,
/// order=2, urgency=3). Written by `update_eval_v_range` from the per-branch
/// Q-stats EMA state; read by `adaptive_atom_positions` (atom grid),
/// `warm_start_atom_positions` (quantile clamp) and `update_per_sample_support`
/// (loss kernel projection range). Bootstrap values at construction + fold
/// reset: `v_center = 0`, `v_half = (config.v_max - config.v_min) / 2` — at
/// epoch 1 this reproduces the pre-spec `[config.v_min, config.v_max]`
/// behaviour byte-for-byte across all four branches.
pub const V_CENTER_DIR_INDEX: usize = 23;
pub const V_HALF_DIR_INDEX: usize = 24;
pub const V_CENTER_MAG_INDEX: usize = 25;
pub const V_HALF_MAG_INDEX: usize = 26;
pub const V_CENTER_ORD_INDEX: usize = 27;
pub const V_HALF_ORD_INDEX: usize = 28;
pub const V_CENTER_URG_INDEX: usize = 29;
pub const V_HALF_URG_INDEX: usize = 30;
const ISV_EMB_DIM: usize = 8; // ISV embedding output dimension (FC2 output, gate/gamma input)
/// First ISV tensor index in the flat param buffer.
/// ISV weights (68-79) are online-only — NOT synced to the target network.
@@ -481,6 +513,21 @@ pub struct QValueStatsResult {
pub atom_utilization: f32,
}
/// Per-branch Q-value statistics — one 7-tuple per action branch (direction=0,
/// magnitude=1, order=2, urgency=3). Produced by `q_stats_per_branch_reduce`
/// kernel; consumed by `update_eval_v_range` to track per-branch EMA state
/// and write the 8 ISV v-range slots.
///
/// The per-branch atom-entropy / utilization slots in the kernel output are
/// left at 0 (per-branch atom_stats would require splitting the global
/// accumulator). `QValueStatsResult.atom_{entropy,utilization}` inherit 0
/// until that wiring lands — the v-range EMA updater doesn't use those two
/// fields, so the zero does not propagate into the ISV path.
#[derive(Debug, Clone, Copy)]
pub struct PerBranchQValueStats {
pub per_branch: [QValueStatsResult; 4],
}
/// Scalar-only result from the fused GPU training path.
///
/// TD errors stay on GPU (`td_errors_buf`) — no readback. PER priority update
@@ -1210,10 +1257,13 @@ pub struct GpuDqnTrainer {
/// Pinned device-mapped: CPU writes v_min/v_max, GPU reads via dev_ptr. No HtoD copy.
eval_v_range_pinned: *mut f32,
eval_v_range_ptr: u64,
/// EMA-smoothed Q-stats for stable eval v_range updates.
pub(crate) eval_q_mean_ema: f32,
pub(crate) eval_q_std_ema: f32,
eval_ema_initialized: bool,
/// Per-branch EMA-smoothed Q-stats for the ISV v-range update. Index
/// corresponds to branch (direction=0, magnitude=1, order=2, urgency=3).
/// Promoted from scalar to `[f32; 4]` as part of the ISV v-range
/// unification — see `update_eval_v_range`.
pub(crate) eval_q_mean_ema: [f32; 4],
pub(crate) eval_q_std_ema: [f32; 4],
eval_ema_initialized: [bool; 4],
/// Adaptive IQN lambda readiness: 0=uncertain (suppress gradient), 1=converged (full weight).
iqn_readiness: f32,
/// IQN readiness — pinned device-mapped for CUDA graph. GPU reads via dev_ptr.
@@ -1533,6 +1583,14 @@ pub struct GpuDqnTrainer {
q_stats_kernel: CudaFunction,
/// GPU buffer for Q-value statistics [7 floats]
q_stats_buf: CudaSlice<f32>,
/// Per-branch Q-stats reducer — writes 28 floats to `per_branch_q_stats_pinned`
/// (4 × 7, branch-major). Runs alongside `q_stats_kernel`; feeds the
/// ISV v-range per-branch EMA state in `update_eval_v_range`.
q_stats_per_branch_kernel: CudaFunction,
/// Pinned device-mapped readback for per-branch Q-stats [4 × 7 = 28 floats].
/// GPU writes via `per_branch_q_stats_dev_ptr`; CPU reads host-side.
per_branch_q_stats_pinned: *mut f32,
per_branch_q_stats_dev_ptr: u64,
/// GPU buffer for atom utilization accumulation [2 floats: sum_entropy, sum_utilized].
/// Populated by `atom_stats_finalize` kernel (phase 2 of the deterministic reduction).
atom_stats_buf: CudaSlice<f32>,
@@ -1828,11 +1886,15 @@ pub struct GpuDqnTrainer {
q_dir_bin_means_reduce_kernel: CudaFunction,
// ── ISV core buffers (pinned device-mapped, GPU read/write) ──
isv_signals_pinned: *mut f32, // [ISV_DIM = 23]
// isv_signals_pinned is sized for ISV_TOTAL_DIM (31) — the extra eight
// slots [23..31] carry per-branch Q-support centres/half-widths. The
// network (w_isv_fc1) and history rotation still consume only the first
// ISV_NETWORK_DIM (23) slots; the tail is broadcast-bus scratchpad only.
isv_signals_pinned: *mut f32, // [ISV_TOTAL_DIM = 31]
isv_signals_dev_ptr: u64,
isv_history_pinned: *mut f32, // [ISV_K * 12 = 48] — history rotates slots [0..11] only
isv_history_dev_ptr: u64,
isv_decay_pinned: *mut f32, // [ISV_DIM = 23] learned decay weights
isv_decay_pinned: *mut f32, // [ISV_NETWORK_DIM = 23] learned decay weights
isv_decay_dev_ptr: u64,
lagged_td_error_pinned: *mut f32, // [1] recursive confidence target
lagged_td_error_dev_ptr: u64,
@@ -2070,31 +2132,32 @@ impl GpuDqnTrainer {
/// Recompute adaptive atom positions from learned spacing parameters.
/// Call once per epoch (positions are slow-moving).
///
/// Per-branch v-range bounds come from the ISV signal bus (slots
/// `V_CENTER_{branch}_INDEX` and `V_HALF_{branch}_INDEX`). The bootstrap
/// written at construction reproduces `[config.v_min, config.v_max]` for
/// every branch before the first `update_eval_v_range` call, so epoch 1
/// behaviour matches the pre-ISV path byte-for-byte.
pub(crate) fn recompute_atom_positions(&self) -> Result<(), MLError> {
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();
// Kernel signature expects `float v_min, float v_max` (f32). Passing
// `self.config.v_{min,max}` (f64) directly routes 8 bytes into 4-byte
// arg slots, corrupting both values and producing unbounded atom
// positions — the observed cause of Q-range explosion to ±144k while
// config bound is ±15. Cast to f32 explicitly to match the kernel ABI.
let v_min = self.config.v_min as f32;
let v_max = self.config.v_max as f32;
let isv_dev_ptr = self.isv_signals_dev_ptr;
for branch in 0..4_usize {
let spacing_ptr = self.ptrs.params_ptr + padded_byte_offset(&param_sizes, 52 + branch);
let out_offset = (branch * na * std::mem::size_of::<f32>()) as u64;
let out_ptr = atom_positions_buf_ptr + out_offset;
let branch_i32 = branch as i32;
unsafe {
self.stream.launch_builder(&self.adaptive_atom_kernel)
.arg(&spacing_ptr)
.arg(&out_ptr)
.arg(&(na as i32))
.arg(&v_min)
.arg(&v_max)
.arg(&branch_i32)
.arg(&isv_dev_ptr)
.launch(LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (256, 1, 1),
@@ -2109,8 +2172,14 @@ impl GpuDqnTrainer {
/// Warm-start C51 atom positions from empirical reward quantiles.
///
/// Called once per epoch after experience collection, before training steps.
/// Writes the same quantile values to all 4 branches — per-branch specialization
/// happens via the existing `step_atom_positions` SGD optimizer.
/// Shared empirical reward quantiles are clamped to each branch's
/// ISV-derived `[centre - half, centre + half]` support before being
/// tiled into `atom_positions_buf`. Prior to the first ISV write, the
/// bootstrap values (centre=0, half=(config.v_max-config.v_min)/2) make
/// the clamp equivalent to the old `config.v_{min,max}` clamp for every
/// branch — epoch 1 behaviour is preserved byte-for-byte. Per-branch
/// specialisation then flows through the existing `step_atom_positions`
/// SGD optimiser.
///
/// The HtoD transfer is ~1us for 4 * num_atoms floats (typically 4 * 51 = 204).
pub(crate) fn warm_start_atom_positions(&mut self, quantiles: &[f32]) -> Result<(), MLError> {
@@ -2119,10 +2188,31 @@ impl GpuDqnTrainer {
return Ok(()); // skip on mismatch
}
// Write same quantiles to all 4 branches
// Fallback bounds used when the ISV bus has not yet been written —
// keep the pre-spec behaviour (full config range, shared across
// branches) as a safety net.
let v_min_f = self.config.v_min;
let v_max_f = self.config.v_max;
let isv_ptr = self.isv_signals_pinned;
let mut host_data = vec![0.0_f32; 4 * na];
for branch in 0..4 {
host_data[branch * na..(branch + 1) * na].copy_from_slice(quantiles);
for branch in 0..4_usize {
let (b_min, b_max) = if isv_ptr.is_null() {
(v_min_f, v_max_f)
} else {
let centre = unsafe { *isv_ptr.add(V_CENTER_DIR_INDEX + 2 * branch) };
let half = unsafe { *isv_ptr.add(V_CENTER_DIR_INDEX + 2 * branch + 1) };
let b_min = centre - half;
let b_max = centre + half;
// Guard against degenerate (half=0) reads; fall back to full
// config range so warm-start never produces zero-range atoms.
if (b_max - b_min).abs() < 1e-6 { (v_min_f, v_max_f) }
else { (b_min, b_max) }
};
let slot = &mut host_data[branch * na..(branch + 1) * na];
for (atom_idx, q) in quantiles.iter().enumerate() {
slot[atom_idx] = q.clamp(b_min, b_max);
}
}
self.stream.memcpy_htod(&host_data, &mut self.atom_positions_buf)
@@ -2306,6 +2396,9 @@ impl Drop for GpuDqnTrainer {
if !self.q_readback_pinned.is_null() {
let _ = unsafe { cudarc::driver::result::free_host(self.q_readback_pinned.cast()) };
}
if !self.per_branch_q_stats_pinned.is_null() {
let _ = unsafe { cudarc::driver::result::free_host(self.per_branch_q_stats_pinned.cast()) };
}
if !self.q_divergence_pinned.is_null() {
let _ = unsafe { cudarc::driver::result::free_host(self.q_divergence_pinned.cast()) };
}
@@ -2932,65 +3025,92 @@ impl GpuDqnTrainer {
Ok(())
}
/// 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, 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.
if !self.eval_ema_initialized {
self.eval_q_mean_ema = 0.0;
self.eval_q_std_ema = q_std.max(0.01);
self.eval_ema_initialized = true;
}
{
let baseline = self.eval_q_std_ema.max(0.001);
let mean_err = (q_mean - self.eval_q_mean_ema).abs();
let std_err = (q_std - self.eval_q_std_ema).abs();
let alpha_mean = (mean_err / (mean_err + baseline)).clamp(0.01, 0.3);
let alpha_std = (std_err / (std_err + baseline)).clamp(0.01, 0.3);
self.eval_q_mean_ema = (1.0 - alpha_mean) * self.eval_q_mean_ema + alpha_mean * q_mean;
self.eval_q_std_ema = (1.0 - alpha_std) * self.eval_q_std_ema + alpha_std * q_std;
}
/// Update eval v_range from observed per-branch Q-value statistics.
/// Called at epoch boundary after `reduce_current_q_stats_per_branch`.
///
/// Per-branch path (ISV v-range unification, spec 2026-04-23):
/// maintains 4 independent `(q_mean_ema, q_std_ema)` pairs and writes 8
/// ISV slots (`V_CENTER_{branch}_INDEX`, `V_HALF_{branch}_INDEX`) — one
/// pair of (centre, half) per branch. Consumers read through the ISV
/// signal bus (`adaptive_atom_positions`, `warm_start_atom_positions`).
///
/// Legacy global `eval_v_range_pinned[2]` is still updated using the
/// direction-branch (branch 0) bounds so any consumer that has not yet
/// migrated to the per-branch ISV bus sees the same value as the
/// "dominant branch" behaviour prior to this change. Full removal of
/// `eval_v_range_pinned` is deferred to the Phase-3 cleanup.
pub fn update_eval_v_range(
&mut self,
per_branch_stats: &PerBranchQValueStats,
per_branch_q_gaps: [f32; 4],
) {
let v_min_f = self.config.v_min;
let v_max_f = self.config.v_max;
let v_range_full = (v_max_f - v_min_f).max(1e-6);
// Prevents atom-grid collapse when the Q-distribution is narrow.
// Spec 2026-04-23: min_half_floor = 0.1 * (v_max - v_min).
let min_half_floor = 0.1_f32 * v_range_full;
let abs_half = 0.5_f32 * v_range_full;
// Keep last_per_branch_q_gaps in sync (per_branch_q_gaps pinned already written in reduce_current_q_stats).
let _ = per_branch_q_gaps; // already on device via per_branch_q_gaps_dev_ptr
for (branch_idx, stats) in per_branch_stats.per_branch.iter().enumerate() {
let q_mean = stats.q_mean;
let q_std = stats.q_variance.max(0.0).sqrt();
let q_gap = per_branch_q_gaps[branch_idx].max(0.0);
// Adaptive-rate EMA with baseline proportional to Q-std.
if !self.eval_ema_initialized[branch_idx] {
self.eval_q_mean_ema[branch_idx] = 0.0;
self.eval_q_std_ema[branch_idx] = q_std.max(0.01);
self.eval_ema_initialized[branch_idx] = true;
} else {
let baseline = self.eval_q_std_ema[branch_idx].max(0.001);
let mean_err = (q_mean - self.eval_q_mean_ema[branch_idx]).abs();
let std_err = (q_std - self.eval_q_std_ema[branch_idx]).abs();
let alpha_mean = (mean_err / (mean_err + baseline)).clamp(0.01, 0.3);
let alpha_std = (std_err / (std_err + baseline)).clamp(0.01, 0.3);
self.eval_q_mean_ema[branch_idx] =
(1.0 - alpha_mean) * self.eval_q_mean_ema[branch_idx] + alpha_mean * q_mean;
self.eval_q_std_ema[branch_idx] =
(1.0 - alpha_std) * self.eval_q_std_ema[branch_idx] + alpha_std * q_std;
}
// Width from max(Q-gap action scale, 3σ band, min floor), clamped
// to the full config range. Floor = 10% of config range — enough
// atom-grid resolution to differentiate actions regardless of the
// observed distribution.
let gap_width = (10.0 * q_gap)
.max(3.0 * self.eval_q_std_ema[branch_idx])
.max(min_half_floor);
let half = gap_width.min(abs_half).max(min_half_floor);
// Centre clamp so the band always fits inside [v_min, v_max].
let center = self.eval_q_mean_ema[branch_idx].clamp(
v_min_f + half,
v_max_f - half,
);
unsafe {
*self.isv_signals_pinned.add(V_CENTER_DIR_INDEX + 2 * branch_idx) = center;
*self.isv_signals_pinned.add(V_CENTER_DIR_INDEX + 2 * branch_idx + 1) = half;
}
// Branch 0 (direction) drives the legacy scalar support — matches
// the dominant-branch Q-scale the single-EMA path used.
if branch_idx == 0 && !self.eval_v_range_pinned.is_null() {
unsafe {
*self.eval_v_range_pinned = center - half;
*self.eval_v_range_pinned.add(1) = center + half;
}
}
}
// GPU kernel: liquid tau RK4/Euler adaptive ODE (zero CPU compute).
let delta_z = self.eval_q_std_ema.max(0.01);
// Uses the direction-branch (b0) Q-std as the global delta_z proxy —
// matches the single-EMA path before per-branch promotion.
let delta_z = self.eval_q_std_ema[0].max(0.01);
if let Err(e) = self.update_liquid_tau(delta_z) {
tracing::warn!("liquid_tau_rk4 kernel failed: {e}");
}
// Width from Q-gap (action differentiation range), not Q-std.
// Floor = max(10 * q_gap, 3 * q_std) — enough atoms to resolve actions.
// With 51 atoms, a width of 10*q_gap gives q_gap/delta_z ≈ 5 atoms of
// resolution between best and worst action.
//
// Clamp the adaptive [v_min, v_max] to the theoretical bounds supplied
// at construction (config.v_{min,max}, derived from reward_scale/gamma).
// Without this clamp the window can drift arbitrarily far from 0: once
// the policy overestimates, eval_q_mean_ema tracks it up → v_max follows
// → C51 atoms span a wider range → even larger TD targets → loss
// explodes geometrically (Fold 1 smoke: final_loss=1.6e16, grad=62B).
//
// Two clamps are needed — one on half-width so the band can't exceed
// the full theoretical range, and one on the centre so the band always
// fits inside [config.v_min, config.v_max]:
let gap_width = (10.0 * q_gap).max(3.0 * self.eval_q_std_ema).max(0.1);
let abs_half = (self.config.v_max - self.config.v_min) * 0.5;
let half = gap_width.min(abs_half);
let center = self.eval_q_mean_ema.clamp(
self.config.v_min + half,
self.config.v_max - half,
);
let v_min = center - half;
let v_max = center + half;
// Pinned device-mapped: CPU writes directly, GPU reads via dev_ptr. No HtoD copy.
unsafe {
*self.eval_v_range_pinned = v_min;
*self.eval_v_range_pinned.add(1) = v_max;
}
let _ = per_branch_q_gaps; // already on device via per_branch_q_gaps_dev_ptr
}
pub fn set_per_sample_support_ptr(&mut self, ptr: u64) {
self.per_sample_support_ptr = ptr;
@@ -6837,12 +6957,37 @@ impl GpuDqnTrainer {
// ── Compile expected Q-value + stats kernels (validation, not in CUDA Graph) ─
let expected_q_kernel = compile_expected_q_kernel(&stream)?;
let q_stats_kernel = compile_q_stats_kernel(&stream)?;
// ISV v-range unification (spec 2026-04-23): per-branch Q-stats
// reducer — emits one 7-tuple per action branch (4 × 7 = 28 floats).
// Drives the per-branch EMA state in `update_eval_v_range` which in
// turn writes the 8 ISV v-range slots.
let q_stats_per_branch_kernel = compile_q_stats_per_branch_kernel(&stream)?;
// Task 2.X "make Full useful" — per-magnitude-bin Q-mean reducer that
// feeds ISV slots [13..16] on the stats cadence.
let q_mag_bin_means_reduce_kernel = compile_q_mag_bin_means_kernel(&stream)?;
let q_dir_bin_means_reduce_kernel = compile_q_dir_bin_means_kernel(&stream)?;
let q_stats_buf = stream.alloc_zeros::<f32>(7)
.map_err(|e| MLError::ModelError(format!("alloc q_stats_f32: {e}")))?;
// Per-branch Q-stats readback — pinned device-mapped, 28 floats (4 × 7).
// GPU writes, CPU reads without sync (one-step lag tolerated; the
// update_eval_v_range path sees the latest value after the next
// replay completes).
let per_branch_q_stats_pinned: *mut f32 = unsafe {
let flags = cudarc::driver::sys::CU_MEMHOSTALLOC_DEVICEMAP;
cudarc::driver::result::malloc_host(28 * std::mem::size_of::<f32>(), flags)
.map_err(|e| MLError::ModelError(format!("pinned per_branch_q_stats alloc: {e}")))?
as *mut f32
};
unsafe { std::ptr::write_bytes(per_branch_q_stats_pinned, 0, 28); }
let per_branch_q_stats_dev_ptr = unsafe {
let mut dp = 0u64;
cudarc::driver::sys::cuMemHostGetDevicePointer_v2(
&mut dp as *mut u64,
per_branch_q_stats_pinned.cast(),
0,
);
dp
};
let atom_stats_buf = stream.alloc_zeros::<f32>(2)
.map_err(|e| MLError::ModelError(format!("alloc atom_stats: {e}")))?;
// Per-block partial sums for the atom-stat deterministic reduction.
@@ -7820,8 +7965,12 @@ impl GpuDqnTrainer {
};
// ── ISV core buffers (pinned device-mapped) ──────────────────────
// Allocation uses `ISV_TOTAL_DIM` so slots [23..31] (per-branch
// Q-support centres and half-widths) have backing storage. The
// network-facing path still consumes only the first `ISV_NETWORK_DIM`
// slots — the scratchpad tail is invisible to `w_isv_fc1`.
let (isv_signals_pinned, isv_signals_dev_ptr) = {
let num_bytes = ISV_DIM * std::mem::size_of::<f32>();
let num_bytes = ISV_TOTAL_DIM * std::mem::size_of::<f32>();
let mut host_ptr: *mut std::ffi::c_void = std::ptr::null_mut();
let mut dev_ptr_out: u64 = 0;
unsafe {
@@ -7830,6 +7979,18 @@ impl GpuDqnTrainer {
let rc2 = cudarc::driver::sys::cuMemHostGetDevicePointer_v2(&mut dev_ptr_out, host_ptr, 0);
assert_eq!(rc2, cudarc::driver::sys::cudaError_enum::CUDA_SUCCESS, "cuMemHostGetDevicePointer for isv_signals");
std::ptr::write_bytes(host_ptr as *mut u8, 0, num_bytes);
// Bootstrap the per-branch v-range slots to the config bounds so
// that any consumer reaching for slots [23..31] before the first
// `update_eval_v_range` call sees identical-to-config behaviour.
// Half-width = (v_max - v_min) / 2, centre = 0 → v_min..v_max.
let v_min_f = config.v_min as f32;
let v_max_f = config.v_max as f32;
let half_bootstrap = 0.5_f32 * (v_max_f - v_min_f);
let sig_ptr = host_ptr as *mut f32;
for b in 0..4usize {
*sig_ptr.add(23 + 2 * b) = 0.0_f32; // v_center
*sig_ptr.add(23 + 2 * b + 1) = half_bootstrap; // v_half
}
}
(host_ptr as *mut f32, dev_ptr_out)
};
@@ -8648,9 +8809,9 @@ impl GpuDqnTrainer {
branch_scales_ptr: 0,
eval_v_range_pinned,
eval_v_range_ptr,
eval_q_mean_ema: 0.0,
eval_q_std_ema: 0.0,
eval_ema_initialized: false,
eval_q_mean_ema: [0.0_f32; 4],
eval_q_std_ema: [0.0_f32; 4],
eval_ema_initialized: [false; 4],
iqn_readiness: 0.0,
iqn_readiness_pinned,
iqn_readiness_dev_ptr,
@@ -8734,6 +8895,9 @@ impl GpuDqnTrainer {
bw_d_glu_gate,
expected_q_kernel,
q_stats_kernel,
q_stats_per_branch_kernel,
per_branch_q_stats_pinned,
per_branch_q_stats_dev_ptr,
q_mag_bin_means_reduce_kernel,
q_mag_means_scratch_pinned,
q_mag_means_scratch_dev_ptr,
@@ -9199,8 +9363,13 @@ impl GpuDqnTrainer {
Ok(())
}
/// Read eval_q_mean_ema.
pub fn eval_q_mean_ema(&self) -> f32 { self.eval_q_mean_ema }
/// Read eval_q_mean_ema as the direction-branch scalar — matches the
/// pre-unification single-EMA behaviour for callers that have not yet
/// migrated to per-branch awareness.
pub fn eval_q_mean_ema(&self) -> f32 { self.eval_q_mean_ema[0] }
/// Read all four per-branch `q_mean_ema` values.
pub fn eval_q_mean_ema_per_branch(&self) -> [f32; 4] { self.eval_q_mean_ema }
/// ISV signals device pointer for adaptive hold enforcement in experience collector.
pub fn isv_signals_dev_ptr(&self) -> u64 { self.isv_signals_dev_ptr }
@@ -9595,8 +9764,13 @@ impl GpuDqnTrainer {
sigma1 / sigma2
}
/// Read eval_q_std_ema.
pub fn eval_q_std_ema(&self) -> f32 { self.eval_q_std_ema }
/// Read eval_q_std_ema as the direction-branch scalar — matches the
/// pre-unification single-EMA behaviour for callers that have not yet
/// migrated to per-branch awareness.
pub fn eval_q_std_ema(&self) -> f32 { self.eval_q_std_ema[0] }
/// Read all four per-branch `q_std_ema` values.
pub fn eval_q_std_ema_per_branch(&self) -> [f32; 4] { self.eval_q_std_ema }
// G6 branch_indep_loss_value + G10 temporal_loss_value removed —
// their buffers were deallocated after V7-gem sub-noise verdict.
@@ -9630,36 +9804,78 @@ impl GpuDqnTrainer {
v
}
/// Set eval_q_mean_ema (for trajectory backtracking restore).
pub fn set_eval_q_mean_ema(&mut self, v: f32) { self.eval_q_mean_ema = v; }
/// Set eval_q_mean_ema for all four branches (trajectory backtracking
/// restore). Scalar setter kept for backwards compatibility — broadcasts
/// the same value to every branch. Callers that snapshot per-branch state
/// should use `set_eval_q_mean_ema_per_branch`.
pub fn set_eval_q_mean_ema(&mut self, v: f32) {
self.eval_q_mean_ema = [v; 4];
self.eval_ema_initialized = [v != 0.0; 4];
}
/// Set eval_q_std_ema (for trajectory backtracking restore).
pub fn set_eval_q_std_ema(&mut self, v: f32) { self.eval_q_std_ema = v; }
/// Set eval_q_std_ema for all four branches (trajectory backtracking
/// restore). Scalar setter kept for backwards compatibility — broadcasts
/// the same value to every branch. Callers that snapshot per-branch state
/// should use `set_eval_q_std_ema_per_branch`.
pub fn set_eval_q_std_ema(&mut self, v: f32) {
self.eval_q_std_ema = [v; 4];
}
/// Per-branch setter for trajectory backtracking restore.
pub fn set_eval_q_mean_ema_per_branch(&mut self, v: [f32; 4]) {
self.eval_q_mean_ema = v;
// Treat any non-zero value as evidence the EMA was observed before.
for b in 0..4 {
self.eval_ema_initialized[b] = v[b] != 0.0 || self.eval_ema_initialized[b];
}
}
/// Per-branch setter for trajectory backtracking restore.
pub fn set_eval_q_std_ema_per_branch(&mut self, v: [f32; 4]) {
self.eval_q_std_ema = v;
}
/// Reset the adaptive C51 v_range state at a fold boundary.
///
/// Clears `eval_q_mean_ema` / `eval_q_std_ema` (the rolling Q-distribution
/// estimate) and resets `eval_v_range_pinned` back to the theoretical
/// `[config.v_min, config.v_max]`. Without this, Fold N+1 inherits Fold N's
/// final tight atom support — if the new fold's Q-distribution doesn't fit
/// the stale range, TD errors explode and training NaN's out (observed:
/// Clears the per-branch `eval_q_mean_ema` / `eval_q_std_ema` arrays and
/// resets both the legacy `eval_v_range_pinned[2]` buffer and the 8
/// per-branch ISV slots (23..30) back to bootstrap values so that fold
/// N+1 starts with atoms spanning the full `[config.v_min, config.v_max]`
/// range on every branch. Without this, Fold N+1 inherits Fold N's final
/// tight atom support — if the new fold's Q-distribution doesn't fit the
/// stale range, TD errors explode and training NaN's out (observed:
/// train-92xbj Fold 1 Epoch 12, grad_norm=inf → NaN loss at step 5).
///
/// The first call to `update_eval_v_range` after this reset will reinitialise
/// the EMA from the new fold's first observed (q_mean, q_std).
/// The first call to `update_eval_v_range` after this reset will
/// reinitialise each branch's EMA from the new fold's first observed
/// `(q_mean, q_std)`.
pub fn reset_eval_v_range_state(&mut self) {
self.eval_q_mean_ema = 0.0;
self.eval_q_std_ema = 0.0;
self.eval_ema_initialized = false;
// Restore the pinned device-mapped range to the theoretical bounds so
// the experience-collection kernels running before the next
// update_eval_v_range call see the wide safe support.
self.eval_q_mean_ema = [0.0_f32; 4];
self.eval_q_std_ema = [0.0_f32; 4];
self.eval_ema_initialized = [false; 4];
// Restore the pinned device-mapped legacy range to the theoretical
// bounds so the experience-collection kernels running before the
// next update_eval_v_range call see the wide safe support.
if !self.eval_v_range_pinned.is_null() {
unsafe {
*self.eval_v_range_pinned = self.config.v_min;
*self.eval_v_range_pinned.add(1) = self.config.v_max;
}
}
// Bootstrap the 8 per-branch ISV slots to centre=0, half=abs_half.
// Matches the construction-time bootstrap — epoch 1 of the next fold
// sees `[config.v_min, config.v_max]` on every branch.
if !self.isv_signals_pinned.is_null() {
let v_min_f = self.config.v_min;
let v_max_f = self.config.v_max;
let half_bootstrap = 0.5_f32 * (v_max_f - v_min_f);
unsafe {
for b in 0..4usize {
*self.isv_signals_pinned.add(V_CENTER_DIR_INDEX + 2 * b) = 0.0_f32;
*self.isv_signals_pinned.add(V_CENTER_DIR_INDEX + 2 * b + 1) = half_bootstrap;
}
}
}
}
/// Set per_branch_q_gap_ema — HtoD for trajectory backtracking restore.
@@ -11198,6 +11414,93 @@ impl GpuDqnTrainer {
Ok(result)
}
/// Reduce current `q_out_buf` to per-branch Q-stats (4 × 7 floats).
///
/// Launches `q_stats_per_branch_reduce` — slices the flat Q-output by
/// branch and computes `[avg_max_q, q_min, q_max, q_mean, q_var,
/// atom_entropy(=0), atom_utilization(=0)]` per branch. Drives the
/// `update_eval_v_range` per-branch EMA state and the 8 ISV v-range
/// slots (spec 2026-04-23).
///
/// Non-destructive to `q_stats_buf` / `q_readback_pinned` — uses a
/// separate pinned readback. Intended to be called alongside
/// `reduce_current_q_stats` at epoch boundary; both share the same
/// `q_out_buf` population so there is no extra forward pass.
pub fn reduce_current_q_stats_per_branch(&mut self)
-> Result<PerBranchQValueStats, MLError>
{
let _eg = EventTrackingGuard::new(self.stream.context());
let total_actions = self.total_actions() as i32;
let n = self.config.batch_size as i32;
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 b0_off: i32 = 0;
let b1_off: i32 = b0;
let b2_off: i32 = b0 + b1;
let b3_off: i32 = b0 + b1 + b2;
let q_out_buf_ptr = self.q_out_buf.raw_ptr();
let out_ptr = self.per_branch_q_stats_dev_ptr;
unsafe {
self.stream
.launch_builder(&self.q_stats_per_branch_kernel)
.arg(&q_out_buf_ptr)
.arg(&out_ptr)
.arg(&n)
.arg(&total_actions)
.arg(&b0_off).arg(&b0)
.arg(&b1_off).arg(&b1)
.arg(&b2_off).arg(&b2)
.arg(&b3_off).arg(&b3)
.launch(LaunchConfig {
grid_dim: (1, 1, 1),
block_dim: (1, 1, 1),
shared_mem_bytes: 0,
})
.map_err(|e| MLError::ModelError(format!("q_stats_per_branch_reduce: {e}")))?;
}
// Synchronous read from pinned host buffer. The kernel is launched on
// the same stream as the C51 graph, so a single stream synchronize
// guarantees the pinned bytes reflect the current launch — matches the
// legacy `compute_q_stats` cadence.
unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); }
let mut host = [0.0_f32; 28];
unsafe {
std::ptr::copy_nonoverlapping(
self.per_branch_q_stats_pinned, host.as_mut_ptr(), 28,
);
}
let mut per_branch = [QValueStatsResult {
avg_max_q: 0.0,
q_min: 0.0,
q_max: 0.0,
q_mean: 0.0,
q_variance: 0.0,
atom_entropy: 0.0,
atom_utilization: 0.0,
}; 4];
for b in 0..4_usize {
let base = b * 7;
per_branch[b] = QValueStatsResult {
avg_max_q: host[base] as f64,
q_min: host[base + 1],
q_max: host[base + 2],
q_mean: host[base + 3],
q_variance: host[base + 4],
atom_entropy: host[base + 5],
atom_utilization: host[base + 6],
};
}
Ok(PerBranchQValueStats { per_branch })
}
/// Launch the xLSTM mLSTM step: reads q_stats_buf[7] + per_branch_q_gaps[4] (pinned),
/// updates persistent matrix memory C[8,8] and normalizer n[8],
/// writes context[8] output. Runs in a single thread (grid=1, block=1).
@@ -14887,6 +15190,22 @@ fn compile_q_stats_kernel(
.map_err(|e| MLError::ModelError(format!("q_stats_reduce load: {e}")))
}
/// Load the per-branch Q-value statistics reducer.
///
/// Emits one 7-tuple per branch (direction, magnitude, order, urgency) for the
/// ISV v-range per-branch EMA update path. See `q_stats_kernel.cu` for the
/// `q_stats_per_branch_reduce` body. Runs alongside the legacy global
/// `q_stats_reduce` — no semantic conflict, both single-threaded deterministic.
fn compile_q_stats_per_branch_kernel(
stream: &Arc<CudaStream>,
) -> Result<CudaFunction, MLError> {
let context = stream.context();
let module = context.load_cubin(Q_STATS_CUBIN.to_vec())
.map_err(|e| MLError::ModelError(format!("q_stats cubin load: {e}")))?;
module.load_function("q_stats_per_branch_reduce")
.map_err(|e| MLError::ModelError(format!("q_stats_per_branch_reduce load: {e}")))
}
/// Load the per-magnitude-bin Q-mean reducer from the q_stats cubin.
///
/// Task 2.X "make Full useful": this kernel computes batch-mean Q-values

View File

@@ -61,6 +61,108 @@ extern "C" __global__ void q_stats_reduce(
out[6] = (atom_stats != NULL) ? atom_stats[1] * inv_n / (float)num_atoms : 0.0f; /* utilization fraction [0,1] */
}
/**
* Per-branch Q-value statistics reduction.
*
* Parallel path to `q_stats_reduce` above. Emits one `[7]` stats tuple per
* action branch (direction, magnitude, order, urgency) — 28 floats total,
* branch-major.
*
* Layout of `out[4 * 7]`:
* out[b*7 + 0] = avg_max_q_b (mean of per-sample max over branch_b's actions)
* out[b*7 + 1] = q_min_b (min Q across all samples × branch_b's actions)
* out[b*7 + 2] = q_max_b
* out[b*7 + 3] = q_mean_b (mean Q across samples × branch_b's actions)
* out[b*7 + 4] = q_var_b (variance)
* out[b*7 + 5] = atom_entropy_b (NaN — atom_stats kernel is global, not per-branch yet)
* out[b*7 + 6] = atom_util_b (NaN — same)
*
* Used by `update_eval_v_range` to maintain per-branch EMA state for the
* ISV v-range broadcast (slots 23..30). Branches have genuinely different
* Q-scales — direction Q ≈ ±50, magnitude Q ≈ ±5 — and a single shared
* range forces atoms into a bad quantization for at least one branch.
*
* Branch slicing: given `q_values[N, total_actions]`, branch b's actions
* occupy indices `[branch_offsets[b], branch_offsets[b] + branch_sizes[b])`.
* Caller passes `branch_offsets[4]` and `branch_sizes[4]`.
*
* Grid=(1,1,1) Block=(1,1,1) — single-threaded determinism matching
* q_stats_reduce. Cost: ~N * total_actions FP32 loads split across 4 branches
* plus a second pass for variance; single-digit microseconds at typical B.
*/
extern "C" __global__ void q_stats_per_branch_reduce(
const float* __restrict__ q_values, /* [N, total_actions] */
float* __restrict__ out, /* [4 * 7] branch-major */
int N,
int total_actions,
int b0_off, int b0_size,
int b1_off, int b1_size,
int b2_off, int b2_size,
int b3_off, int b3_size)
{
if (threadIdx.x != 0) return;
const int branch_offsets[4] = { b0_off, b1_off, b2_off, b3_off };
const int branch_sizes[4] = { b0_size, b1_size, b2_size, b3_size };
for (int b = 0; b < 4; b++) {
const int off = branch_offsets[b];
const int sz = branch_sizes[b];
/* Bounds check — invalid branch produces zeros, not garbage. */
if (sz <= 0 || off < 0 || off + sz > total_actions || N <= 0) {
for (int k = 0; k < 7; k++) out[b * 7 + k] = 0.0f;
continue;
}
float branch_min = 1e30f;
float branch_max = -1e30f;
float branch_sum = 0.0f;
float sum_max_q = 0.0f;
const int total = N * sz;
/* First pass: min/max/mean/avg_max. */
for (int i = 0; i < N; i++) {
const float* row = q_values + (long long)i * total_actions + off;
float row_max = -1e30f;
for (int a = 0; a < sz; a++) {
float v = row[a];
if (v < branch_min) branch_min = v;
if (v > branch_max) branch_max = v;
if (v > row_max) row_max = v;
branch_sum += v;
}
sum_max_q += row_max;
}
const float mean = (total > 0) ? branch_sum / (float)total : 0.0f;
/* Second pass: variance. */
float var_sum = 0.0f;
for (int i = 0; i < N; i++) {
const float* row = q_values + (long long)i * total_actions + off;
for (int a = 0; a < sz; a++) {
float d = row[a] - mean;
var_sum += d * d;
}
}
const float variance = (total > 0) ? var_sum / (float)total : 0.0f;
const float avg_max = (N > 0) ? sum_max_q / (float)N : 0.0f;
out[b * 7 + 0] = avg_max;
out[b * 7 + 1] = branch_min;
out[b * 7 + 2] = branch_max;
out[b * 7 + 3] = mean;
out[b * 7 + 4] = variance;
/* Atom entropy / util remain global-only for now. Per-branch atom_stats
* would require splitting the compute_expected_q accumulator the same
* way; deferred to a follow-up since ISV v-range only needs the first
* five stats. Leave 0 — consumers check for it. */
out[b * 7 + 5] = 0.0f;
out[b * 7 + 6] = 0.0f;
}
}
/**
* Per-branch magnitude-Q mean reducer: computes batch-mean Q-values for each
* of the 3 magnitude bins separately (Quarter=0, Half=1, Full=2), plus the

View File

@@ -2940,9 +2940,26 @@ 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, per_branch_q_gaps: [f32; 4]) {
self.trainer.update_eval_v_range(q_mean, q_std, q_gap, per_branch_q_gaps);
/// Update eval v_range from observed per-branch Q-value statistics.
/// Spec 2026-04-23: writes 8 ISV slots (centre + half-width per branch)
/// and, for backward compatibility, the legacy scalar `eval_v_range_pinned`
/// from the direction-branch bounds.
pub(crate) fn update_eval_v_range(
&mut self,
per_branch_stats: &crate::cuda_pipeline::gpu_dqn_trainer::PerBranchQValueStats,
per_branch_q_gaps: [f32; 4],
) {
self.trainer.update_eval_v_range(per_branch_stats, per_branch_q_gaps);
}
/// Reduce per-branch Q-stats from the current `q_out_buf` — 4 × 7 floats
/// covering `[avg_max_q, q_min, q_max, q_mean, q_var, atom_entropy(=0),
/// atom_utilization(=0)]` per branch. Feeds `update_eval_v_range` for the
/// ISV v-range per-branch EMA bus.
pub(crate) fn reduce_current_q_stats_per_branch(&mut self)
-> Result<crate::cuda_pipeline::gpu_dqn_trainer::PerBranchQValueStats, crate::MLError>
{
self.trainer.reduce_current_q_stats_per_branch()
}
/// 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

@@ -1722,12 +1722,12 @@ impl DQNTrainer {
self.epoch_atom_entropy = stats.atom_entropy;
self.epoch_atom_utilization = stats.atom_utilization;
// 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);
// Per-branch Q-gaps computed in reduce_current_q_stats (48B DtoH readback).
// Per-branch Q-stats for the ISV v-range EMA updater.
// Per-branch Q-gaps pinned in reduce_current_q_stats (48B DtoH readback).
let per_branch_q_gaps = fused.get_per_branch_q_gaps();
fused.update_eval_v_range(stats.q_mean, q_std, q_gap, per_branch_q_gaps);
if let Ok(per_branch_stats) = fused.reduce_current_q_stats_per_branch() {
fused.update_eval_v_range(&per_branch_stats, per_branch_q_gaps);
}
}
}
} // end if false (Q-stats disabled per-step)
@@ -1756,10 +1756,18 @@ impl DQNTrainer {
// the only thing that updates the buffer — so it stays at
// zero and the downstream health composition believes
// q_gap=0 even when raw epoch q_gap peaks above 1.0.
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);
let per_branch_q_gaps = fused.get_per_branch_q_gaps();
fused.update_eval_v_range(stats.q_mean, q_std, q_gap, per_branch_q_gaps);
match fused.reduce_current_q_stats_per_branch() {
Ok(per_branch_stats) => {
fused.update_eval_v_range(&per_branch_stats, per_branch_q_gaps);
}
Err(e) => {
tracing::warn!(
"reduce_current_q_stats_per_branch failed at epoch-boundary: {e} — \
ISV v-range update skipped this step"
);
}
}
}
}

View File

@@ -0,0 +1,280 @@
# ISV-Unified Q-Value Support Range (per-branch)
> **Problem statement:** Three Q-support consumers (atom grid, quantile warm-start, loss projection) hold independent opinions about the valid Q range for each of the 4 action branches. When they disagree, atoms become a bad quantizer: probability mass collapses onto 1-2 boundary atoms, the policy can't differentiate actions, and reward signal doesn't flow into useful Q-gradient information.
## Observation driving this design
- `train-fpxnw` (current run, Q-clamp fix applied): atom utilization **11-15%** across 57 fold-0 epochs. Probability concentrates on 1-2 of 51 atoms.
- Peak fold-0 Sharpe trending down as correctness fixes land: **79 → 68 → 48**. Each fix removed an accidental exploration crutch (wide Q range = wide atom grid = noisy Q-differentiation = effective exploration).
- Hypothesis under test: the Sharpe regression is not caused by our fixes being wrong — it's caused by *remaining* range inconsistencies between consumers. Atoms sit in one range, loss projects onto another, quantile-init fills a third. Each fix exposed one layer of the mismatch but left the others.
## Design principles
1. **One signal, one source of truth** for Q-support range per branch.
2. **Adaptive, temporal, signal-driven** — per the `adaptive_not_tuned` rule. EMA-tracked from observed Q-stats.
3. **Safety-clamped** to static `config.v_{min,max}` — hard rail, not tuning dial.
4. **Per-branch** — direction, magnitude, order, urgency have genuinely different Q scales; forcing a shared range recreates the exact mismatch we're fixing, just moved one level deeper.
5. **ISV bus broadcast** — established pattern for cross-kernel signals; consumers read from the bus they already receive.
## Current state (before this spec)
| Component | Q-range source | Type |
|-----------|----------------|------|
| `atom_positions_buf[4, num_atoms]` | `adaptive_atom_positions` kernel with `config.v_{min,max}` | **Static, per-config** |
| `per_sample_support_buf[N, 3]` | `eval_v_range_pinned`, tiled per epoch | **Adaptive, single-EMA, shared across branches** |
| `warm_start_atom_positions` quantile clamp | `config.v_{min,max}` (commit `d38a8cf99`) | **Static, per-config** |
| `q_stats_reduce` kernel output | Global min/max/mean/var across all branches | **Not per-branch** |
The four components disagree along two axes:
- **Static vs adaptive**: atoms + quantile = static; projection = adaptive.
- **Per-branch vs pooled**: nothing is per-branch today; magnitude's Q=±5 and direction's Q=±50 share atoms and support bounds.
## Proposed state
| Component | New Q-range source | Type |
|-----------|-------------------|------|
| `atom_positions_buf[4, num_atoms]` | ISV per-branch center/half-width | **Adaptive, per-branch** |
| `per_sample_support_buf` — deprecated | ISV per-branch center/half-width (read per-branch inside loss kernel) | **Adaptive, per-branch** |
| `warm_start_atom_positions` quantile clamp | ISV per-branch center/half-width (read per-branch on host) | **Adaptive, per-branch** |
| `q_stats_reduce` kernel output | Per-branch `q_mean, q_std, q_gap` → 4 sets of stats | **Per-branch** |
One source of truth: the ISV bus. Broadcast written by one path, read by three consumers.
## ISV slot allocation
Current ISV layout: 23 slots (slot 22 = `sharpe_ema`). Extend to 31:
| Slot | Name | Semantics |
|------|------|-----------|
| 23 | `v_center_dir` | EMA Q-center, direction branch |
| 24 | `v_half_dir` | EMA Q-half-width, direction branch |
| 25 | `v_center_mag` | EMA Q-center, magnitude branch |
| 26 | `v_half_mag` | EMA Q-half-width, magnitude branch |
| 27 | `v_center_ord` | EMA Q-center, order branch |
| 28 | `v_half_ord` | EMA Q-half-width, order branch |
| 29 | `v_center_urg` | EMA Q-center, urgency branch |
| 30 | `v_half_urg` | EMA Q-half-width, urgency branch |
**Why center + half-width, not min + max:** the adaptive EMA computes these two quantities naturally. Center is mean-Q EMA; half-width is `max(k·std, k·gap, min_floor)`. Consumers compute `v_min = center - half_width`, `v_max = center + half_width` themselves. Storing (center, half-width) keeps the evolution independent — center can shift while width stays constant.
**Bootstrap values** (pre any Q observations, set at construction + fold-boundary reset):
- `v_center_* = 0.0`
- `v_half_* = (config.v_max - config.v_min) / 2.0`
At epoch 1 these produce `[config.v_min, config.v_max]` for every branch — byte-identical to current behavior.
**Safety clamp** on every ISV write:
- `center = center_raw.clamp(config.v_min + min_half, config.v_max - min_half)`
- `half = half_raw.clamp(min_half_floor, (config.v_max - config.v_min) / 2.0)`
where `min_half_floor = 0.1 * (config.v_max - config.v_min)` prevents atom-grid collapse. Guarantees atoms always span a usable fraction of the config range.
## Prerequisite: per-branch Q-stats kernel
The ISV write path needs per-branch `q_mean`, `q_std`, `q_gap` statistics. Currently `q_stats_reduce` produces one global tuple across all 108 actions. This must be reworked first.
**Option A (extend existing kernel):** modify `q_stats_reduce` to emit `[4 × 7]` instead of `[7]`. Output format per branch: `[avg_max_q, q_min, q_max, q_mean, q_var, atom_entropy, atom_utilization]`.
**Option B (separate per-branch kernel):** add `q_stats_per_branch_reduce` alongside the existing one. Keeps the legacy kernel for any global consumers, adds a parallel path.
**Choice: Option A.** The global stats are already computed implicitly in the per-branch reduction (sum-across-branches of per-branch sums). Callers that need the global tuple can aggregate from the per-branch output. Avoids kernel duplication.
Branch offsets into `q_out_buf[N, total_actions]`:
- direction: `[0, b0)`
- magnitude: `[b0, b0+b1)`
- order: `[b0+b1, b0+b1+b2)`
- urgency: `[b0+b1+b2, b0+b1+b2+b3)`
Pass `b0, b1, b2, b3` to the kernel; iterate per sample, per branch, reduce.
**Readback shape change:** host side `QValueStatsResult` becomes `[QValueStatsResult; 4]` (or a new `PerBranchQValueStats`). Propagate the type change through `reduce_current_q_stats`, `update_eval_v_range`, consumers in `training_loop.rs`.
## Write path
Only `GpuDqnTrainer::update_eval_v_range` writes the ISV slots. It already runs at epoch boundary after `compute_q_stats`. New code flow:
```rust
pub fn update_eval_v_range(&mut self, per_branch_stats: &[QValueStatsResult; 4]) {
for (branch_idx, stats) in per_branch_stats.iter().enumerate() {
// Adaptive-rate EMA per branch
let baseline = self.eval_q_std_ema[branch_idx].max(0.001);
let mean_err = (stats.q_mean - self.eval_q_mean_ema[branch_idx]).abs();
let std_err = (stats.q_std() - self.eval_q_std_ema[branch_idx]).abs();
let alpha_mean = (mean_err / (mean_err + baseline)).clamp(0.01, 0.3);
let alpha_std = (std_err / (std_err + baseline)).clamp(0.01, 0.3);
self.eval_q_mean_ema[branch_idx] =
(1.0 - alpha_mean) * self.eval_q_mean_ema[branch_idx] + alpha_mean * stats.q_mean;
self.eval_q_std_ema[branch_idx] =
(1.0 - alpha_std) * self.eval_q_std_ema[branch_idx] + alpha_std * stats.q_std();
// Derived center + half-width
let gap_width = (10.0 * stats.q_gap()).max(3.0 * self.eval_q_std_ema[branch_idx]).max(min_half_floor);
let abs_half = (self.config.v_max - self.config.v_min) as f32 * 0.5;
let half = gap_width.min(abs_half).max(min_half_floor);
let center = self.eval_q_mean_ema[branch_idx].clamp(
(self.config.v_min as f32) + half,
(self.config.v_max as f32) - half,
);
// Pinned device-mapped: writes visible to next kernel replay
unsafe {
*self.isv_signals_pinned.add(23 + 2 * branch_idx) = center;
*self.isv_signals_pinned.add(24 + 2 * branch_idx) = half;
}
}
}
```
Per-branch EMA state replaces the current single `eval_q_mean_ema, eval_q_std_ema` scalars — they become `[f32; 4]`.
## Read-path changes
**1. `adaptive_atom_positions` kernel (CUDA)**
Current signature:
```cpp
extern "C" __global__ void adaptive_atom_positions(
const float* __restrict__ spacing_raw, // [NA]
float* __restrict__ positions_out, // [NA]
int num_atoms,
float v_min, // per-branch in call site loop, from config
float v_max
);
```
New signature:
```cpp
extern "C" __global__ void adaptive_atom_positions(
const float* __restrict__ spacing_raw, // [NA]
float* __restrict__ positions_out, // [NA]
int num_atoms,
int branch_idx, // 0..4
const float* __restrict__ isv_signals // read slots 23+2*branch_idx, 24+2*branch_idx
);
```
Kernel body:
```cpp
float center = isv_signals[23 + 2 * branch_idx];
float half = isv_signals[24 + 2 * branch_idx];
float v_min = center - half;
float v_max = center + half;
// (rest unchanged — softmax + cumsum scaled by v_min, v_max)
```
Rust caller `recompute_atom_positions` changes: instead of passing config.v_min/v_max as f32, pass `branch_idx` and `isv_signals_dev_ptr`. The f64→f32 ABI fix from commit `768cc7d82` becomes moot (no more scalar f32 arg for bounds).
**2. `warm_start_atom_positions` (Rust host side)**
Current: reads reward quantiles, clamps with `hyperparams.v_{min,max}`, writes to atom_positions_buf.
New: reads per-branch bounds from `isv_signals_pinned` (pinned, zero-cost host read), clamps each of the 4 branch quantile arrays independently:
```rust
for branch_idx in 0..4 {
let center = unsafe { *isv_pinned.add(23 + 2 * branch_idx) };
let half = unsafe { *isv_pinned.add(24 + 2 * branch_idx) };
let v_min = center - half;
let v_max = center + half;
for (atom_idx, q) in branch_quantiles[branch_idx].iter_mut().enumerate() {
*q = q.clamp(v_min, v_max);
}
}
```
Currently `compute_reward_quantiles` returns a single `Vec<f32>` of size num_atoms shared across branches. Either:
- Keep shared quantiles (the reward distribution is the same regardless of branch), just clamp per-branch.
- Extract per-branch quantiles (requires per-branch reward attribution). **Out of scope** — shared quantiles with per-branch clamp is sufficient for the fix.
**3. C51 loss + MSE loss kernels**
Currently read `per_sample_support_buf[N, 3]` which is tiled from the single `eval_v_range_pinned`. For per-branch atom positions, the projection support in the loss kernel must also be per-branch.
Two options:
- **Keep `per_sample_support_buf` but make it `[N, 4, 3]`** (per-sample, per-branch). Tiling cost: 4× but still one HtoD per epoch. Minimal kernel change — index `per_sample_support[sample, branch_idx, {min,max,dz}]`.
- **Remove `per_sample_support_buf`, read from ISV directly inside the loss kernel.** Smaller kernel arg surface, one fewer buffer. But every Bellman-projection iteration does an extra ISV read per atom.
**Choice: option 1** (per-sample × per-branch support tile). Simpler kernel diff, caches well, matches the atom_positions layout. Tiling happens on epoch boundary from ISV slots.
**4. `update_per_sample_support` (Rust) — rename & restructure**
`update_per_sample_support(v_min, v_max)` becomes `update_per_sample_support_per_branch()` — reads all 8 ISV slots, tiles into `[N, 4, 3]`.
## Cross-fold semantics
`reset_eval_v_range_state()` extends to:
- Per-branch EMA state reset to 0.0 (center) and config-default (std).
- ISV slots 23..30 reset to bootstrap values.
Fold N+1 starts with atoms spanning full config range per branch. After epoch 1's per-branch Q-stats, adapts to new fold's distribution. No cross-fold bleed.
## Feedback-loop safety
Same safety rails as before, now per-branch:
1. **Hard clamp** at `config.v_{min,max}` — per-branch half-width cannot exceed `(v_max - v_min) / 2`.
2. **Minimum half-width floor** — prevents atom-grid collapse when Q-distribution is narrow.
3. **Anti-LR controller**, **Q-divergence EMA**, **gradient clipping** — unchanged, all still apply.
A per-branch runaway (e.g., magnitude Q grows without bound) is observable via `isv[24+2k]` hitting the config-clamp ceiling. Diagnosable directly from HEALTH_DIAG, which already reads ISV slots.
## Test plan
1. **Unit test: per-branch stats extraction** — feed known `q_out_buf` into `q_stats_reduce`, assert output matches per-branch hand-computed stats.
2. **Smoke test: ISV write-read round-trip** — run `update_eval_v_range` with synthetic per-branch stats, read back ISV slots 23..30, assert matches expected center + half.
3. **Smoke test: atom range tracks ISV** — inject ISV slots with known values, run `recompute_atom_positions`, assert atom_positions span `[center - half, center + half]` for each branch.
4. **Smoke test: quantile clamp uses per-branch bounds** — inject ISV values that differ per branch, verify `warm_start_atom_positions` clamps each branch's atoms to that branch's range.
5. **Smoke test: fold-boundary reset** — call `reset_for_fold`, assert all 8 ISV slots at bootstrap values.
6. **Smoke test: atom utilization** — target metric. Expect ≥40% after fix (up from 11-15%).
7. **L40S 60-epoch run** — compare Sharpe trajectory, peak Sharpe, atom util vs `train-fpxnw` baseline.
## Implementation phases
**Phase 0 — prerequisite: per-branch Q-stats kernel (~1 day)**
- Modify `q_stats_reduce` to emit `[4 × 7]` stats.
- Propagate `PerBranchQValueStats` type through Rust side.
- Unit-test equivalence: aggregated per-branch stats == old global stats.
- Commit as a data-structure refactor, no semantic behavior change.
**Phase 1 — ISV plumbing (~half day, no functional change at epoch 1)**
- Grep-update `ISV_DIM` 23 → 31.
- Add per-branch EMA state `[f32; 4]` for mean + std.
- Bootstrap values set at construction.
- `update_eval_v_range` writes ISV slots 23..30.
- `reset_eval_v_range_state` resets slots.
- Consumers still reach for `config.v_{min,max}` — no semantic change.
- Commit: "ISV plumbing, zero behavior change."
**Phase 2 — consumer migration (~half day, breaks behavior)**
- `adaptive_atom_positions` kernel: new signature, reads ISV.
- `recompute_atom_positions` Rust caller: passes branch_idx + isv_dev_ptr.
- `warm_start_atom_positions`: per-branch clamp from ISV pinned.
- `update_per_sample_support` rename + restructure to `[N, 4, 3]` tiling.
- C51/MSE loss kernels: index per-branch into support tile.
- Commit: "v-range unified via ISV, per-branch atom grids."
**Phase 3 — validation + cleanup (~half day + 2.5h GPU)**
- Run smoke tests 1-6.
- L40S run, compare to `train-fpxnw`.
- Delete `eval_v_range_pinned` if redundant (probably is).
Total: ~2.5 days engineering + 1 validation run.
## Success criteria
All four must hold:
1. **Correctness:** ISV slots 23..30 written from one place, read by three consumers; no other path reaches for `config.v_{min,max}` as a runtime range except through ISV.
2. **Atom utilization improved:** ≥40% on fold 0 (up from 11-15% currently).
3. **Sharpe not regressed:** peak fold-0 Sharpe ≥ current +48.19.
4. **Fold-1 stability holds:** Q-range at fold 1 epoch 1 is within ISV-derived bounds (no ±333k or ±144k excursions as in pre-fix runs).
## Out of scope
- Removing `eval_v_range_pinned` (deferred to Phase 3 cleanup).
- Per-branch reward quantile extraction (shared quantile with per-branch clamp is sufficient).
- Cross-fold inherit (fold-reset is the proven default; revisit only if convergence speed becomes a bottleneck).
- Additional ISV signals beyond the 8 slots for v-range.
## Revision history
- **v1 (earlier today)**: proposed 2 ISV slots (global center + half-width). Revised after recognizing that direction-branch Q-scale (~±50) differs from magnitude-branch (~±5) — a shared range recreates the exact mismatch the fix targets.
- **v2 (this version)**: 8 slots (4 centers + 4 half-widths). Per-branch Q-stats kernel added as explicit prerequisite.