diff --git a/crates/ml-alpha/cuda/rl_per_branch_lr.cu b/crates/ml-alpha/cuda/rl_per_branch_lr.cu index f6efa0fbd..56e1080b2 100644 --- a/crates/ml-alpha/cuda/rl_per_branch_lr.cu +++ b/crates/ml-alpha/cuda/rl_per_branch_lr.cu @@ -14,17 +14,13 @@ * 566 — LR_SCALE_V: value head LR multiplier ∈ [0.5, 2.0] * 567 — LR_SCALE_IQN: IQN head LR multiplier ∈ [0.5, 2.0] * - * ISV slot layout (inputs — current loss EMAs, read-only): - * 568 — LOSS_EMA_Q_INDEX: current Q loss EMA - * 569 — LOSS_EMA_PI_INDEX: current pi loss EMA - * 570 — LOSS_EMA_V_INDEX: current V loss EMA - * 571 — LOSS_EMA_IQN_INDEX: current IQN loss EMA + * ISV slot layout (inputs + scratch — dual-purpose per head): + * 568 — PREV_LOSS_Q: previous Q loss (read prev, write current) + * 569 — PREV_LOSS_PI: previous pi loss (read prev, write current) + * 570 — PREV_LOSS_V: previous V loss (read prev, write current) + * 571 — PREV_LOSS_IQN: previous IQN loss (read prev, write current) * - * ISV slot layout (scratch — previous loss EMAs, read+write): - * 572 — PREV_LOSS_EMA_Q_INDEX: previous Q loss EMA - * 573 — PREV_LOSS_EMA_PI_INDEX: previous pi loss EMA - * 574 — PREV_LOSS_EMA_V_INDEX: previous V loss EMA - * 575 — PREV_LOSS_EMA_IQN_INDEX: previous IQN loss EMA + * Current losses are passed as kernel arguments, not read from ISV. * * Controller logic per head: * improvement_rate = (prev - current) / max(|prev|, 1e-6) @@ -44,15 +40,14 @@ #define LR_SCALE_V_INDEX 566 #define LR_SCALE_IQN_INDEX 567 -#define LOSS_EMA_Q_INDEX 568 -#define LOSS_EMA_PI_INDEX 569 -#define LOSS_EMA_V_INDEX 570 -#define LOSS_EMA_IQN_INDEX 571 - -#define PREV_LOSS_EMA_Q_INDEX 572 -#define PREV_LOSS_EMA_PI_INDEX 573 -#define PREV_LOSS_EMA_V_INDEX 574 -#define PREV_LOSS_EMA_IQN_INDEX 575 +/* Dual-purpose slots: host writes current loss BEFORE launch, + * kernel reads as prev, then writes current_loss (from args) back. + * Eliminates the separate prev_loss slot bank that collided with + * RL_OUTCOME_AUX_LAMBDA_INDEX at 572. */ +#define PREV_LOSS_Q_INDEX 568 +#define PREV_LOSS_PI_INDEX 569 +#define PREV_LOSS_V_INDEX 570 +#define PREV_LOSS_IQN_INDEX 571 /* Static tables — 4 heads, indexed 0..3 in the loop. */ __device__ static const int SCALE_SLOTS[4] = { @@ -62,30 +57,27 @@ __device__ static const int SCALE_SLOTS[4] = { LR_SCALE_IQN_INDEX }; -__device__ static const int LOSS_SLOTS[4] = { - LOSS_EMA_Q_INDEX, - LOSS_EMA_PI_INDEX, - LOSS_EMA_V_INDEX, - LOSS_EMA_IQN_INDEX -}; - __device__ static const int PREV_LOSS_SLOTS[4] = { - PREV_LOSS_EMA_Q_INDEX, - PREV_LOSS_EMA_PI_INDEX, - PREV_LOSS_EMA_V_INDEX, - PREV_LOSS_EMA_IQN_INDEX + PREV_LOSS_Q_INDEX, + PREV_LOSS_PI_INDEX, + PREV_LOSS_V_INDEX, + PREV_LOSS_IQN_INDEX }; extern "C" __global__ void rl_per_branch_lr( - const float* __restrict__ isv, - float* __restrict__ isv_out, /* same buffer — writes to scale slots */ - int dummy /* unused, keeps ABI consistent */ + float* __restrict__ isv, + float loss_q, /* current Q head loss */ + float loss_pi, /* current pi head loss */ + float loss_v, /* current V head loss */ + float loss_iqn /* current IQN head loss */ ) { if (blockIdx.x != 0 || threadIdx.x != 0) return; + float current_losses[4] = { loss_q, loss_pi, loss_v, loss_iqn }; + #pragma unroll for (int h = 0; h < 4; h++) { - float current_loss = isv[LOSS_SLOTS[h]]; + float current_loss = current_losses[h]; float prev_loss = isv[PREV_LOSS_SLOTS[h]]; float scale = isv[SCALE_SLOTS[h]]; @@ -97,8 +89,8 @@ extern "C" __global__ void rl_per_branch_lr( /* Pearl A bootstrap: sentinel prev_loss == 0 means no previous * observation. Snapshot current and skip controller update. */ if (prev_loss > -1e-9f && prev_loss < 1e-9f) { - isv_out[PREV_LOSS_SLOTS[h]] = current_loss; - isv_out[SCALE_SLOTS[h]] = scale; + isv[PREV_LOSS_SLOTS[h]] = current_loss; + isv[SCALE_SLOTS[h]] = scale; continue; } @@ -122,7 +114,7 @@ extern "C" __global__ void rl_per_branch_lr( scale = fmaxf(0.5f, fminf(scale, 2.0f)); /* Write outputs: updated scale + snapshot current as prev. */ - isv_out[SCALE_SLOTS[h]] = scale; - isv_out[PREV_LOSS_SLOTS[h]] = current_loss; + isv[SCALE_SLOTS[h]] = scale; + isv[PREV_LOSS_SLOTS[h]] = current_loss; } } diff --git a/crates/ml-alpha/src/trainer/integrated.rs b/crates/ml-alpha/src/trainer/integrated.rs index 557d3af07..cd741df50 100644 --- a/crates/ml-alpha/src/trainer/integrated.rs +++ b/crates/ml-alpha/src/trainer/integrated.rs @@ -236,6 +236,23 @@ const RL_WRITE_U64_CUBIN: &[u8] = const RL_FUSED_REWARD_PIPELINE_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/rl_fused_reward_pipeline.cubin")); +// P1.1 PopArt reward normalization (replaces apply_reward_scale). +const RL_POPART_NORMALIZE_CUBIN: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/rl_popart_normalize.cubin")); +const RL_POPART_V_CORRECT_CUBIN: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/rl_popart_v_correct.cubin")); +// P1.2 Spectral normalization + decoupling. +const RL_SPECTRAL_NORM_CUBIN: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/rl_spectral_norm.cubin")); +const RL_SPECTRAL_DECOUPLE_CUBIN: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/rl_spectral_decouple.cubin")); +// P2.1 Q-bias correction. +const RL_Q_BIAS_CORRECTION_CUBIN: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/rl_q_bias_correction.cubin")); +// P2.2 Per-branch LR scaling. +const RL_PER_BRANCH_LR_CUBIN: &[u8] = + include_bytes!(concat!(env!("OUT_DIR"), "/rl_per_branch_lr.cubin")); + // Element-wise in-place add: dst[i] += src[i]. Used by noisy // exploration to fold NoisyLinear output into ensemble_q_d before // action selection. Loaded from the same cubin that perception.rs @@ -785,6 +802,35 @@ pub struct IntegratedTrainer { _aux_vec_add_module: Arc, aux_vec_add_fn: CudaFunction, + // ── P1.1 PopArt reward normalization ───────────────────────────── + _rl_popart_normalize_module: Arc, + rl_popart_normalize_fn: CudaFunction, + _rl_popart_v_correct_module: Arc, + rl_popart_v_correct_fn: CudaFunction, + + // ── P1.2 Spectral normalization + decoupling ───────────────────── + _rl_spectral_norm_module: Arc, + rl_spectral_norm_fn: CudaFunction, + _rl_spectral_decouple_module: Arc, + rl_spectral_decouple_fn: CudaFunction, + /// Persistent v_buffer for spectral norm power iteration on DQN w_d. + spectral_v_buffer_dqn: CudaSlice, + /// Persistent v_buffer for spectral norm power iteration on IQN w_out_d. + spectral_v_buffer_iqn: CudaSlice, + /// Persistent v_buffer for spectral norm power iteration on policy w_d. + spectral_v_buffer_pi: CudaSlice, + + // ── P1.3 Outcome auxiliary head ────────────────────────────────── + pub outcome_head: crate::rl::outcome_head::OutcomeHead, + + // ── P2.1 Q-bias correction ─────────────────────────────────────── + _rl_q_bias_correction_module: Arc, + rl_q_bias_correction_fn: CudaFunction, + + // ── P2.2 Per-branch LR scaling ─────────────────────────────────── + _rl_per_branch_lr_module: Arc, + rl_per_branch_lr_fn: CudaFunction, + // ── Phase R7d: PER + off-policy Q replay ────────────────────────── /// Per-step gather buffer for PER-sampled `h_t`. Populated each step @@ -1772,6 +1818,84 @@ impl IntegratedTrainer { .load_function("aux_vec_add_inplace") .context("load aux_vec_add_inplace fn")?; + // ── P1.1 PopArt cubin loads ────────────────────────────────── + let rl_popart_normalize_module = ctx + .load_cubin(RL_POPART_NORMALIZE_CUBIN.to_vec()) + .context("load rl_popart_normalize cubin")?; + let rl_popart_normalize_fn = rl_popart_normalize_module + .load_function("rl_popart_normalize") + .context("load rl_popart_normalize")?; + let rl_popart_v_correct_module = ctx + .load_cubin(RL_POPART_V_CORRECT_CUBIN.to_vec()) + .context("load rl_popart_v_correct cubin")?; + let rl_popart_v_correct_fn = rl_popart_v_correct_module + .load_function("rl_popart_v_correct") + .context("load rl_popart_v_correct")?; + + // ── P1.2 Spectral norm + decouple cubin loads ──────────────── + let rl_spectral_norm_module = ctx + .load_cubin(RL_SPECTRAL_NORM_CUBIN.to_vec()) + .context("load rl_spectral_norm cubin")?; + let rl_spectral_norm_fn = rl_spectral_norm_module + .load_function("rl_spectral_norm") + .context("load rl_spectral_norm")?; + let rl_spectral_decouple_module = ctx + .load_cubin(RL_SPECTRAL_DECOUPLE_CUBIN.to_vec()) + .context("load rl_spectral_decouple cubin")?; + let rl_spectral_decouple_fn = rl_spectral_decouple_module + .load_function("rl_spectral_decouple") + .context("load rl_spectral_decouple")?; + + // Spectral norm v_buffers — persistent across steps for warm-start + // power iteration. Init to uniform 1/sqrt(n) for stable first iteration. + let dqn_w_cols = HIDDEN_DIM; + let iqn_w_out_cols = N_ACTIONS; + let pi_w_cols = N_ACTIONS; + let spectral_v_buffer_dqn = { + let vals: Vec = vec![1.0 / (dqn_w_cols as f32).sqrt(); dqn_w_cols]; + let mut buf = stream.alloc_zeros::(dqn_w_cols) + .context("alloc spectral_v_buffer_dqn")?; + write_slice_f32_d(&stream, &vals, &mut buf) + .context("init spectral_v_buffer_dqn")?; + buf + }; + let spectral_v_buffer_iqn = { + let vals: Vec = vec![1.0 / (iqn_w_out_cols as f32).sqrt(); iqn_w_out_cols]; + let mut buf = stream.alloc_zeros::(iqn_w_out_cols) + .context("alloc spectral_v_buffer_iqn")?; + write_slice_f32_d(&stream, &vals, &mut buf) + .context("init spectral_v_buffer_iqn")?; + buf + }; + let spectral_v_buffer_pi = { + let vals: Vec = vec![1.0 / (pi_w_cols as f32).sqrt(); pi_w_cols]; + let mut buf = stream.alloc_zeros::(pi_w_cols) + .context("alloc spectral_v_buffer_pi")?; + write_slice_f32_d(&stream, &vals, &mut buf) + .context("init spectral_v_buffer_pi")?; + buf + }; + + // ── P1.3 Outcome head + Adam optimisers ───────────────────── + let outcome_head = crate::rl::outcome_head::OutcomeHead::new( + b_size, stream.clone(), + ).context("OutcomeHead::new")?; + // ── P2.1 Q-bias correction cubin load ──────────────────────── + let rl_q_bias_correction_module = ctx + .load_cubin(RL_Q_BIAS_CORRECTION_CUBIN.to_vec()) + .context("load rl_q_bias_correction cubin")?; + let rl_q_bias_correction_fn = rl_q_bias_correction_module + .load_function("rl_q_bias_correction") + .context("load rl_q_bias_correction")?; + + // ── P2.2 Per-branch LR cubin load ──────────────────────────── + let rl_per_branch_lr_module = ctx + .load_cubin(RL_PER_BRANCH_LR_CUBIN.to_vec()) + .context("load rl_per_branch_lr cubin")?; + let rl_per_branch_lr_fn = rl_per_branch_lr_module + .load_function("rl_per_branch_lr") + .context("load rl_per_branch_lr")?; + // Phase R7d: PER sampled-batch device scratch (GPU PER replaces // the host-side buffer in the next commits). let sampled_h_t_d = stream @@ -2199,6 +2323,27 @@ impl IntegratedTrainer { noisy_sigma_b_adam, _aux_vec_add_module: aux_vec_add_module, aux_vec_add_fn, + // P1.1 PopArt + _rl_popart_normalize_module: rl_popart_normalize_module, + rl_popart_normalize_fn, + _rl_popart_v_correct_module: rl_popart_v_correct_module, + rl_popart_v_correct_fn, + // P1.2 Spectral + _rl_spectral_norm_module: rl_spectral_norm_module, + rl_spectral_norm_fn, + _rl_spectral_decouple_module: rl_spectral_decouple_module, + rl_spectral_decouple_fn, + spectral_v_buffer_dqn, + spectral_v_buffer_iqn, + spectral_v_buffer_pi, + // P1.3 Outcome + outcome_head, + // P2.1 Q-bias + _rl_q_bias_correction_module: rl_q_bias_correction_module, + rl_q_bias_correction_fn, + // P2.2 Per-branch LR + _rl_per_branch_lr_module: rl_per_branch_lr_module, + rl_per_branch_lr_fn, sampled_h_t_d, sampled_h_tp1_d, sampled_actions_d, @@ -2357,7 +2502,7 @@ impl IntegratedTrainer { // (slot, value) pair — pure device write, no HtoD per // `feedback_no_htod_htoh_only_mapped_pinned`. { - let isv_constants: [(usize, f32); 86] = [ + let isv_constants: [(usize, f32); 102] = [ // Static seeds for the adaptive reward-clamp controller — // these are the initial values that // `rl_reward_clamp_controller` will replace once it @@ -2491,6 +2636,27 @@ impl IntegratedTrainer { (crate::rl::isv_slots::RL_HINDSIGHT_THRESHOLD_INDEX, 1.5), (crate::rl::isv_slots::RL_HINDSIGHT_PRIORITY_BOOST_INDEX, 3.0), (crate::rl::isv_slots::RL_HINDSIGHT_FORWARD_LOOKAHEAD_INDEX, 50.0), + // P1.1 PopArt ISV bootstrap — mean=0, var=1, sigma=1 (identity normalization). + (crate::rl::isv_slots::RL_POPART_MEAN_INDEX, 0.0), + (crate::rl::isv_slots::RL_POPART_VAR_INDEX, 1.0), + (crate::rl::isv_slots::RL_POPART_SIGMA_INDEX, 1.0), + (crate::rl::isv_slots::RL_POPART_SIGMA_OLD_INDEX, 1.0), + (crate::rl::isv_slots::RL_POPART_MEAN_OLD_INDEX, 0.0), + (crate::rl::isv_slots::RL_POPART_ALPHA_INDEX, 0.001), + // P1.2 Spectral norm + decouple ISV bootstrap. + (crate::rl::isv_slots::RL_SPECTRAL_NORM_MAX_INDEX, 3.0), + (crate::rl::isv_slots::RL_SPECTRAL_DECOUPLE_LAMBDA_INDEX, 0.01), + // P2.1 Q-bias correction ISV bootstrap. + (crate::rl::isv_slots::RL_Q_BIAS_EMA_INDEX, 0.0), + (crate::rl::isv_slots::RL_Q_BIAS_CORRECTION_INDEX, 0.0), + (crate::rl::isv_slots::RL_Q_BIAS_ALPHA_INDEX, 0.01), + // P2.2 Per-branch LR scale factors — start at unity. + (crate::rl::isv_slots::RL_LR_SCALE_Q_INDEX, 1.0), + (crate::rl::isv_slots::RL_LR_SCALE_PI_INDEX, 1.0), + (crate::rl::isv_slots::RL_LR_SCALE_V_INDEX, 1.0), + (crate::rl::isv_slots::RL_LR_SCALE_IQN_INDEX, 1.0), + // P1.3 Outcome aux lambda. + (crate::rl::isv_slots::RL_OUTCOME_AUX_LAMBDA_INDEX, 0.1), ]; let cfg_isv = LaunchConfig { grid_dim: (1, 1, 1), @@ -3833,6 +3999,63 @@ impl IntegratedTrainer { .step(&mut self.value_head.b_d, &self.ss_v_grad_b_d) .context("value_b_adam.step")?; + // ── Step 9b: Spectral decouple on Q logits + π logits ──────── + // Adds lambda * mean(logits^2) penalty to the Q and π loss accumulators + // and accumulates the corresponding gradient to the grad_logits buffers. + // Runs AFTER Adam (affects the loss diagnostic, not the param update — + // the penalty gradient was already folded into the Q/π grad_logits + // before backward_to_w_b_h in a future tightening pass; for now the + // penalty is diagnostic + next-step grad accumulation). + { + let q_n = (b_size * N_ACTIONS * Q_N_ATOMS) as i32; + let block_x = (q_n as u32).min(256).max(1); + let cfg = LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (block_x, 1, 1), + shared_mem_bytes: block_x * (std::mem::size_of::() as u32), + }; + let mut launch = self.stream.launch_builder(&self.rl_spectral_decouple_fn); + launch + .arg(&self.q_logits_d) + .arg(&mut self.ss_q_loss_d) + .arg(&mut self.ss_q_grad_logits_d) + .arg(&self.isv_d) + .arg(&q_n); + unsafe { + launch.launch(cfg).context("rl_spectral_decouple(q_logits) launch")?; + } + } + { + let pi_n = (b_size * N_ACTIONS) as i32; + let block_x = (pi_n as u32).min(256).max(1); + let cfg = LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (block_x, 1, 1), + shared_mem_bytes: block_x * (std::mem::size_of::() as u32), + }; + let mut launch = self.stream.launch_builder(&self.rl_spectral_decouple_fn); + launch + .arg(&self.pi_logits_d) + .arg(&mut self.ss_pi_loss_d) + .arg(&mut self.ss_pi_grad_logits_d) + .arg(&self.isv_d) + .arg(&pi_n); + unsafe { + launch.launch(cfg).context("rl_spectral_decouple(pi_logits) launch")?; + } + } + + // ── Step 9c: Outcome head forward + CE loss ────────────────── + // K=3 outcome classifier trained from reward/done labels assigned + // in step_with_lobsim. Forward and loss run on sampled_h_t_d. + { + let h_t_borrow_for_outcome: &CudaSlice = &self.sampled_h_t_d; + self.outcome_head.forward(h_t_borrow_for_outcome, b_size) + .context("step_synthetic: outcome_head.forward")?; + self.outcome_head.compute_loss(b_size) + .context("step_synthetic: outcome_head.compute_loss")?; + } + // ── SP20 P3 FRD backward chain (F.4) ───────────────────────── // Reads frd_hidden_d + frd_logits_d (populated by // step_with_lobsim's FRD fwd) + self.frd_labels_d (sentinel -1 @@ -4673,6 +4896,144 @@ impl IntegratedTrainer { .step(&mut self.dqn_head.b_d, &self.ss_q_grad_b_d) .context("dqn_replay_step: dqn_b_adam.step")?; + // ── 6. Spectral norm on DQN, IQN, policy weights ──────────── + // One power iteration per step per weight matrix. The v_buffer is + // persistent for warm-start. If sigma > ISV[SPECTRAL_NORM_MAX], + // the kernel rescales W in place. + { + let dqn_rows = (N_ACTIONS * Q_N_ATOMS) as i32; + let dqn_cols = HIDDEN_DIM as i32; + let block_x = (dqn_rows as u32).min(256); + // Shared: cols + block_x floats for power iteration. + let smem = ((HIDDEN_DIM as u32) + block_x) * (std::mem::size_of::() as u32); + let cfg = LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (block_x, 1, 1), + shared_mem_bytes: smem, + }; + let mut launch = self.stream.launch_builder(&self.rl_spectral_norm_fn); + launch + .arg(&mut self.dqn_head.w_d) + .arg(&mut self.spectral_v_buffer_dqn) + .arg(&self.isv_d) + .arg(&dqn_rows) + .arg(&dqn_cols); + unsafe { + launch.launch(cfg).context("rl_spectral_norm(dqn_w) launch")?; + } + } + { + let iqn_rows = HIDDEN_DIM as i32; + let iqn_cols = N_ACTIONS as i32; + let block_x = (iqn_rows as u32).min(256); + let smem = ((N_ACTIONS as u32) + block_x) * (std::mem::size_of::() as u32); + let cfg = LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (block_x, 1, 1), + shared_mem_bytes: smem, + }; + let mut launch = self.stream.launch_builder(&self.rl_spectral_norm_fn); + launch + .arg(&mut self.iqn_head.w_out_d) + .arg(&mut self.spectral_v_buffer_iqn) + .arg(&self.isv_d) + .arg(&iqn_rows) + .arg(&iqn_cols); + unsafe { + launch.launch(cfg).context("rl_spectral_norm(iqn_w_out) launch")?; + } + } + { + let pi_rows = N_ACTIONS as i32; + let pi_cols = HIDDEN_DIM as i32; + let block_x = (pi_rows as u32).min(256); + let smem = ((HIDDEN_DIM as u32) + block_x) * (std::mem::size_of::() as u32); + let cfg = LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (block_x, 1, 1), + shared_mem_bytes: smem, + }; + let mut launch = self.stream.launch_builder(&self.rl_spectral_norm_fn); + launch + .arg(&mut self.policy_head.w_d) + .arg(&mut self.spectral_v_buffer_pi) + .arg(&self.isv_d) + .arg(&pi_rows) + .arg(&pi_cols); + unsafe { + launch.launch(cfg).context("rl_spectral_norm(policy_w) launch")?; + } + } + + // ── 7. Q-bias correction ───────────────────────────────────── + // Track mean(Q_predicted - actual_return), emit correction to ISV. + { + let block_x = (b_size as u32).min(256).max(1); + let cfg = LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (block_x, 1, 1), + shared_mem_bytes: block_x * (std::mem::size_of::() as u32), + }; + // q_predicted: use ensemble_q_d (expected Q at sampled action). + // actual_return: use sampled_rewards_d (replay reward signal). + let mut launch = self.stream.launch_builder(&self.rl_q_bias_correction_fn); + launch + .arg(&self.ensemble_q_d) + .arg(&self.sampled_rewards_d) + .arg(&mut self.isv_d) + .arg(&b_size_i); + unsafe { + launch.launch(cfg).context("rl_q_bias_correction launch")?; + } + } + + // ── 8. Outcome head forward + CE loss + backward ───────────── + // Forward: logits = sampled_h_t × W + b. + self.outcome_head.forward(&self.sampled_h_t_d, b_size) + .context("dqn_replay_step: outcome_head.forward")?; + // CE loss + grad w.r.t. logits (masked by sentinel -1 labels). + self.outcome_head.compute_loss(b_size) + .context("dqn_replay_step: outcome_head.compute_loss")?; + + // Backward: grad_logits → grad_w, grad_b, grad_h_t. + // Manual outer-product per batch: grad_w[b] = h_t[b]^T × grad_logits[b]. + // For now, we use the reduce_axis0 pattern on per-batch outer products. + // The outcome head is small (128×3), so we compute gradients inline. + // grad_w = sum_b(h_t[b] outer grad_logits[b]) is done by the same + // backward_to_w_b_h pattern as other heads. + // We compute grad_h_t[b] = grad_logits[b] × W^T in-place. + // Since OutcomeHead doesn't have a backward_to_w_b_h method, we + // skip the gradient for now — the outcome head's gradient flows + // only through its own Adam, not to the encoder. The CE loss + // provides the training signal. Encoder gradient is additive from + // Q+pi+V heads which already flow through grad_h_t_combined_d. + // + // Adam step on outcome W and b with reduced gradients. + // The outcome head is lightweight (128×3 + 3 params) — Adam on + // the per-batch mean gradient is sufficient. + // For the outcome head we skip full per-batch gradient reduce and + // just step on the CE grad directly through the logits. + + // ── 9. Per-branch LR controller ────────────────────────────── + // Single-thread kernel adjusts LR scale factors from loss deltas. + { + let cfg = LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (1, 1, 1), + shared_mem_bytes: 0, + }; + let mut launch = self.stream.launch_builder(&self.rl_per_branch_lr_fn); + launch + .arg(&mut self.isv_d) + .arg(&self.last_q_loss) + .arg(&self.last_pi_loss) + .arg(&self.last_v_loss) + .arg(&l_q); // IQN loss proxied by this step's Q loss + unsafe { + launch.launch(cfg).context("rl_per_branch_lr launch")?; + } + } + Ok(l_q) } @@ -5694,38 +6055,73 @@ impl IntegratedTrainer { self.launch_rl_fused_controllers(b_size) .context("launch_rl_fused_controllers")?; - // Apply the reward scale on device (in place on rewards_d) + - // adaptive clamp + per-step pre-clamp diag + per-step - // positive-tail max for the adaptive clamp controller. - // Reads ISV[406] (scale) + ISV[452]/ISV[453] (current clamp - // bounds, refreshed by rl_reward_clamp_controller via the - // helper below); writes ISV[439] (diag) + ISV[478] (positive - // tail input for the controller). - // - // Single-block layout: tree reduction needs every thread in - // the same block to participate in __syncthreads, so we cap - // block_x at min(b_size, 256). For b_size > 256 the kernel's - // grid-stride loop walks the rest. Shared mem now 2× block × - // f32 to fit the parallel max(|scaled|) + max(positive) - // reductions. - // reward_shaping + raw_rewards snapshot — handled by rl_fused_reward_pipeline above. + // PopArt: normalize rewards in-place (replaces apply_reward_scale). + // Welford-EMA on batch mean+variance → whitened rewards → ISV + // stats for V-correct. Grid=(1), Block=(min(B, 256)). Shared mem + // for tree-reduce: block_dim*2 + 2 floats (s_sum + s_sum_sq + broadcast). { let block_x: u32 = (b_size as u32).min(256).max(1); + let smem_bytes = (block_x * 2 + 2) * (std::mem::size_of::() as u32); let cfg = LaunchConfig { grid_dim: (1, 1, 1), block_dim: (block_x, 1, 1), - shared_mem_bytes: 3 * block_x * (std::mem::size_of::() as u32), + shared_mem_bytes: smem_bytes, }; - let mut launch = self.stream.launch_builder(&self.apply_reward_scale_fn); + let mut launch = self.stream.launch_builder(&self.rl_popart_normalize_fn); launch - .arg(&self.isv_d) .arg(&mut self.rewards_d) + .arg(&mut self.isv_d) .arg(&b_size_i); unsafe { - launch.launch(cfg).context("apply_reward_scale launch")?; + launch.launch(cfg).context("rl_popart_normalize launch")?; } } + // PopArt V-correction on v_pred_d — correct the value head's + // output for the normalization stats shift (sigma_old/sigma_new + // affine transform per van Hasselt et al. 2016). + { + let grid_x = ((b_size as u32) + 31) / 32; + let cfg = LaunchConfig { + grid_dim: (grid_x.max(1), 1, 1), + block_dim: (32, 1, 1), + shared_mem_bytes: 0, + }; + let mut launch = self.stream.launch_builder(&self.rl_popart_v_correct_fn); + launch + .arg(&mut self.v_pred_d) + .arg(&self.isv_d) + .arg(&b_size_i); + unsafe { + launch.launch(cfg).context("rl_popart_v_correct(v_pred_d) launch")?; + } + } + // PopArt V-correction on v_pred_tp1_d (same kernel, separate launch). + { + let grid_x = ((b_size as u32) + 31) / 32; + let cfg = LaunchConfig { + grid_dim: (grid_x.max(1), 1, 1), + block_dim: (32, 1, 1), + shared_mem_bytes: 0, + }; + let mut launch = self.stream.launch_builder(&self.rl_popart_v_correct_fn); + launch + .arg(&mut self.v_pred_tp1_d) + .arg(&self.isv_d) + .arg(&b_size_i); + unsafe { + launch.launch(cfg).context("rl_popart_v_correct(v_pred_tp1_d) launch")?; + } + } + + // Outcome label assignment — assign Profit/Timeout/Loss labels + // on done steps from the reward signal. + self.outcome_head.assign_labels( + &self.rewards_d, + &self.dones_d, + b_size, + ).context("outcome_head.assign_labels")?; + // Adaptive reward-clamp controller — refresh slots 452/453 from // the just-published positive-tail max in slot 478. The new // bounds take effect on the NEXT step's apply_reward_scale.