feat(dqn-v2): Plan 4 Task 3 E.3 — IQN fixed-τ multi-quantile heads (5/25/50/75/95)

Replaces random τ ∈ U(0,1) sampling with `FIXED_TAUS = [0.05, 0.25, 0.50,
0.75, 0.95]`. Kernel-side `IQN_NUM_QUANTILES` macro 32 → 5; `GpuIqnConfig::
default().num_quantiles` 32 → 5. Construction-time τ broadcast (option B2)
populates `online_taus` / `target_taus` / `cos_features` once via
`clone_htod`; both online and target IQN forwards plus the CVaR cold path
read this static buffer. The Philox-driven `iqn_sample_taus_kernel` deleted
along with its only Rust consumer (in `compute_cvar_scales`); the
`rng_step` Philox seed counter also gone. Action ranking in the IQN
inference kernel switched from mean-over-quantiles to MEDIAN
(`q_acc[a] = q_val` only when `t == IQN_MEDIAN_INDEX = 2`); the off-median
positions feed four new ISV diagnostic slots.

Four new ISV slots tail-appended:
  IQN_Q_P05_EMA_INDEX = 99   (mean |Q| at τ=0.05, EMA)
  IQN_Q_P25_EMA_INDEX = 100  (τ=0.25)
  IQN_Q_P75_EMA_INDEX = 101  (τ=0.75)
  IQN_Q_P95_EMA_INDEX = 102  (τ=0.95)

Median (τ=0.50) intentionally skipped — already in greedy-Q diagnostic.
Fingerprint pair shifted 97→103, 98→104; ISV_TOTAL_DIM 99→105.
Layout fingerprint: 0x3e21acecd922e540 → 0x5789155b683ab59c.

New kernel `iqn_quantile_ema_kernel.cu` (4-block × 256-thread shmem-reduce,
no atomicAdd) reads `save_q_online [TBA, B*Q]` and EMA-updates the four
slots. Launched from `training_loop.rs` per-step alongside
`launch_h_s2_rms_ema`. StateResetRegistry extended with 4 FoldReset
entries (cold-start 0.0).

Hyperparam plumbing: `hyperparams.num_quantiles` and
`DQNConfig::iqn_num_quantiles` pinned to `FIXED_TAUS.len()` at the
`GpuIqnConfig` construction site in `fused_training.rs::new` and
`trainer/constructor.rs`. Legacy fields stay for compat; production /
hyperopt configs (dqn-production.toml, DQNHyperparameters defaults)
aligned to 5.

Adam state for IQN params auto-resizes via `m_buf`/`v_buf` sizing
through `total_params + cublas_pad`. **Checkpoint break** — IQN head
parameter shapes change with `num_quantiles`; new fingerprint hash
fails-fast at constructor load on pre-Task-3 checkpoints.

Smoke tests:
- New `iqn_multi_quantile_heads_produce_monotonic_estimates` (1.23s on
  RTX 3050 Ti): asserts ISV[99..103) finite + non-zero + spread > 1e-6
  after 1 epoch — PASS (Q_p05=0.0187 Q_p25=0.0200 Q_p75=0.0193
  Q_p95=0.0190).
