diff --git a/crates/ml-alpha/src/trainer/integrated.rs b/crates/ml-alpha/src/trainer/integrated.rs index c061ce359..c1730a1d2 100644 --- a/crates/ml-alpha/src/trainer/integrated.rs +++ b/crates/ml-alpha/src/trainer/integrated.rs @@ -343,6 +343,12 @@ pub struct IntegratedTrainer { /// DQN distributional Q-head (Phase C). pub dqn_head: DqnHead, + /// IQN distributional Q-head — runs alongside C51 for ensemble + /// action selection. Forward produces `Q(s, τ, a)` for N_TAU + /// quantile samples; expected Q is combined with C51's E[Q] via + /// ISV[RL_IQN_ENSEMBLE_ALPHA_INDEX]. + pub iqn_head: crate::rl::iqn::IqnHead, + /// PPO policy head (Phase D). pub policy_head: PolicyHead, @@ -354,6 +360,11 @@ pub struct IntegratedTrainer { /// `AdamW` wrapper; each instance owns its own m / v / step counter. pub dqn_w_adam: AdamW, pub dqn_b_adam: AdamW, + /// IQN head Adam optimisers — w_embed/b_embed/w_out/b_out. + pub iqn_w_embed_adam: AdamW, + pub iqn_b_embed_adam: AdamW, + pub iqn_w_out_adam: AdamW, + pub iqn_b_out_adam: AdamW, pub policy_w_adam: AdamW, pub policy_b_adam: AdamW, pub value_w_adam: AdamW, @@ -634,6 +645,17 @@ pub struct IntegratedTrainer { /// h_t-as-proxy comment at the call site). pub h_tp1_d: CudaSlice, + // ── IQN per-step device buffers ────────────────────────────────── + /// Per-step quantile fractions `[B × N_TAU]` sampled from U(0,1). + /// Populated each step from the host RNG via mapped-pinned upload. + pub iqn_tau_d: CudaSlice, + /// IQN forward output `[B × N_TAU × N_ACTIONS]`. + pub iqn_q_values_d: CudaSlice, + /// IQN expected Q `[B × N_ACTIONS]` = mean over tau dimension. + pub iqn_expected_q_d: CudaSlice, + /// Host RNG for tau sampling (deterministic from dqn_seed + offset). + iqn_tau_rng: rand_chacha::ChaCha8Rng, + // ── Phase R7d: PER + off-policy Q replay ────────────────────────── /// Prioritised Experience Replay buffer (off-policy DQN training). /// Per `feedback_always_per` PER is always-on. Wired in R7d after @@ -758,6 +780,16 @@ impl IntegratedTrainer { ) .context("DqnHead::new")?; + let iqn_head = crate::rl::iqn::IqnHead::new( + dev, + crate::rl::iqn::IqnHeadConfig { + hidden_dim: HIDDEN_DIM, + seed: cfg.dqn_seed.wrapping_add(100), + n_tau: 32, + }, + ) + .context("IqnHead::new")?; + let policy_head = PolicyHead::new( dev, PpoHeadsConfig { @@ -792,6 +824,14 @@ impl IntegratedTrainer { AdamW::new(dev, dqn_head.w_d.len(), lr_placeholder).context("dqn_w_adam")?; let dqn_b_adam = AdamW::new(dev, dqn_head.b_d.len(), lr_placeholder).context("dqn_b_adam")?; + let iqn_w_embed_adam = + AdamW::new(dev, iqn_head.w_embed_d.len(), lr_placeholder).context("iqn_w_embed_adam")?; + let iqn_b_embed_adam = + AdamW::new(dev, iqn_head.b_embed_d.len(), lr_placeholder).context("iqn_b_embed_adam")?; + let iqn_w_out_adam = + AdamW::new(dev, iqn_head.w_out_d.len(), lr_placeholder).context("iqn_w_out_adam")?; + let iqn_b_out_adam = + AdamW::new(dev, iqn_head.b_out_d.len(), lr_placeholder).context("iqn_b_out_adam")?; let policy_w_adam = AdamW::new(dev, policy_head.w_d.len(), lr_placeholder).context("policy_w_adam")?; let policy_b_adam = @@ -1317,6 +1357,22 @@ impl IntegratedTrainer { .alloc_zeros::(b_size * HIDDEN_DIM) .context("alloc h_tp1_d")?; + // IQN per-step buffers: tau samples, Q(s,τ,a), E[Q(s,a)]. + let iqn_n_tau = iqn_head.n_tau(); + let iqn_tau_d = stream + .alloc_zeros::(b_size * iqn_n_tau) + .context("alloc iqn_tau_d")?; + let iqn_q_values_d = stream + .alloc_zeros::(b_size * iqn_n_tau * N_ACTIONS) + .context("alloc iqn_q_values_d")?; + let iqn_expected_q_d = stream + .alloc_zeros::(b_size * N_ACTIONS) + .context("alloc iqn_expected_q_d")?; + let iqn_tau_rng = + ::seed_from_u64( + cfg.dqn_seed.wrapping_add(0x1C_A0_7A_00), + ); + // Phase R7d: PER buffer + sampled-batch device scratch. // The replay buffer holds per-transition device payloads // (h_t / next_h_t as small `CudaSlice(HIDDEN_DIM)` slices @@ -1391,10 +1447,15 @@ impl IntegratedTrainer { cfg, perception, dqn_head, + iqn_head, policy_head, value_head, dqn_w_adam, dqn_b_adam, + iqn_w_embed_adam, + iqn_b_embed_adam, + iqn_w_out_adam, + iqn_b_out_adam, policy_w_adam, policy_b_adam, value_w_adam, @@ -1534,6 +1595,10 @@ impl IntegratedTrainer { advantages_d, returns_d, h_tp1_d, + iqn_tau_d, + iqn_q_values_d, + iqn_expected_q_d, + iqn_tau_rng, replay, n_step_buffer: (0..b_size).map(|_| Vec::with_capacity(16)).collect(), sampled_h_t_d, @@ -3464,6 +3529,36 @@ impl IntegratedTrainer { .forward(&self.sampled_h_tp1_d, b_size, &mut q_logits_tp1_sampled_d) .context("dqn_replay_step: dqn_head.forward(sampled_h_tp1)")?; + // ── 1b. IQN forward on sampled_h_t (forward-only, loss in + // follow-up). Verifies the head runs on replay data without + // crashing; the expected_q output can feed an ensemble Q later. + { + let n_tau = self.iqn_head.n_tau(); + let tau_len = b_size * n_tau; + let tau_host: Vec = (0..tau_len) + .map(|_| self.iqn_tau_rng.gen::()) + .collect(); + write_slice_f32_d(&self.stream, &tau_host, &mut self.iqn_tau_d) + .context("dqn_replay_step: iqn_tau upload")?; + self.iqn_head + .forward( + &self.sampled_h_t_d, + &self.iqn_tau_d, + b_size, + n_tau, + &mut self.iqn_q_values_d, + ) + .context("dqn_replay_step: iqn_head.forward(sampled_h_t)")?; + self.iqn_head + .expected_q( + &self.iqn_q_values_d, + b_size, + n_tau, + &mut self.iqn_expected_q_d, + ) + .context("dqn_replay_step: iqn_head.expected_q")?; + } + // ── 2. Double-DQN argmax on online Q at h_tp1 ─────────────── { let cfg_argmax = LaunchConfig { @@ -3718,6 +3813,37 @@ impl IntegratedTrainer { .forward(&self.h_tp1_d, &self.isv_d, b_size, &mut v_pred_tp1_d) .context("step_with_lobsim: value_head.forward(h_tp1) for true V(s_{t+1})")?; + // ── IQN forward alongside C51 ──────────────────────────────── + // Sample tau ~ U(0,1) on host, upload via mapped-pinned, then + // run forward + expected_q reduction. The expected Q output is + // available for the Q→π distillation ensemble in a follow-up. + { + let n_tau = self.iqn_head.n_tau(); + let tau_len = b_size * n_tau; + let tau_host: Vec = (0..tau_len) + .map(|_| self.iqn_tau_rng.gen::()) + .collect(); + write_slice_f32_d(&self.stream, &tau_host, &mut self.iqn_tau_d) + .context("step_with_lobsim: iqn_tau upload")?; + self.iqn_head + .forward( + h_t_borrow, + &self.iqn_tau_d, + b_size, + n_tau, + &mut self.iqn_q_values_d, + ) + .context("step_with_lobsim: iqn_head.forward(h_t)")?; + self.iqn_head + .expected_q( + &self.iqn_q_values_d, + b_size, + n_tau, + &mut self.iqn_expected_q_d, + ) + .context("step_with_lobsim: iqn_head.expected_q")?; + } + // ── Step 2b: Forward π logits for log_pi_old. ───────────────── let mut pi_logits_d = self.stream.alloc_zeros::(b_size * N_ACTIONS)?; self.policy_head