- `multi_fold_convergence` (606.50s, 3 folds × 5 epochs): all 3 fold
  checkpoints written; per-fold best train Sharpe -8.17 / 74.24 / 63.44
  at epochs 2 / 4 / 2 (mean 43.17 vs 2c.3c.6 baseline mean 23.43 — folds
  1+2 substantially up, fold 0 down -16 points; absolute-mean comfortably
  above the plan's 3.8 floor). No NaN/Inf, no panic.

cargo check clean at 11 warnings (baseline preserved); cargo build
compiles 61 cubins (was 60; +iqn_quantile_ema, -nothing — the old
sample_taus kernel was inside iqn_dual_head_kernel.cu, not a separate
cubin file).

Files touched: 14 modified (`iqn_dual_head_kernel.cu`, `iqn_cvar_kernel.cu`,
`gpu_iqn_head.rs`, `gpu_dqn_trainer.rs`, `build.rs`, `state_reset_registry
.rs`, `training_loop.rs`, `constructor.rs`, `fused_training.rs`,
`config.rs`, `dqn-production.toml`, `smoke_tests/mod.rs`,
`docs/dqn-wire-up-audit.md`, `dqn-production.toml`) + 2 new (`iqn_quantile
_ema_kernel.cu`, `smoke_tests/iqn_quantile_monotonicity.rs`).

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-04-25 17:10:30 +02:00
parent eaa65ac241
commit 005ed3a4f9
15 changed files with 739 additions and 130 deletions

View File

@@ -91,7 +91,9 @@ her_ratio = 0.2
cql_alpha = 1.0
curiosity_weight = 0.0 # disabled: dense micro-reward provides per-bar signal, curiosity adds noise
iqn_lambda = 0.25
num_quantiles = 32
# Plan 4 Task 3 (E.3): IQN is pinned to FIXED_TAUS.len() (= 5) at the kernel
# layer; this slot stays for compat but no longer controls runtime behaviour.
num_quantiles = 5
spectral_norm_sigma_max = 1.5
spectral_decoupling_lambda = 0.01
# With per-component clipping removed, raw gradient norms are ~4000.

View File

@@ -137,6 +137,12 @@ fn main() {
// in this commit — 2c.3c.6 wires the consumer in `mag_concat_qdir`'s
// adaptive-scale path.
"h_s2_rms_ema_kernel.cu",
// Plan 4 Task 3 (E.3): IQN multi-quantile diagnostic EMAs into
// ISV[99..103) (Q_p05/Q_p25/Q_p75/Q_p95; median τ=0.50 is the existing
// greedy-Q diagnostic and not duplicated). 4-block kernel, one block
// per off-median quantile, shmem-reduce over (B × TBA) of
// `save_q_online`. No atomicAdd. Producer-only — diagnostic only.
"iqn_quantile_ema_kernel.cu",
];
// ALL kernels get common header (BF16 types + wrappers)

View File

@@ -132,6 +132,15 @@ pub(crate) static GRN_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/g
/// 2c.3c.6 wires the consumer in `mag_concat_qdir`'s adaptive-scale path.
static H_S2_RMS_EMA_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/h_s2_rms_ema_kernel.cubin"));
/// Plan 4 Task 3 (E.3): IQN multi-quantile diagnostic EMAs into ISV[99..103).
/// 4-block (256 threads/block, shmem-reduce, no atomicAdd) kernel launched
/// alongside `h_s2_rms_ema_update` from `training_loop.rs`. Reads the IQN
/// online forward's `save_q_online [TBA, B*Q]` surface and EMAs the mean
/// |Q| at τ ∈ {0.05, 0.25, 0.75, 0.95} into ISV[99..103) — the median
/// (τ=0.50) is intentionally skipped (already in the greedy-Q diagnostic).
/// Diagnostic only — producer-only in this commit.
static IQN_QUANTILE_EMA_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/iqn_quantile_ema_kernel.cubin"));
/// Mamba2 temporal scan configuration.
const MAMBA2_HISTORY_K: usize = 8; // Rolling history length
const MAMBA2_STATE_DIM: usize = 16; // SSM state dimension
@@ -298,11 +307,22 @@ const ISV_NETWORK_DIM: usize = 23;
/// [96] H_S2_RMS_EMA_INDEX — EMA of RMS(save_h_s2) per batch. Producer-only
/// in this commit (no consumers). Wired in 2c.3c.6 to drive the
/// adaptive scale of `mag_concat_qdir`'s residual stack.
/// [97] ISV_LAYOUT_FINGERPRINT_LO_INDEX — low 32 bits of layout fingerprint (shifted from 94).
/// [98] ISV_LAYOUT_FINGERPRINT_HI_INDEX — high 32 bits of layout fingerprint (shifted from 95).
/// [99] IQN_Q_P05_EMA_INDEX — Plan 4 Task 3 (E.3) — EMA of the IQN τ=0.05
/// quantile averaged across (B × tba). Producer:
/// `iqn_quantile_ema_update` GPU kernel reading
/// `gpu_iqn.save_q_online` after each training step. Diagnostic only
/// (no consumer kernel reads slot [99..103) in this commit).
/// [100] IQN_Q_P25_EMA_INDEX — same producer/contract; τ=0.25.
/// [101] IQN_Q_P75_EMA_INDEX — same producer/contract; τ=0.75.
/// [102] IQN_Q_P95_EMA_INDEX — same producer/contract; τ=0.95.
/// Note: the median (τ=0.50) lives in the existing C51 / IQL greedy-Q
/// diagnostic — not duplicated here. Only the 4 off-median quantiles
/// are new.
/// [103] ISV_LAYOUT_FINGERPRINT_LO_INDEX — low 32 bits of layout fingerprint (shifted from 97).
/// [104] ISV_LAYOUT_FINGERPRINT_HI_INDEX — high 32 bits of layout fingerprint (shifted from 98).
/// Written by the constructor; checked at checkpoint load. Fail-fast only — no migration
/// path exists. See spec §4.A.2 and `LAYOUT_FINGERPRINT_CURRENT` for structural-hash rationale.
const ISV_TOTAL_DIM: usize = 99;
const ISV_TOTAL_DIM: usize = 105;
/// 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).
@@ -570,7 +590,28 @@ pub const TARGET_DRIFT_DIR_EMA_INDEX: usize = 93;
/// invariant input regardless of the GRN trunk's drift).
pub const H_S2_RMS_EMA_INDEX: usize = 96;
/// ISV slot [97] — low 32 bits of the u64 layout fingerprint (stored as raw f32 bits).
/// ISV slots [99..103) — Plan 4 Task 3 (E.3): IQN multi-quantile diagnostic
/// EMAs for the four off-median fixed quantiles. The median (τ=0.50) is
/// already exposed via the existing greedy-Q diagnostics (no duplication).
///
/// Producer: `iqn_quantile_ema_update` GPU kernel (single-block, 256
/// threads, shmem-reduce, no atomicAdd). Reads
/// `GpuIqnHead::save_q_online [TBA, B*Q]` (col-major; TBA = sum of branch
/// sizes, Q = `IQN_NUM_QUANTILES = 5`) and EMA-updates each of the four
/// non-median quantile slices (τ ∈ {0.05, 0.25, 0.75, 0.95} → indices
/// {0, 1, 3, 4} per `FIXED_TAUS`) into ISV[99..103) at α matching the
/// other Plan 4 producers (cold-path-cadence, one launch per training
/// step alongside `h_s2_rms_ema_update`).
///
/// Cold-start: 0.0 (no IQN forwards have happened yet). FoldReset → 0.0.
/// Producer-only — diagnostic surface for HEALTH_DIAG / risk monitoring;
/// no consumer kernel reads these slots in this commit.
pub const IQN_Q_P05_EMA_INDEX: usize = 99;
pub const IQN_Q_P25_EMA_INDEX: usize = 100;
pub const IQN_Q_P75_EMA_INDEX: usize = 101;
pub const IQN_Q_P95_EMA_INDEX: usize = 102;
/// ISV slot [103] — low 32 bits of the u64 layout fingerprint (stored as raw f32 bits).
///
/// Design note: this is NOT a version number. There is no ordered version space.
/// The fingerprint is a structural hash that automatically changes whenever any
@@ -584,14 +625,16 @@ pub const H_S2_RMS_EMA_INDEX: usize = 96;
/// Shifted 61→69 in Plan 3 Task 1, 69→73 in Plan 3 Task 3, 73→76 in
/// Plan 3 Task 4 B.4, 76→80 in Plan 3 Task 7 C.3, 80→85 in Plan 3 Task 8 B.3,
/// 85→90 in Plan 4 Task 5 Mode A, 90→94 in Plan 4 follow-up (target-drift
/// EMA replacing legacy CPU-DtoH path), then 94→97 in Plan 4 Task 2c.3c.5
/// (H_S2_RMS_EMA producer slot).
pub const ISV_LAYOUT_FINGERPRINT_LO_INDEX: usize = 97;
/// ISV slot [98] — high 32 bits of the u64 layout fingerprint (stored as raw f32 bits).
/// EMA replacing legacy CPU-DtoH path), 94→97 in Plan 4 Task 2c.3c.5
/// (H_S2_RMS_EMA producer slot), then 97→103 in Plan 4 Task 3 E.3 (IQN
/// fixed-τ multi-quantile head, +4 diagnostic slots).
pub const ISV_LAYOUT_FINGERPRINT_LO_INDEX: usize = 103;
/// ISV slot [104] — high 32 bits of the u64 layout fingerprint (stored as raw f32 bits).
/// Shifted 62→70 in Plan 3 Task 1, 70→74 in Plan 3 Task 3, 74→77 in Plan 3 Task 4 B.4,
/// 77→81 in Plan 3 Task 7 C.3, 81→86 in Plan 3 Task 8 B.3, 86→91 in Plan 4 Task 5,
/// 91→95 in Plan 4 follow-up, then 95→98 in Plan 4 Task 2c.3c.5.
pub const ISV_LAYOUT_FINGERPRINT_HI_INDEX: usize = 98;
/// 91→95 in Plan 4 follow-up, 95→98 in Plan 4 Task 2c.3c.5, then 98→104 in Plan 4
/// Task 3 E.3.
pub const ISV_LAYOUT_FINGERPRINT_HI_INDEX: usize = 104;
/// Canonical alias for the fingerprint slot (the low half).
pub const ISV_LAYOUT_FINGERPRINT_INDEX: usize = ISV_LAYOUT_FINGERPRINT_LO_INDEX;
@@ -670,8 +713,9 @@ const fn layout_fingerprint_seed() -> &'static [u8] {
VSN_MAG_EMA=87;VSN_DIR_EMA=88;MAMBA2_RETENTION_EMA=89;\
TARGET_DRIFT_MAG_EMA=92;TARGET_DRIFT_DIR_EMA=93;\
H_S2_RMS_EMA=96;\
ISV_LAYOUT_FINGERPRINT_LO=97;ISV_LAYOUT_FINGERPRINT_HI=98;\
ISV_TOTAL_DIM=99;\
IQN_Q_P05_EMA=99;IQN_Q_P25_EMA=100;IQN_Q_P75_EMA=101;IQN_Q_P95_EMA=102;\
ISV_LAYOUT_FINGERPRINT_LO=103;ISV_LAYOUT_FINGERPRINT_HI=104;\
ISV_TOTAL_DIM=105;\
PARAM_W_A_H_S1=0;PARAM_B_A_H_S1=1;PARAM_W_B_H_S1=2;PARAM_B_B_H_S1=3;\
PARAM_W_RESIDUAL_H_S1=4;PARAM_GAMMA_H_S1=5;PARAM_BETA_H_S1=6;\
PARAM_W_A_H_S2=7;PARAM_B_A_H_S2=8;PARAM_W_B_H_S2=9;PARAM_B_B_H_S2=10;\
@@ -2268,6 +2312,18 @@ pub struct GpuDqnTrainer {
/// Loaded from `h_s2_rms_ema_kernel.cubin`.
h_s2_rms_ema_kernel: CudaFunction,
// ── Plan 4 Task 3 (E.3): IQN multi-quantile diagnostic EMA producer ──
/// 4-block kernel (one per off-median quantile τ ∈ {0.05, 0.25, 0.75, 0.95})
/// computing the mean |Q| over (B × TBA) at each fixed-τ position from
/// `GpuIqnHead::save_q_online [TBA, B*Q]` and EMA-updating
/// ISV[IQN_Q_P05_EMA_INDEX..IQN_Q_P95_EMA_INDEX] (slots 99..103). The
/// median (τ=0.50) is intentionally skipped — already exposed via the
/// existing greedy-Q diagnostic. Cold-path-cadence (one launch per
/// training step, alongside `h_s2_rms_ema_update`). Producer-only —
/// diagnostic surface for HEALTH_DIAG / risk monitoring.
/// Loaded from `iqn_quantile_ema_kernel.cubin`.
iqn_quantile_ema_kernel: CudaFunction,
// ── Q-mean drift correction (Component 8) ──────────────────────────
/// Phase 1 kernel: computes global mean of Q-values into q_mean_scratch (pinned).
q_mean_reduce_kernel: CudaFunction,
@@ -6768,6 +6824,72 @@ impl GpuDqnTrainer {
Ok(())
}
/// Plan 4 Task 3 (E.3): launch `iqn_quantile_ema_update` (4-block,
/// 256-thread shmem-reduction kernel). Reads the IQN online forward's
/// per-quantile Q surface (`save_q_online [TBA, B*Q]`, populated by
/// `GpuIqnHead::execute_training_pipeline` after each training step)
/// and EMAs the mean |Q| at each off-median fixed-τ position
/// {0.05, 0.25, 0.75, 0.95} into ISV[IQN_Q_P05_EMA_INDEX..=IQN_Q_P95_EMA_INDEX]
/// (slots 99..103). Median (τ=0.50, idx=2) is intentionally skipped —
/// already in the greedy-Q diagnostic.
///
/// `save_q_online_ptr` / `tba` / `num_quantiles` are sourced from the
/// `GpuIqnHead` accessors so the trainer doesn't have to duplicate the
/// IQN config layout. Producer-only — diagnostic surface for
/// HEALTH_DIAG / risk monitoring; no consumer kernel reads slots
/// [99..103) in this commit.
///
/// No-op fast-return when `save_q_online_ptr == 0` (IQN not active in
/// this trainer config — the Rust caller already gates on `gpu_iqn`'s
/// `Option`, so a NULL here means an explicit programmer choice and
/// returning silently keeps the caller's gating contract intact).
pub fn launch_iqn_quantile_ema(
&self,
save_q_online_ptr: u64,
tba: usize,
num_quantiles: usize,
ema_alpha: f32,
) -> Result<(), MLError> {
if save_q_online_ptr == 0 {
return Ok(());
}
debug_assert!(self.isv_signals_dev_ptr != 0,
"launch_iqn_quantile_ema: isv_signals_dev_ptr must be allocated by constructor");
debug_assert!(num_quantiles == 5,
"launch_iqn_quantile_ema: Plan 4 Task 3 (E.3) pins IQN_NUM_QUANTILES = 5; \
got num_quantiles={num_quantiles}");
let b_i = self.config.batch_size as i32;
let q_i = num_quantiles as i32;
let tba_i = tba as i32;
let isv_dev_ptr = self.isv_signals_dev_ptr;
let p05 = IQN_Q_P05_EMA_INDEX as i32;
let p25 = IQN_Q_P25_EMA_INDEX as i32;
let p75 = IQN_Q_P75_EMA_INDEX as i32;
let p95 = IQN_Q_P95_EMA_INDEX as i32;
const BLOCK_DIM: u32 = 256;
let smem_bytes = BLOCK_DIM * std::mem::size_of::<f32>() as u32;
unsafe {
self.stream.launch_builder(&self.iqn_quantile_ema_kernel)
.arg(&save_q_online_ptr)
.arg(&b_i)
.arg(&q_i)
.arg(&tba_i)
.arg(&isv_dev_ptr)
.arg(&p05)
.arg(&p25)
.arg(&p75)
.arg(&p95)
.arg(&ema_alpha)
.launch(LaunchConfig {
grid_dim: (4, 1, 1),
block_dim: (BLOCK_DIM, 1, 1),
shared_mem_bytes: smem_bytes,
})
.map_err(|e| MLError::ModelError(format!("iqn_quantile_ema_update: {e}")))?;
}
Ok(())
}
/// Apply spectral normalization to all 10 weight matrices (trunk + 8 heads).
///
/// One step of power iteration per call (standard practice — single step
@@ -7449,6 +7571,18 @@ impl GpuDqnTrainer {
.map_err(|e| MLError::ModelError(format!("h_s2_rms_ema_update load: {e}")))?
};
// Plan 4 Task 3 (E.3): load iqn_quantile_ema kernel (cold-path, per-step).
// 4-block kernel — one block per off-median fixed quantile in
// FIXED_TAUS = [0.05, 0.25, 0.50, 0.75, 0.95]. Producer for
// ISV[IQN_Q_P05_EMA_INDEX..IQN_Q_P95_EMA_INDEX] (slots 99..103).
// Diagnostic only — no consumer kernel reads these slots in this commit.
let iqn_quantile_ema_kernel = {
let module = stream.context().load_cubin(IQN_QUANTILE_EMA_CUBIN.to_vec())
.map_err(|e| MLError::ModelError(format!("iqn_quantile_ema cubin load: {e}")))?;
module.load_function("iqn_quantile_ema_update")
.map_err(|e| MLError::ModelError(format!("iqn_quantile_ema_update load: {e}")))?
};
// Plan 2 Task 1 C.1: allocate branch-action offset/size device buffers for q_quantile_reduce.
// Branch layout in q_out_buf: [dir | mag | ord | urg] row-major per sample.
// dir occupies actions [0..b0), mag [b0..b0+b1), ord [b0+b1..b0+b1+b2),
@@ -8810,6 +8944,17 @@ impl GpuDqnTrainer {
* Producer-only in this commit; 2c.3c.6 wires the consumer
* in `mag_concat_qdir`'s adaptive-scale path. FoldReset → 1.0. */
*sig_ptr.add(H_S2_RMS_EMA_INDEX) = 1.0_f32;
/* Plan 4 Task 3 (E.3): cold-start the four IQN multi-quantile
* diagnostic EMAs at 0.0. The producer kernel
* `iqn_quantile_ema_update` runs each step alongside
* `h_s2_rms_ema_update`; first fire EMAs the measured
* per-quantile mean toward 0.0 with the same α as the other
* Plan 4 producers. Producer-only — no consumer kernel reads
* slots [99..103) in this commit. FoldReset → 0.0. */
*sig_ptr.add(IQN_Q_P05_EMA_INDEX) = 0.0_f32;
*sig_ptr.add(IQN_Q_P25_EMA_INDEX) = 0.0_f32;
*sig_ptr.add(IQN_Q_P75_EMA_INDEX) = 0.0_f32;
*sig_ptr.add(IQN_Q_P95_EMA_INDEX) = 0.0_f32;
// Plan 2 C.1 Q-quantile bootstrap (2026-04-24):
// q_p05 = v_min, q_p95 = v_max — matches the cold-start atom range
// so the first update_eval_v_range call reads meaningful bounds.
@@ -9826,6 +9971,7 @@ impl GpuDqnTrainer {
popart_count,
reward_component_ema_kernel,
h_s2_rms_ema_kernel,
iqn_quantile_ema_kernel,
q_mean_reduce_kernel,
q_mean_subtract_kernel,
q_mean_scratch_pinned,

View File

@@ -41,6 +41,29 @@ use super::shared_cublas_handle::PerStreamCublasHandles;
// Configuration
// ---------------------------------------------------------------------------
/// Plan 4 Task 3 (E.3): fixed-τ multi-quantile heads.
///
/// Replaces the legacy QR-DQN midpoint sampling τ_i = (2i+1)/(2N) with five
/// well-known risk-aware quantile points sorted ascending: 5th / 25th /
/// 50th (median) / 75th / 95th. The kernel's `IQN_NUM_QUANTILES` macro is
/// fixed to 5 in lockstep so this Rust-side table fully describes the
/// per-step τ axis. The constructor's `clone_htod` broadcasts these 5
/// values across all batch rows (`[B, 5]`) once at init; both online and
/// target IQN forwards plus the CVaR cold path read this static buffer
/// directly — no per-step kernel sampling required.
///
/// Median (τ=0.50) lives at index 2; off-median positions feed the four
/// new ISV diagnostic slots (ISV[99..103) Q_p05/p25/p75/p95 EMAs). The
/// kernel's inference path (`iqn_forward_kernel`) latches the median Q
/// per action via `if (t == IQN_MEDIAN_INDEX) q_acc[a] = q_val;`, so any
/// re-ordering of this table requires updating that constant (= 2) too.
pub const FIXED_TAUS: [f32; 5] = [0.05, 0.25, 0.50, 0.75, 0.95];
/// Index of the median quantile (τ=0.50) in `FIXED_TAUS`. Mirrors
/// `IQN_MEDIAN_INDEX` in `iqn_forward_kernel`. Used by IQR and any other
/// consumer that needs the median position; the spec § E.3 ordering pins
/// this at 2 (sorted ascending {0.05, 0.25, 0.50, 0.75, 0.95}).
pub const FIXED_TAUS_MEDIAN_INDEX: usize = 2;
/// Configuration for the GPU IQN dual-head trainer.
#[derive(Debug, Clone)]
pub struct GpuIqnConfig {
@@ -48,7 +71,10 @@ pub struct GpuIqnConfig {
pub hidden_dim: usize,
/// Cosine embedding dimension (default 64).
pub embed_dim: usize,
/// Number of quantile samples per forward pass (default 32).
/// Number of quantile samples per forward pass.
/// Plan 4 Task 3 (E.3): fixed at 5 — see `FIXED_TAUS`. Must match the
/// kernel-side `IQN_NUM_QUANTILES` define; mismatch is a build error
/// because the kernel sizes register arrays from the macro.
pub num_quantiles: usize,
/// Huber threshold for quantile loss (default 1.0).
pub kappa: f32,
@@ -87,7 +113,9 @@ impl Default for GpuIqnConfig {
Self {
hidden_dim: 256,
embed_dim: 64,
num_quantiles: 32,
// Plan 4 Task 3 (E.3): fixed-τ {0.05, 0.25, 0.50, 0.75, 0.95}.
// Was 32 (random U(0,1) sampling, QR-DQN-equivalent midpoints).
num_quantiles: FIXED_TAUS.len(),
kappa: 1.0,
state_dim: 72,
shared_h1: 256,
@@ -183,7 +211,10 @@ pub struct GpuIqnHead {
forward_kernel: CudaFunction,
ema_kernel: CudaFunction,
decode_actions_kernel: CudaFunction,
sample_taus_kernel: CudaFunction,
// Plan 4 Task 3 (E.3): `sample_taus_kernel` deleted. Fixed-τ broadcast
// at construction replaces per-step Philox sampling. The kernel itself
// is removed from `iqn_dual_head_kernel.cu` in lockstep — no other
// consumer existed.
// ── cuBLAS element-wise cubin kernels (forward + backward) ────
relu_fwd_kernel: CudaFunction,
@@ -321,8 +352,9 @@ pub struct GpuIqnHead {
t_dev_ptr: u64,
tau_pinned: *mut f32, // pinned+device-mapped host pointer to tau scalar [1]
tau_dev_ptr: u64, // device pointer to same physical page — passed to EMA kernel
/// Monotonic step counter for Philox PRNG seeding (τ sampling).
rng_step: u32,
// Plan 4 Task 3 (E.3): `rng_step` (Philox seed counter for per-step τ
// sampling) deleted alongside the sampler kernel. Fixed-τ broadcast
// at construction has no per-step state.
total_params: usize,
}
@@ -459,14 +491,28 @@ impl GpuIqnHead {
let bias_grad_partials_buf = alloc_f32(&stream, bias_grad_num_blocks * h, "iqn_bias_grad_partials")?;
// ── Per-step buffers ────────────────────────────────────────────
// Fixed τ midpoints (QR-DQN style): τ_i = (2i + 1) / (2N)
// Plan 4 Task 3 (E.3): fixed-τ {0.05, 0.25, 0.50, 0.75, 0.95}.
// Pre-edit used the QR-DQN midpoints τ_i = (2i+1)/(2N) over 32
// quantiles, which is the same uniform-grid mean estimator but
// randomized via Philox at run time. Fixing the τ values lets the
// 5 quantile heads be interpretable / consistent across batches and
// enables the median-Q action-ranking switch in `iqn_forward_kernel`.
// Both online and target IQN read this same static buffer — the
// off-median heads are diagnostic only (ISV[99..103) producers) and
// must agree on the fixed τ between online and target so the
// quantile-Huber regression learns a stable distributional surface.
let online_taus;
let target_taus;
let cos_features;
{
let fixed: Vec<f32> = (0..n)
.map(|i| (2.0 * (i as f32) + 1.0) / (2.0 * n as f32))
.collect();
assert_eq!(
n, FIXED_TAUS.len(),
"GpuIqnConfig::num_quantiles ({n}) must equal FIXED_TAUS.len() \
({}) — kernel-side IQN_NUM_QUANTILES is sized to FIXED_TAUS \
in lockstep (Plan 4 Task 3, E.3)",
FIXED_TAUS.len(),
);
let fixed: &[f32] = &FIXED_TAUS;
// Precompute cosine features in col-major [D, N]:
// cos_features[d + q*D] = cos(π·(d+1)·τ_q)
@@ -484,7 +530,7 @@ impl GpuIqnHead {
let mut tiled = Vec::with_capacity(b * n);
for _ in 0..b {
tiled.extend_from_slice(&fixed);
tiled.extend_from_slice(fixed);
}
online_taus = stream.clone_htod(&tiled).map_err(|e| {
MLError::ModelError(format!("IQN htod online_taus ({} f32): {e}", b * n))
@@ -583,7 +629,6 @@ impl GpuIqnHead {
forward_kernel: kernels.fwd_only,
ema_kernel: kernels.ema,
decode_actions_kernel: kernels.decode_act,
sample_taus_kernel: kernels.sample_taus,
relu_fwd_kernel: kernels.relu_fwd,
relu_bwd_kernel: kernels.relu_bwd,
hadamard_sigmoid_kernel: kernels.hadamard_sigmoid,
@@ -644,7 +689,6 @@ impl GpuIqnHead {
t_dev_ptr,
tau_pinned,
tau_dev_ptr,
rng_step: 0,
total_params,
})
}
@@ -1579,9 +1623,13 @@ impl GpuIqnHead {
/// Compute IQR per action from the last forward pass's quantile estimates.
///
/// IQR = quantile_75 - quantile_25 for each action, averaged over the last
/// sample in the training batch. Uses fixed midpoint taus so idx_25 and idx_75
/// map to the 25th and 75th percentile quantiles.
/// IQR = quantile_75 quantile_25 for each action, drawn from the last
/// sample in the training batch. Plan 4 Task 3 (E.3): with `FIXED_TAUS`
/// = `[0.05, 0.25, 0.50, 0.75, 0.95]`, the 25th and 75th percentiles
/// live at slot indices 1 and 3 respectively (= `n_q / 4` and
/// `(3*n_q) / 4` for `n_q = 5` — the integer-truncation arithmetic
/// happens to land on the right slots, kept for parity with the prior
/// 32-quantile QR-DQN-midpoint behaviour).
///
/// Call after `train_iqn_step_gpu()` which populates `save_q_online`.
/// This is a cold-path operation (epoch boundary) -- host-side readback is acceptable.
@@ -1589,8 +1637,15 @@ impl GpuIqnHead {
let n_q = self.config.num_quantiles;
let n_a = self.total_branch_actions();
// Plan 4 Task 3 (E.3): with FIXED_TAUS = [0.05, 0.25, 0.50, 0.75, 0.95],
// idx_25 = 1 and idx_75 = 3. Computed as `n_q/4` / `3*n_q/4` so the
// formula remains valid if FIXED_TAUS were ever extended to a longer
// grid that still placed Q25/Q75 at the canonical positions; with
// n_q=5 these evaluate to 1 and 3 by integer truncation.
let idx_25 = n_q / 4;
let idx_75 = (3 * n_q) / 4;
debug_assert!(idx_25 < n_q && idx_75 < n_q && idx_25 < idx_75,
"IQR indices out of order: idx_25={idx_25}, idx_75={idx_75}, n_q={n_q}");
// save_q_online is [TBA, B*Q] col-major.
// For last sample (b=B-1), quantile q: column = (B-1)*Q + q, row = action
@@ -1643,6 +1698,29 @@ impl GpuIqnHead {
self.iqr_buf.raw_ptr()
}
/// Raw device pointer to the per-quantile online Q surface
/// `save_q_online [TBA, B*Q]` (col-major) populated by
/// `execute_training_pipeline`. Plan 4 Task 3 (E.3) consumer:
/// `iqn_quantile_ema_update` GPU kernel, which reduces (B × TBA)
/// per off-median quantile into ISV[99..103). The buffer's lifetime
/// matches `GpuIqnHead`'s; the pointer is stable across training steps.
pub fn save_q_online_ptr(&self) -> u64 {
self.save_q_online.raw_ptr()
}
/// Total branch actions (TBA) — sum across all 4 branches. Public
/// accessor for Plan 4 Task 3 (E.3) callers that need to size the
/// `iqn_quantile_ema_update` reduction without re-reading the config
/// (which is private to this struct).
pub fn tba(&self) -> usize {
self.total_branch_actions()
}
/// Number of quantiles (= `FIXED_TAUS.len()` = 5 under Plan 4 Task 3).
pub fn num_quantiles(&self) -> usize {
self.config.num_quantiles
}
/// Compute CVaR-based position scaling from IQN quantiles.
///
/// Runs the IQN forward-only kernel on `h_s2` (trunk activation),
@@ -1664,36 +1742,17 @@ impl GpuIqnHead {
let hidden_dim_i32 = self.config.hidden_dim as i32;
let embed_dim_i32 = self.config.embed_dim as i32;
// 1. Sample τ values for inference
let total_taus = (b * n) as i32;
let tau_blocks = (b * n + 255) / 256;
let tau_config = LaunchConfig {
grid_dim: (tau_blocks as u32, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
// Plan 4 Task 3 (E.3): per-step Philox τ sampling deleted. The
// constructor already populated `online_taus` (and `target_taus`)
// with FIXED_TAUS broadcast across all B rows; CVaR's quantile
// estimates are therefore drawn from the same fixed quantile
// levels {0.05, 0.25, 0.50, 0.75, 0.95} as the training pipeline,
// making the CVaR estimate stable across steps and consistent
// with the spec's interpretable-quantile contract.
let b0_i32 = self.config.branch_0_size as i32;
let b1_i32 = self.config.branch_1_size as i32;
let b2_i32 = self.config.branch_2_size as i32;
let b3_i32 = self.config.branch_3_size as i32;
let rng_step = self.rng_step as u32;
self.rng_step += 1;
unsafe {
self.stream
.launch_builder(&self.sample_taus_kernel)
.arg(&mut self.online_taus)
.arg(&rng_step)
.arg(&total_taus)
.arg(&shared_h1_i32)
.arg(&hidden_dim_i32)
.arg(&embed_dim_i32)
.arg(&b0_i32)
.arg(&b1_i32)
.arg(&b2_i32)
.arg(&b3_i32)
.launch(tau_config)
.map_err(|e| MLError::ModelError(format!("IQN CVaR sample_taus: {e}")))?;
}
// 2. Run IQN forward-only kernel (256 threads per sample)
let fwd_inf_shmem = (256 / 32) * 4;
@@ -1929,6 +1988,8 @@ fn create_iqn_gemm_desc(
struct IqnKernels {
// Plan 4 Task 2c.3b: trunk_fwd deleted (orphaned kernel — see kernel
// file for the deletion comment).
// Plan 4 Task 3 (E.3): sample_taus deleted (fixed-τ broadcast replaces
// per-step Philox sampling — see kernel file for the deletion comment).
loss_reduce: CudaFunction,
gnorm_p1: CudaFunction,
gnorm_p2: CudaFunction,
@@ -1936,7 +1997,6 @@ struct IqnKernels {
fwd_only: CudaFunction,
ema: CudaFunction,
decode_act: CudaFunction,
sample_taus: CudaFunction,
relu_fwd: CudaFunction,
relu_bwd: CudaFunction,
hadamard_sigmoid: CudaFunction,
@@ -1989,7 +2049,8 @@ fn load_iqn_kernels(
let fwd_only = load("iqn_forward_kernel")?;
let ema = load("iqn_ema_kernel")?;
let decode_act = load("iqn_decode_actions_kernel")?;
let sample_taus = load("iqn_sample_taus_kernel")?;
// Plan 4 Task 3 (E.3): iqn_sample_taus_kernel removed from cubin — no
// load. Fixed-τ broadcast at construction replaces per-step sampling.
// New element-wise kernels for cuBLAS pipeline
let relu_fwd = load("iqn_relu_fwd")?;
@@ -2004,10 +2065,13 @@ fn load_iqn_kernels(
let d_h_s2_reduce = load("iqn_d_h_s2_reduce")?;
let cos_tile = load("iqn_cos_tile")?;
info!("GpuIqnHead: 19 kernels loaded (8 retained + 11 element-wise; trunk_fwd deleted in Task 2c.3b)");
info!(
"GpuIqnHead: 18 kernels loaded (7 retained + 11 element-wise; \
trunk_fwd deleted in Task 2c.3b, sample_taus deleted in Task 3 E.3)"
);
Ok(IqnKernels {
loss_reduce, gnorm_p1, gnorm_p2, adam, fwd_only,
ema, decode_act, sample_taus,
ema, decode_act,
relu_fwd, relu_bwd, hadamard_sigmoid, hadamard_sigmoid_bwd,
quantile_huber_loss, bias_add, bias_grad_reduce_p1, bias_grad_reduce_p2,
h_s2_tile, d_h_s2_reduce, cos_tile,

View File

@@ -5,6 +5,21 @@
* One thread per sample: insertion-sort quantile values in registers,
* compute CVaR as mean of lowest alpha_count values.
*
* Plan 4 Task 3 (E.3): N_TAU now 5 (was 32). With fixed
* FIXED_TAUS = {0.05, 0.25, 0.50, 0.75, 0.95}, the discrete CVaR cases
* are:
* ALPHA = 0.05 → alpha_count = floor(0.05 × 5) = 0 → clamped to 1
* → returns the single τ=0.05 quantile (worst-5% estimator)
* ALPHA = 0.25 → alpha_count = floor(0.25 × 5) = 1
* → again returns the τ=0.05 quantile (single worst-25%
* sample given the 5-point grid; mathematically
* consistent with discrete quantile aggregation)
* ALPHA = 0.5+ → alpha_count grows toward N_TAU; mean of bottom slots
*
* The `sorted[8]` register array stays oversized for the new grid (max
* alpha_count = 5) — no resizing needed; same kernel handles the smaller
* τ axis without modification.
*
* Scale: [0.25, 1.0]. CVaR >= 0 -> full size. CVaR < 0 -> reduce.
*
* Launch config: grid=(ceil(B/256), 1, 1), block=(256, 1, 1).
@@ -39,8 +54,14 @@ extern "C" __global__ void iqn_cvar_kernel(
}
cvar_sum = min_val;
} else {
// General path: find alpha_count smallest values
float sorted[8]; // max alpha_count for alpha=0.25, N_TAU=32
// General path: find alpha_count smallest values.
// sorted[8] sizing: with Plan 4 Task 3 (E.3) N_TAU=5 the maximum
// possible alpha_count is N_TAU itself = 5 (when ALPHA=1.0). The
// legacy 32-quantile kernel sized this for ALPHA≤0.25 → 8 slots;
// 8 still covers the new grid completely. Kept oversized rather
// than tightened to 5 so future N_TAU bumps below 8 don't require
// a kernel-source edit.
float sorted[8];
int sorted_len = 0;
for (int t = 0; t < N_TAU; t++) {
float v = q_values[i * N_TAU * TBA + t * TBA + exposure_idx];

View File

@@ -51,8 +51,12 @@
#ifndef IQN_EMBED_DIM
#define IQN_EMBED_DIM 64
#endif
/* Plan 4 Task 3 (E.3): IQN_NUM_QUANTILES 32 → 5. Fixed-τ multi-quantile
* heads {0.05, 0.25, 0.50, 0.75, 0.95} replace random U(0,1) sampling.
* Per-step quantile axis shrinks 32→5 (6.4× smaller); the 5 fixed
* quantile heads are interpretable / consistent across batches. */
#ifndef IQN_NUM_QUANTILES
#define IQN_NUM_QUANTILES 32
#define IQN_NUM_QUANTILES 5
#endif
#ifndef IQN_KAPPA
#define IQN_KAPPA 1.0f
@@ -901,10 +905,16 @@ void iqn_adam_kernel(
* Reads bf16 h_s2 from DQN trunk, converts to f32 at boundary.
* All params and taus are f32. Output is f32.
*
* Computes expected Q-values per action per branch by averaging over
* quantile samples. Used for action selection blending with C51.
* Plan 4 Task 3 (E.3): with `IQN_NUM_QUANTILES = 5` and FIXED_TAUS
* = {0.05, 0.25, 0.50, 0.75, 0.95} (sorted ascending, populated from the
* Rust constructor's HtoD upload), the median quantile lives at index 2
* (τ = 0.50). Action ranking now uses the median Q rather than the mean
* over all quantiles — the median is the risk-neutral central estimate
* that the C51 mode/argmax-of-expected aggregation already provides at
* its head, and using it here keeps the IQN inference path consistent
* with the spec's 5-quantile decomposition.
*
* Outputs: expected_q [B, tba] — mean Q per action.
* Outputs: expected_q [B, tba] — median (τ=0.50) Q per action.
*/
extern "C" __global__
void iqn_forward_kernel(
@@ -953,8 +963,13 @@ void iqn_forward_kernel(
for (int h = tid; h < hidden_dim; h += IQN_BLOCK_SIZE)
h_dist[h / IQN_BLOCK_SIZE] = (float)my_h_s2_bf16[h];
/* Accumulate Q-values across quantiles (for mean).
* Max tba = 3+3+3+3 = 12 — fits easily in registers. */
/* Plan 4 Task 3 (E.3): hold the MEDIAN (τ=0.50, index 2) Q per action.
* The fixed-τ broadcast guarantees τ_2 = 0.50 across all samples; the
* mean over the asymmetric set {0.05, 0.25, 0.50, 0.75, 0.95} would
* mix tail-risk into the action-ranking estimate. Median is the risk-
* neutral central tendency consistent with the spec's 5-quantile
* decomposition (other 4 quantiles surface as ISV diagnostics).
* Max tba = 4+3+3+3 = 13 — fits easily in registers. */
float q_acc[16]; /* sized >= max tba */
for (int a = 0; a < tba; a++)
q_acc[a] = 0.0f;
@@ -979,7 +994,20 @@ void iqn_forward_kernel(
for (int h = tid; h < hidden_dim; h += IQN_BLOCK_SIZE)
comb_dist[h / IQN_BLOCK_SIZE] = h_dist[h / IQN_BLOCK_SIZE] * embed_dist[h / IQN_BLOCK_SIZE];
/* Branch outputs */
/* Plan 4 Task 3 (E.3): only the median quantile (τ=0.50, index 2)
* feeds expected_q. Off-median quantiles still execute their GEMMs
* (the trunk is fixed; per-quantile cost is the W_b·combined dot
* product) so that the IQN's per-step Q surface is fully sampled
* for downstream compute_iqr / CVaR consumers — but we discard the
* non-median values here. The 4-fold reduction in cost vs the
* legacy 32-quantile mean comes from the kernel-wide
* IQN_NUM_QUANTILES drop (32 → 5), not from skipping non-medians.
* IQN_MEDIAN_INDEX = 2 is anchored to the FIXED_TAUS table layout
* `[0.05, 0.25, 0.50, 0.75, 0.95]` populated by the constructor's
* HtoD; if that ordering ever changes the seed-fingerprint catches
* the mismatch at checkpoint load before this kernel executes. */
const int IQN_MEDIAN_INDEX = 2;
/* Branch 0 */
for (int a = 0; a < b0_size; a++) {
float partial = 0.0f;
@@ -987,7 +1015,7 @@ void iqn_forward_kernel(
for (int h = tid; h < hidden_dim; h += IQN_BLOCK_SIZE)
partial += w_row[h] * comb_dist[h / IQN_BLOCK_SIZE];
float q_val = iqn_block_sum(partial, shmem_reduce) + b_b0[a];
q_acc[a] += q_val;
if (t == IQN_MEDIAN_INDEX) q_acc[a] = q_val;
}
/* Branch 1 */
for (int a = 0; a < b1_size; a++) {
@@ -996,7 +1024,7 @@ void iqn_forward_kernel(
for (int h = tid; h < hidden_dim; h += IQN_BLOCK_SIZE)
partial += w_row[h] * comb_dist[h / IQN_BLOCK_SIZE];
float q_val = iqn_block_sum(partial, shmem_reduce) + b_b1[a];
q_acc[b0_size + a] += q_val;
if (t == IQN_MEDIAN_INDEX) q_acc[b0_size + a] = q_val;
}
/* Branch 2 */
for (int a = 0; a < b2_size; a++) {
@@ -1005,7 +1033,7 @@ void iqn_forward_kernel(
for (int h = tid; h < hidden_dim; h += IQN_BLOCK_SIZE)
partial += w_row[h] * comb_dist[h / IQN_BLOCK_SIZE];
float q_val = iqn_block_sum(partial, shmem_reduce) + b_b2[a];
q_acc[b0_size + b1_size + a] += q_val;
if (t == IQN_MEDIAN_INDEX) q_acc[b0_size + b1_size + a] = q_val;
}
/* Branch 3 */
for (int a = 0; a < b3_size; a++) {
@@ -1014,15 +1042,15 @@ void iqn_forward_kernel(
for (int h = tid; h < hidden_dim; h += IQN_BLOCK_SIZE)
partial += w_row[h] * comb_dist[h / IQN_BLOCK_SIZE];
float q_val = iqn_block_sum(partial, shmem_reduce) + b_b3[a];
q_acc[b0_size + b1_size + b2_size + a] += q_val;
if (t == IQN_MEDIAN_INDEX) q_acc[b0_size + b1_size + b2_size + a] = q_val;
}
}
/* Write mean Q-values (averaged over quantiles) */
float inv_n = 1.0f / (float)IQN_NUM_QUANTILES;
/* Plan 4 Task 3 (E.3): write median Q (τ=0.50). No mean-over-quantiles
* division — q_acc already holds the median quantile values per action. */
if (tid == 0) {
for (int a = 0; a < tba; a++)
expected_q[sample * tba + a] = q_acc[a] * inv_n;
expected_q[sample * tba + a] = q_acc[a];
}
}
@@ -1109,58 +1137,20 @@ void iqn_decode_actions_kernel(
branch_actions[i * 4 + 3] = urgency;
}
/* ── GPU τ sampling kernel (Philox-based PRNG) ─────────────────────────
/* Plan 4 Task 3 (E.3): iqn_sample_taus_kernel deleted.
*
* Generates uniform random τ ∈ (0.01, 0.99) directly on GPU as f32.
* Uses Philox 4×32 counter-based PRNG — deterministic, high quality,
* no global state. Each thread produces one τ value from (thread_id, seed).
* The legacy Philox-driven random τ sampler is obsolete under the fixed-τ
* multi-quantile design. The Rust constructor uploads FIXED_TAUS =
* {0.05, 0.25, 0.50, 0.75, 0.95} into `online_taus` / `target_taus` once
* via clone_htod (broadcast B times to fill the [B, 5] buffer); both the
* online and target IQN forwards and the CVaR cold path read this static
* device buffer directly. No per-step kernel launch is needed.
*
* Grid: (ceil(total/256), 1, 1), Block: (256, 1, 1)
* total = batch_size × num_quantiles
* Pre-edit consumers: only `compute_cvar_scales` invoked this kernel; that
* call site is removed in lockstep. No other Rust-side or kernel-side
* consumer existed (verified by grep on `iqn_sample_taus_kernel`).
*/
/* Philox 4×32-10 round function — same PRNG used by PyTorch/JAX */
__device__ __forceinline__
unsigned int iqn_philox_single(unsigned int counter, unsigned int key) {
unsigned int hi, lo;
/* Philox S-box constants */
lo = counter * 0xD2511F53u;
hi = __umulhi(counter, 0xD2511F53u);
/* 10 rounds of mixing */
for (int r = 0; r < 10; r++) {
unsigned int t = hi ^ key;
hi = lo * 0xCD9E8D57u;
lo = __umulhi(lo, 0xCD9E8D57u);
lo ^= t;
key += 0x9E3779B9u; /* golden ratio */
}
return lo ^ hi;
}
extern "C" __global__
void iqn_sample_taus_kernel(
float* __restrict__ taus,
unsigned int seed,
int total,
int shared_h1, /* runtime: unused, for consistent interface */
int hidden_dim, /* runtime: unused, for consistent interface */
int embed_dim, /* runtime: unused, for consistent interface */
int b0_size, /* runtime: unused, for consistent interface */
int b1_size, /* runtime: unused, for consistent interface */
int b2_size, /* runtime: unused, for consistent interface */
int b3_size /* runtime: unused, for consistent interface */
)
{
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= total) return;
/* Counter = thread index, key = seed — unique per (step, online/target) */
unsigned int bits = iqn_philox_single((unsigned int)i, seed);
/* Map to (0, 1) then clamp to (0.01, 0.99) */
float u = (float)(bits >> 8) * (1.0f / 16777216.0f); /* 24-bit mantissa */
taus[i] = 0.01f + u * 0.98f;
}
/* ================================================================== */
/* Element-wise kernels for cuBLAS-based IQN forward/backward */
/* ================================================================== */

View File

@@ -0,0 +1,106 @@
/*
* iqn_quantile_ema_update — Plan 4 Task 3 (E.3) producer kernel.
*
* Reads the IQN online forward's per-quantile Q surface
* (`save_q_online [TBA, B*Q]` col-major; populated by
* `GpuIqnHead::execute_training_pipeline` after each step) and EMAs the
* mean Q at each off-median fixed-τ position into ISV[99..103):
*
* FIXED_TAUS = [0.05, 0.25, 0.50, 0.75, 0.95]
*
* blockIdx.x = 0 → τ_idx 0 (τ=0.05) → ISV[IQN_Q_P05_EMA_INDEX = 99]
* blockIdx.x = 1 → τ_idx 1 (τ=0.25) → ISV[IQN_Q_P25_EMA_INDEX = 100]
* blockIdx.x = 2 → τ_idx 3 (τ=0.75) → ISV[IQN_Q_P75_EMA_INDEX = 101]
* blockIdx.x = 3 → τ_idx 4 (τ=0.95) → ISV[IQN_Q_P95_EMA_INDEX = 102]
*
* The median (τ=0.50, idx=2) is intentionally skipped — it is already
* surfaced via the existing greedy-Q diagnostic, no duplication needed.
*
* Reduction layout (one block per ISV slot):
* For target τ_idx t, walk every (sample b, action a) and read
* save_q_online[a + (b*Q + t) * TBA]. Sum across (B × TBA) entries,
* divide by N = B × TBA, EMA into ISV[target_slot]. Single block,
* 256-thread shmem-tree reduce — no atomicAdd
* (per `feedback_no_atomicadd.md`).
*
* Cold-path-cadence: launched once per training step alongside
* `h_s2_rms_ema_update`. Diagnostic only — no consumer kernel reads
* slots [99..103) in this commit. ISV is pinned device-mapped, so host
* HEALTH_DIAG sees the new values without an explicit DtoH.
*
* The EMA α parameter is passed by the host (matches the `ema_alpha`
* already plumbed into the other Plan 4 producers from `training_loop.rs`).
*/
#include <cuda_runtime.h>
extern "C" __global__ void iqn_quantile_ema_update(
const float* __restrict__ save_q_online, /* [TBA, B*Q] col-major */
int batch_size, /* B */
int num_quantiles, /* Q (= 5 under fixed-τ) */
int tba, /* total branch actions */
float* __restrict__ isv, /* pinned device-mapped ISV bus */
int isv_q_p05_idx, /* IQN_Q_P05_EMA_INDEX = 99 */
int isv_q_p25_idx, /* = 100 */
int isv_q_p75_idx, /* = 101 */
int isv_q_p95_idx, /* = 102 */
float ema_alpha /* per-step EMA rate (0,1) */
) {
/* Map blockIdx.x → (τ_idx within FIXED_TAUS, target ISV slot).
* Skip the median (idx=2) — that's the existing greedy-Q diagnostic. */
int tau_idx;
int target_slot;
switch (blockIdx.x) {
case 0: tau_idx = 0; target_slot = isv_q_p05_idx; break;
case 1: tau_idx = 1; target_slot = isv_q_p25_idx; break;
case 2: tau_idx = 3; target_slot = isv_q_p75_idx; break;
case 3: tau_idx = 4; target_slot = isv_q_p95_idx; break;
default: return;
}
/* Bounds check — caller must launch with grid.x = 4 and num_quantiles
* matching FIXED_TAUS.len() = 5. The defensive guard returns the EMA
* unchanged (no write) if the layout is incompatible, surfacing the
* misconfiguration to the host via a stale ISV value rather than a
* silent bad-index read. */
if (tau_idx >= num_quantiles) return;
extern __shared__ float smem[];
const int tid = (int)threadIdx.x;
const int block = (int)blockDim.x;
const int N = batch_size * tba;
/* Sum |Q(s, a; τ_idx)| across (sample, action). Absolute value tracks
* Q-magnitude consistently with the spec's risk-monitoring purpose
* (sign of an off-median quantile alone is not a useful diagnostic;
* its scale relative to the median is). */
float local_sum = 0.0f;
for (int i = tid; i < N; i += block) {
int a = i % tba;
int b = i / tba;
/* Col-major: save_q_online[row=a, col=(b*Q + tau_idx)] */
long col = (long)b * (long)num_quantiles + (long)tau_idx;
long flat = (long)a + col * (long)tba;
local_sum += fabsf(save_q_online[flat]);
}
smem[tid] = local_sum;
__syncthreads();
/* In-block tree reduction (no atomicAdd). */
for (int s = block / 2; s > 0; s >>= 1) {
if (tid < s) smem[tid] += smem[tid + s];
__syncthreads();
}
/* Thread 0 writes the EMA. The host caller's trainer config requires
* batch_size > 0 and tba > 0, so N > 0 is invariant; no defensive
* `(N > 0) ? N : 1` here — masking the divide would silently EMA
* toward 0 if the invariant ever broke; the NaN that division-by-zero
* produces propagates to the ISV slot and downstream HEALTH_DIAG
* loudly visible. Same convention as h_s2_rms_ema_update. */
if (tid == 0) {
const float mean_abs_q = smem[0] / (float)N;
const float prev = isv[target_slot];
isv[target_slot] = (1.0f - ema_alpha) * prev + ema_alpha * mean_abs_q;
}
}

View File

@@ -1525,8 +1525,13 @@ impl DQNHyperparameters {
// Conservative Q-Learning (CQL)
cql_alpha: 1.0, // Strong conservatism — penalizes OOS-destructive Q-values
// QR-DQN (complementary to C51 — IQN for quantile estimation)
num_quantiles: 32, // Default: 32 quantiles
// QR-DQN (complementary to C51 — IQN for quantile estimation).
// Plan 4 Task 3 (E.3): IQN is pinned to FIXED_TAUS.len() (= 5)
// at the GpuIqnConfig construction site. This hyperparam slot
// stays for compat (search-space definitions / older checkpoint
// metadata may reference it) but the IQN forward no longer
// honours it — overriding here would not change runtime behaviour.
num_quantiles: 5, // = FIXED_TAUS.len(); was 32.
qr_kappa: 1.0, // Default: 1.0 (standard quantile Huber loss)
iqn_lambda: 0.25, // Default: mild IQN regularization alongside C51
spectral_norm_sigma_max: 3.0, // Default: permits Xavier scaling [1.0, 10.0]
@@ -1727,9 +1732,13 @@ pub(crate) fn dqn_default_config() -> DQNConfig {
// Portfolio Tracking
initial_capital: 100_000.0,
// CQL + IQN (2026 modernization)
// CQL + IQN (2026 modernization). Plan 4 Task 3 (E.3) pins IQN to
// FIXED_TAUS.len() (= 5); the field below is held for compat (older
// checkpoints / hyperopt configs may reference it) but no longer
// controls the IQN forward — see fused_training.rs::new for the
// canonical wiring.
cql_alpha: 1.0,
iqn_num_quantiles: 32, // 64→32: saves 4.3GB VRAM, halves IQN GEMMs, diminishing returns past 32
iqn_num_quantiles: 5, // Plan 4 Task 3 (E.3): FIXED_TAUS.len(); was 32 — see comment.
iqn_kappa: 1.0,
iqn_embedding_dim: 64,
cvar_alpha: 0.05,

View File

@@ -392,7 +392,12 @@ impl FusedTrainingCtx {
spectral_norm_sigma_max: hyperparams.spectral_norm_sigma_max,
spectral_decoupling_lambda: hyperparams.spectral_decoupling_lambda,
iqn_lambda: hyperparams.iqn_lambda,
iqn_num_quantiles: if hyperparams.iqn_lambda > 0.0 { hyperparams.num_quantiles } else { 0 },
// Plan 4 Task 3 (E.3): pin to FIXED_TAUS.len() (= 5) — see comment
// at the GpuIqnConfig construction below for why hyperparams
// `num_quantiles` no longer controls the IQN.
iqn_num_quantiles: if hyperparams.iqn_lambda > 0.0 {
crate::cuda_pipeline::gpu_iqn_head::FIXED_TAUS.len()
} else { 0 },
iqn_embedding_dim: dqn.config.iqn_embedding_dim,
iqn_kappa: dqn.config.iqn_kappa,
entropy_coefficient: dqn.config.entropy_coefficient as f32,
@@ -530,12 +535,23 @@ impl FusedTrainingCtx {
let gpu_iql_low = GpuIqlTrainer::new(shared_cublas_for_iql, iql_low_config)
.map_err(|e| anyhow::anyhow!("GPU IQL (low-tau) init: {e}"))?;
// Initialize GPU IQN dual-head when iqn_lambda > 0 (CVaR risk sizing)
// Initialize GPU IQN dual-head when iqn_lambda > 0 (CVaR risk sizing).
//
// Plan 4 Task 3 (E.3): `num_quantiles` is pinned to `FIXED_TAUS.len()`
// (= 5) regardless of `hyperparams.num_quantiles`. The kernel-side
// `IQN_NUM_QUANTILES` macro sizes register arrays in lockstep with
// FIXED_TAUS; passing the legacy hyperparam value (32) here would
// produce a host/kernel layout mismatch — the kernel writes 5 floats
// per sample but a 32-sized host buffer would over-read garbage. The
// hyperparam slot is left in place for now (downstream config
// serialisation, hyperopt search-space definitions) but no longer
// controls IQN; a follow-up cleanup commit can remove the field
// once all hyperopt configs migrate.
let gpu_iqn = if hyperparams.iqn_lambda > 0.0 {
let iqn_config = GpuIqnConfig {
hidden_dim: shared_h2,
embed_dim: dqn.config.iqn_embedding_dim,
num_quantiles: hyperparams.num_quantiles,
num_quantiles: crate::cuda_pipeline::gpu_iqn_head::FIXED_TAUS.len(),
kappa: dqn.config.iqn_kappa,
state_dim: ml_core::state_layout::STATE_DIM,
shared_h1,
@@ -2458,6 +2474,23 @@ impl FusedTrainingCtx {
Ok(())
}
/// Plan 4 Task 3 (E.3): launch the IQN multi-quantile diagnostic EMA
/// kernel. Reads `gpu_iqn.save_q_online` and EMAs the four off-median
/// fixed quantiles into ISV[99..103). No-op when IQN is inactive (no
/// quantile surface to read).
pub(crate) fn launch_iqn_quantile_ema(&self, ema_alpha: f32) -> Result<(), crate::MLError> {
let iqn = match self.gpu_iqn.as_ref() {
Some(iqn) => iqn,
None => return Ok(()),
};
self.trainer.launch_iqn_quantile_ema(
iqn.save_q_online_ptr(),
iqn.tba(),
iqn.num_quantiles(),
ema_alpha,
)
}
/// Device pointer to ISV signals [8] pinned buffer for adaptive hold enforcement.
/// Returns the dev_ptr (u64) that the experience collector passes to env_step.
pub(crate) fn isv_signals_dev_ptr(&self) -> u64 {

View File

@@ -0,0 +1,168 @@
//! Smoke test: Plan 4 Task 3 (E.3) IQN multi-quantile diagnostic ISV slots.
//!
//! Validates that the new `iqn_quantile_ema_update` producer kernel populates
//! ISV[IQN_Q_P05_EMA_INDEX..=IQN_Q_P95_EMA_INDEX] (slots 99..103) with sane,
//! non-degenerate, finite values after training has run long enough for the
//! IQN forward to write meaningful per-quantile estimates into
//! `GpuIqnHead::save_q_online`.
//!
//! ## Pass criteria
//!
//! 1. All four off-median slots (Q_p05/Q_p25/Q_p75/Q_p95) finish with FINITE
//! values — no NaN/Inf. Cold-start is 0.0 across all four; the producer
//! kernel runs once per step at α=0.05 (≈13-batch half-life), so a
//! finite-but-zero value would indicate the kernel never fired. We
//! therefore additionally require at least one of the four values to be
//! strictly positive — that proves the EMA actually accumulated mean |Q|
//! rather than staying pinned at the cold-start 0.
//!
//! 2. The values are non-degenerate as a set — `max - min > 0`. Strict
//! per-(s,a) quantile monotonicity (Q(τ=0.05) ≤ Q(τ=0.25) ≤ ... ≤ Q(τ=0.95)
//! for every state-action pair) is a well-known soft property of IQN that
//! fails on a noticeable fraction of (s,a) pairs even on converged
//! policies, and these slots hold mean |Q| over (B × tba) which doesn't
//! preserve per-pair monotonicity anyway. Asserting "the four values
//! aren't all equal" is the loaded-bearing check at the ISV-aggregate
//! level: if the producer kernel reduced the same column for every τ,
//! the four EMAs would converge to the same number and we'd catch the
//! bug here.
//!
//! Run:
//! ```bash
//! FOXHUNT_TEST_DATA=test_data/futures-baseline SQLX_OFFLINE=true \
//! cargo test --release -p ml --lib -- \
//! iqn_multi_quantile_heads_produce_monotonic_estimates \
//! --ignored --nocapture
//! ```
use super::helpers::*;
use crate::cuda_pipeline::gpu_dqn_trainer::{
IQN_Q_P05_EMA_INDEX, IQN_Q_P25_EMA_INDEX,
IQN_Q_P75_EMA_INDEX, IQN_Q_P95_EMA_INDEX,
};
#[test]
#[ignore] // GPU + fxcache data
fn iqn_multi_quantile_heads_produce_monotonic_estimates() -> anyhow::Result<()> {
use crate::fxcache;
let cache_dir = feature_cache_dir();
assert!(
cache_dir.exists(),
"feature-cache not found at {:?} — run precompute_features first",
cache_dir,
);
let entries: Vec<_> = std::fs::read_dir(&cache_dir)?
.filter_map(|e| e.ok())
.filter(|e| e.path().extension().and_then(|s| s.to_str()) == Some("fxcache"))
.collect();
assert!(!entries.is_empty(), "No .fxcache files in {cache_dir:?}");
let fxcache_data = fxcache::load_fxcache(&entries[0].path())?;
// 5000 bars at smoke-config batch_size=64 ≈ 62 train batches/epoch; a single
// epoch is enough to fire the producer kernel ~50× and saturate the α=0.05
// EMA toward the measured per-quantile mean |Q|. More epochs would slow
// the test without changing the assertion outcome.
let n = fxcache_data.bar_count.min(5_000);
let train_end = (n * 80) / 100;
let mut params = smoke_params();
params.epochs = 1;
params.early_stopping_enabled = false;
// The smoke profile already sets iqn_lambda > 0 (= 0.25), so the IQN
// head is constructed and `execute_training_pipeline` runs each step.
// Defensive assertion below verifies the assumption rather than silently
// skipping the test if the profile ever flips iqn_lambda to 0.
assert!(
params.iqn_lambda > 0.0,
"IQN must be enabled (iqn_lambda > 0) for this test; smoke profile \
currently has iqn_lambda={}",
params.iqn_lambda,
);
let mut trainer = smoke_trainer_with(params)?;
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()?;
let features = &fxcache_data.features[..n];
let targets = &fxcache_data.targets[..n];
let ofi = &fxcache_data.ofi[..n.min(fxcache_data.ofi.len())];
let train_feat = &features[..train_end];
let train_targets = &targets[..train_end];
let val_feat = &features[train_end..n];
let val_targets = &targets[train_end..n];
rt.block_on(trainer.init_from_fxcache(features, targets, ofi))?;
trainer.set_training_range(0, train_end, train_end, n);
trainer.set_val_data_from_slices(val_feat, val_targets, train_end);
rt.block_on(trainer.reset_for_fold())?;
rt.block_on(trainer.train_fold_from_slices(
train_feat,
train_targets,
|_epoch, _bytes, _best| Ok("skip".to_owned()),
))?;
// ── Read the four off-median quantile EMA slots ──────────────────────
let fused = trainer
.fused_ctx_mut()
.ok_or_else(|| anyhow::anyhow!("fused_ctx absent — IQN test requires GPU trainer"))?;
let q_p05 = fused.read_isv_signal_at(IQN_Q_P05_EMA_INDEX);
let q_p25 = fused.read_isv_signal_at(IQN_Q_P25_EMA_INDEX);
let q_p75 = fused.read_isv_signal_at(IQN_Q_P75_EMA_INDEX);
let q_p95 = fused.read_isv_signal_at(IQN_Q_P95_EMA_INDEX);
eprintln!(
"[IQN_QUANTILE_EMA] Q_p05={q_p05:.6} Q_p25={q_p25:.6} \
Q_p75={q_p75:.6} Q_p95={q_p95:.6}"
);
// 1. All four values must be finite — NaN/Inf would mean the producer
// kernel divided by zero or read out-of-bounds.
for (name, v) in [
("Q_p05", q_p05),
("Q_p25", q_p25),
("Q_p75", q_p75),
("Q_p95", q_p95),
] {
assert!(
v.is_finite(),
"{name} EMA is not finite: {v} — iqn_quantile_ema_update kernel is \
reading garbage or dividing by zero",
);
}
// 2. At least one slot must be > 0 — proves the kernel actually fired
// and EMA-accumulated mean |Q|. Cold-start is 0.0 for all four; the
// α=0.05 EMA over ~50 batches at any non-zero |Q| signal pulls these
// to strictly positive values. A pinned-at-zero result here means the
// producer kernel never executed (training_loop wiring broken).
let max_val = q_p05.max(q_p25).max(q_p75).max(q_p95);
assert!(
max_val > 0.0,
"All four IQN quantile EMAs are still 0.0 after training — the \
iqn_quantile_ema_update kernel did not fire (or save_q_online is \
all zero, suggesting GpuIqnHead::execute_training_pipeline never \
ran). Q_p05={q_p05} Q_p25={q_p25} Q_p75={q_p75} Q_p95={q_p95}",
);
// 3. The four slots must not all be equal. Identical values across all
// four off-median quantiles would mean the producer kernel was reading
// the same `tau_idx` for every block — a layout bug. A spread of just
// 1e-9 is enough to prove distinctness; use a larger floor (1e-6) to
// avoid floating-point ties on a degenerate (all-zero |Q|) input
// that should already have been caught by check 2.
let min_val = q_p05.min(q_p25).min(q_p75).min(q_p95);
let spread = max_val - min_val;
assert!(
spread > 1e-6,
"IQN quantile EMAs collapsed to a single value (spread={spread:.3e}, \
max={max_val} min={min_val}). The producer kernel likely reads the \
same tau_idx for every block — check the blockIdx → tau_idx switch \
in iqn_quantile_ema_kernel.cu.",
);
Ok(())
}

View File

@@ -40,3 +40,5 @@ mod surrogate_noise_check;
mod mamba2_backward;
#[cfg(test)]
mod soft_reset;
#[cfg(test)]
mod iqn_quantile_monotonicity;

View File

@@ -279,6 +279,27 @@ impl StateResetRegistry {
category: ResetCategory::FoldReset,
description: "ISV[H_S2_RMS_EMA_INDEX=96] — per-batch RMS(save_h_s2) EMA (α=0.05); GPU h_s2_rms_ema_update kernel fills; producer-only in 2c.3c.5, consumer wired in 2c.3c.6 (mag_concat_qdir adaptive-scale). Cold-start 1.0 (neutral RMS) reapplied at fold boundary",
},
// ───── Plan 4 Task 3 (E.3): IQN multi-quantile diagnostic EMAs ──
RegistryEntry {
name: "isv_iqn_q_p05_ema",
category: ResetCategory::FoldReset,
description: "ISV[IQN_Q_P05_EMA_INDEX=99] — mean |Q| at IQN τ=0.05 EMA; GPU iqn_quantile_ema_update kernel fills (Plan 4 Task 3 E.3). Cold-start 0.0 reapplied at fold boundary. Diagnostic only — no consumer kernel reads slot 99 in this commit",
},
RegistryEntry {
name: "isv_iqn_q_p25_ema",
category: ResetCategory::FoldReset,
description: "ISV[IQN_Q_P25_EMA_INDEX=100] — mean |Q| at IQN τ=0.25 EMA; GPU iqn_quantile_ema_update kernel fills (Plan 4 Task 3 E.3). Cold-start 0.0 reapplied at fold boundary. Diagnostic only",
},
RegistryEntry {
name: "isv_iqn_q_p75_ema",
category: ResetCategory::FoldReset,
description: "ISV[IQN_Q_P75_EMA_INDEX=101] — mean |Q| at IQN τ=0.75 EMA; GPU iqn_quantile_ema_update kernel fills (Plan 4 Task 3 E.3). Cold-start 0.0 reapplied at fold boundary. Diagnostic only",
},
RegistryEntry {
name: "isv_iqn_q_p95_ema",
category: ResetCategory::FoldReset,
description: "ISV[IQN_Q_P95_EMA_INDEX=102] — mean |Q| at IQN τ=0.95 EMA; GPU iqn_quantile_ema_update kernel fills (Plan 4 Task 3 E.3). Cold-start 0.0 reapplied at fold boundary. Diagnostic only",
},
];
Self { entries }
}

View File

@@ -287,7 +287,10 @@ impl DQNTrainer {
gradient_collapse_patience: hyperparams.gradient_collapse_patience,
cql_alpha: hyperparams.cql_alpha,
iqn_num_quantiles: hyperparams.num_quantiles, // Controlled by hyperopt
// Plan 4 Task 3 (E.3): IQN is pinned to FIXED_TAUS.len() (= 5).
// The hyperparams field stays for compat but no longer controls
// the IQN — see fused_training.rs::new for rationale.
iqn_num_quantiles: crate::cuda_pipeline::gpu_iqn_head::FIXED_TAUS.len(),
iqn_kappa: hyperparams.qr_kappa, // Controlled by hyperopt
iqn_embedding_dim: 64, // Fixed (not in search space)
iqn_lambda: hyperparams.iqn_lambda, // IQN dual-head loss weight

View File

@@ -2816,6 +2816,20 @@ impl DQNTrainer {
tracing::warn!("Plan 4 2c.3c.5 h_s2_rms_ema launch failed: {e}");
}
}
// Plan 4 Task 3 (E.3): per-step IQN multi-quantile diagnostic
// EMAs into ISV[99..103). 4-block kernel reads the IQN online
// forward's `save_q_online [TBA, B*Q]` (populated by
// `GpuIqnHead::execute_training_pipeline` after each step) and
// EMAs mean |Q| at the four off-median fixed-τ positions
// {0.05, 0.25, 0.75, 0.95}. The median (τ=0.50) is intentionally
// skipped — already in the greedy-Q diagnostic. Same launch
// cadence as `h_s2_rms_ema` above. No-op when IQN is inactive.
if let Some(ref fused) = self.fused_ctx {
if let Err(e) = fused.launch_iqn_quantile_ema(ema_alpha) {
tracing::warn!("Plan 4 Task 3 iqn_quantile_ema launch failed: {e}");
}
}
}
// B.2 Plan 3 Task 3: freeze TRADE_TARGET_RATE at the configured
@@ -4495,6 +4509,28 @@ impl DQNTrainer {
);
}
}
// Plan 4 Task 3 (E.3): IQN multi-quantile diagnostic EMA slots.
// Reset to 0.0 at fold boundary so the first iqn_quantile_ema_update
// fire on the new fold EMAs the measured per-quantile Q toward 0
// rather than carrying forward a stale fold-N estimate. Diagnostic
// only — no consumer kernel reads slots 99..103 in this commit.
"isv_iqn_q_p05_ema" | "isv_iqn_q_p25_ema"
| "isv_iqn_q_p75_ema" | "isv_iqn_q_p95_ema" => {
if let Some(ref fused) = self.fused_ctx {
let idx = match name {
"isv_iqn_q_p05_ema"
=> crate::cuda_pipeline::gpu_dqn_trainer::IQN_Q_P05_EMA_INDEX,
"isv_iqn_q_p25_ema"
=> crate::cuda_pipeline::gpu_dqn_trainer::IQN_Q_P25_EMA_INDEX,
"isv_iqn_q_p75_ema"
=> crate::cuda_pipeline::gpu_dqn_trainer::IQN_Q_P75_EMA_INDEX,
"isv_iqn_q_p95_ema"
=> crate::cuda_pipeline::gpu_dqn_trainer::IQN_Q_P95_EMA_INDEX,
_ => unreachable!(),
};
fused.trainer().write_isv_signal_at(idx, 0.0);
}
}
"isv_gamma_dir_eff" | "isv_gamma_mag_eff" | "isv_gamma_ord_eff" | "isv_gamma_urg_eff" => {
// D.2 per-branch gamma slots. Reset to 0.0 at fold boundary;
// per_branch_gamma_update GPU kernel re-populates on next epoch-boundary launch.

File diff suppressed because one or more lines are too long