feat(rl): real Q/π/V forward+backward + per-head Adam (Phase E.2)

Replaces Phase E.1's placeholder host-side loss scalars with actual GPU
kernel calls. Each of Q, π, V heads now runs:

- Forward kernel on h_t → head outputs (q_logits / pi_logits / v_pred)
- Backward kernel chain → grad_w / grad_b / grad_h_t (per-batch scratch
  reduced by reduce_axis0)
- Adam step using the existing project-wide adamw_step cubin via the
  AdamW wrapper (each head owns 2 instances — one for w, one for b)

New CUDA kernels:

- v_head_fwd_bwd.cu: scalar V head forward + MSE backward. Per-batch
  loss scratch + per-batch grad_w / grad_b scratch — caller reduces
  via reduce_axis0. Per-(batch, c) grad_h_t writes are sole-writer so
  no atomicAdd needed.
- grad_h_accumulate.cu: element-wise grad_h_encoder += λ × grad_h_head
  combiner. Loaded and ready; consumed by Phase E.3 once the encoder
  backward entry point lands.

Extends existing kernels (additive, doesn't break prior callers):

- dqn_distributional_q.cu: adds `dqn_grad_w_b_h_t` that maps the C51
  bwd kernel's grad_logits output into per-batch grad_w / grad_b
  scratch + OVERWRITE grad_h_t. Mirrors aux_heads_bwd's thread-role
  pattern (one thread per HIDDEN_DIM channel, sole writer per slot).
- ppo_clipped_surrogate.cu: adds `ppo_policy_logits_fwd` (standalone
  linear-projection forward producing pi_logits for the surrogate
  kernel to consume) AND `ppo_grad_w_b_h_t` (same backward pattern as
  the DQN extension, with output dim N_ACTIONS=9).

All new kernels honour feedback_no_atomicadd: per-batch grad scratch
reduced by the existing reduce_axis0 path, never atomics. The Phase C/D
loss-scalar atomicAdds (dqn / ppo bwd accumulators) stay as their
loud-flagged deferral.

Per-head Adam state:

- DqnHead: 2 AdamW instances (w_d, b_d). Re-uses the existing
  adamw_step cubin from `trainer::optim::AdamW`.
- PolicyHead: same.
- ValueHead: same.
- Each AdamW owns independent m / v buffers and a device-resident
  step counter (per pearl_no_host_branches_in_captured_graph: counter
  advancement is a captured kernel, not host scalar).

PerceptionTrainer: adds `h_t_view()` accessor — borrows the h_t_d
field that `forward_encoder` populates. IntegratedTrainer uses this
to dispatch Q/π/V head kernels on the same encoder representation
without re-borrowing self.perception mutably between the encoder
forward and the head dispatches (disjoint field borrows).

ENCODER BACKWARD: Phase E.2 DEFERS the encoder-side gradient combine
to Phase E.3 (where the LobSim integration provides a clean entry
point for a new PerceptionTrainer backward-encoder method). The Q/π/V
kernels DO emit grad_h_t but it is not yet wired into the encoder
backward path. The encoder still learns from BCE+aux only in E.2;
the new heads update their own weights via per-head Adam. E.3 lands
the missing piece via `launch_grad_h_accumulate` (wired and ready).

Bellman target distribution: Phase E.2 uses a deterministic
single-atom-mass projection (host-side, mapped-pinned upload). Phase
E.3 replaces with the proper categorical projection kernel that
consumes γ from ISV[400] and the target net's bootstrap atom values.

build.rs: registers `v_head_fwd_bwd` + `grad_h_accumulate` cubins.
Cache-bust v21 → v22 with a verbose changelog entry covering this
phase's contract changes.

The integrated_trainer_smoke #[ignore]-gated GPU test remains gated;
Phase E.3 activates it alongside LobSim. Non-GPU host tests
(loss_balance default + bootstrap) continue to pass.

Companion to Phases A-E.1 (commits 6a46ded7d 9ec43fdb9 56efd96cb
9732a667c 729f110e0).
This commit is contained in:
jgrusewski
2026-05-22 23:25:42 +02:00
parent 729f110e00
commit 2665669b58
9 changed files with 1395 additions and 73 deletions

View File

@@ -37,9 +37,11 @@ const KERNELS: &[&str] = &[
"ppo_clipped_surrogate", // RL Phase D: PPO clipped-surrogate + entropy bonus + value MSE fwd/bwd
"rl_ppo_clip_controller", // RL Phase D: ISV controller emitting ε to ISV[RL_PPO_CLIP_INDEX=402]
"rl_entropy_coef_controller", // RL Phase D: ISV controller emitting entropy bonus weight to ISV[RL_ENTROPY_COEF_INDEX=403]
"v_head_fwd_bwd", // RL Phase E.2: scalar V(s) head fwd + MSE bwd (linear layer; per-batch scratch + reduce_axis0)
"grad_h_accumulate", // RL Phase E.2: element-wise grad_h_encoder += λ × grad_h_head accumulator (one head at a time, serialised by stream)
];
// Cache bust v21 (2026-05-22): RL Phase D — ppo_clipped_surrogate.cu (PPO L_π = -min(r·A, clip(r,1±ε)·A) + L_v = (V-R)² + L_ent = -coef·H(π) fwd/bwd over 9 actions; ε read from ISV[402], coef from ISV[403]), rl_ppo_clip_controller.cu (ε→ISV[402], anchored on KL(π_new‖π_old) EMA, target KL=0.01), rl_entropy_coef_controller.cu (coef→ISV[403], anchored on entropy deficit vs 0.7·ln(9)). Per pearl_controller_anchors_isv_driven, ε + coef are not constants — Wiener-α adaptive with first-observation bootstrap (ε=0.2, coef=0.01).
// Cache bust v22 (2026-05-22): RL Phase E.2 — v_head_fwd_bwd.cu (scalar V head fwd + MSE bwd, per-batch scratch for grad_w/grad_b — reduced by caller via reduce_axis0; no atomicAdd), grad_h_accumulate.cu (element-wise grad_h_encoder += λ × grad_h_head; trainer folds Q/π/V per-head grad_h_t contributions into encoder grad slot via scaled +=; deferred to Phase E.3 — encoder backward integration). ALSO extends dqn_distributional_q.cu with dqn_grad_w_b_h_t (mirrors aux_heads_bwd pattern: grad_logits → grad_w/grad_b/grad_h_t for the C51 head; per-batch scratch) and ppo_clipped_surrogate.cu with ppo_grad_w_b_h_t (same pattern for PPO π head). All new kernels honour feedback_no_atomicadd — per-batch grad scratch reduced by the existing reduce_axis0 path, never atomics.
fn main() {
println!("cargo:rerun-if-changed=build.rs");

View File

@@ -218,3 +218,82 @@ extern "C" __global__ void dqn_distributional_q_bwd(
atomicAdd(loss_out, ce);
}
}
// ─────────────────────────────────────────────────────────────────────
// dqn_grad_w_b_h_t (Phase E.2):
// Given per-batch grad_logits from `dqn_distributional_q_bwd`, propagate
// to weight / bias gradients (per-batch scratch — caller reduces via
// reduce_axis0) and to grad_h_t (per-(batch, c) sole writer; OVERWRITE).
//
// Mathematics (per batch b, hidden unit c):
// grad_w_per_batch[b, a*Q_N_ATOMS + z, c]
// = grad_logits[b, a, z] × h_t[b, c]
// grad_b_per_batch[b, a*Q_N_ATOMS + z]
// = grad_logits[b, a, z]
// grad_h_t[b, c]
// = Σ_{a, z} grad_logits[b, a, z] × w[a*Q_N_ATOMS + z, c]
//
// Block layout (mirrors aux_heads_bwd):
// grid = (B, 1, 1)
// block = (HIDDEN_DIM = 128, 1, 1)
// shared_mem_bytes = (N_ACTIONS * Q_N_ATOMS) * sizeof(float)
// — stages grad_logits row [a*Q_N_ATOMS + z] once
//
// Thread c is the sole writer of grad_w_per_batch[batch, k, c] (for all k
// in [0, N_ACTIONS * Q_N_ATOMS)) and grad_h_t[batch, c]. The first
// N_ACTIONS * Q_N_ATOMS threads (= 189 ≤ 128 → spills past block size)
// would need the bias write — instead, atom-style threads aren't
// available; we use a tile-wise loop where the first 128 threads
// cover the first 128 bias slots and a stride-loop covers the
// remaining 189 - 128 = 61. Per `feedback_no_atomicadd.md`: each k
// has a single thread writing it (one stride loop, single writer per k).
//
// No atomicAdd. No nvrtc.
// ─────────────────────────────────────────────────────────────────────
extern "C" __global__ void dqn_grad_w_b_h_t(
const float* __restrict__ w, // [N_ACTIONS * Q_N_ATOMS, HIDDEN_DIM]
const float* __restrict__ h_t, // [B * HIDDEN_DIM]
const float* __restrict__ grad_logits, // [B * N_ACTIONS * Q_N_ATOMS]
int B,
float* __restrict__ grad_w_per_batch, // [B * N_ACTIONS * Q_N_ATOMS * HIDDEN_DIM]
float* __restrict__ grad_b_per_batch, // [B * N_ACTIONS * Q_N_ATOMS]
float* __restrict__ grad_h_t // [B * HIDDEN_DIM] (OVERWRITE)
) {
// grad_logits row for this batch: [N_ACTIONS * Q_N_ATOMS] = 189 floats.
extern __shared__ float s_gl[]; // [N_ACTIONS * Q_N_ATOMS]
const int batch = blockIdx.x;
const int c = threadIdx.x;
if (batch >= B) return;
const int K_OUT = N_ACTIONS * Q_N_ATOMS; // 189
// Cooperative stage of grad_logits[batch, :] into shared. Stride-loop
// over K_OUT in steps of HIDDEN_DIM (block size).
for (int k = c; k < K_OUT; k += HIDDEN_DIM) {
s_gl[k] = grad_logits[batch * K_OUT + k];
}
__syncthreads();
// ── grad_b_per_batch ─ stride-loop, sole writer per k ────────────
for (int k = c; k < K_OUT; k += HIDDEN_DIM) {
grad_b_per_batch[batch * K_OUT + k] = s_gl[k];
}
if (c >= HIDDEN_DIM) return;
// ── grad_w + grad_h_t (thread role c = hidden unit) ──────────────
const float h_bc = h_t[batch * HIDDEN_DIM + c];
float gh_acc = 0.0f;
for (int k = 0; k < K_OUT; ++k) {
const float g_k = s_gl[k];
// grad_w[batch, k, c] = g_k × h_bc — sole writer per (batch, k, c).
grad_w_per_batch[(long long) batch * K_OUT * HIDDEN_DIM
+ (long long) k * HIDDEN_DIM + c]
= g_k * h_bc;
gh_acc += g_k * w[k * HIDDEN_DIM + c];
}
// OVERWRITE — caller folds via grad_h_accumulate_scaled.
grad_h_t[batch * HIDDEN_DIM + c] = gh_acc;
}

View File

@@ -0,0 +1,38 @@
// grad_h_accumulate.cu — element-wise scaled gradient accumulator
// (Phase E.2 of the integrated RL trainer).
//
// Combines per-head encoder-input gradients into the shared encoder
// grad slot via scaled +=:
//
// grad_h_encoder[i] += lambda × grad_h_head[i]
//
// Used by IntegratedTrainer to fold the Q / π / V heads'
// `grad_h_t [B × HIDDEN_DIM]` outputs into the encoder backward
// input, on top of the BCE + aux contribution already deposited by
// the PerceptionTrainer's own grad chain.
//
// One thread per element. Each thread is the sole writer of its
// destination slot (within a single launch) — no atomicAdd per
// `feedback_no_atomicadd.md`. Multiple launches in sequence on the
// same destination are safe because each launch reads-modifies-writes
// the same slot serially via stream ordering.
//
// Constraints honoured:
// * `feedback_no_atomicadd.md` — sole-writer per (i)
// * `feedback_no_nvrtc.md` — pre-compiled cubin via build.rs
// * `pearl_no_host_branches_in_captured_graph` — `n` is a captured
// scalar; `lambda` is a captured scalar (the trainer reads it from
// the host ISV mirror BEFORE the capture region).
//
// Launch: grid=(ceil(n/256), 1, 1); block=(256, 1, 1); shared=0.
extern "C" __global__ void grad_h_accumulate_scaled(
const float* __restrict__ grad_h_head, // [n] — input head gradient
float lambda, // scalar weight (host-passed)
int n, // total elements
float* __restrict__ grad_h_encoder // [n] — encoder grad accumulator (additive)
) {
const int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= n) return;
grad_h_encoder[i] += lambda * grad_h_head[i];
}

View File

@@ -58,6 +58,51 @@
#define RL_ENTROPY_COEF_INDEX 403
// ─────────────────────────────────────────────────────────────────────
// ppo_policy_logits_fwd (Phase E.2):
// Standalone linear forward producing per-batch logits from h_t. The
// surrogate kernel takes these logits as input; previously the
// integrated trainer had no path to materialise them outside of
// the surrogate compute.
//
// logits[b, a] = b_pi[a] + Σ_c W_pi[a, c] × h_t[b, c]
//
// Block layout:
// grid = (B, N_ACTIONS, 1)
// block = (HIDDEN_DIM, 1, 1)
// shared_mem_bytes = 0 (each thread independently strides; thread 0
// performs the per-output dot product. The
// block over HIDDEN_DIM is sized for backward
// compatibility with grad_w_b_h_t's reuse of
// the same h_t layout in shared.)
//
// Single output slot per block (b, a) → thread 0 reduces inline. Other
// threads idle (block size mirrors backward's HIDDEN_DIM for caller
// uniformity; perf is non-critical since N_ACTIONS=9 means 9×B blocks).
// ─────────────────────────────────────────────────────────────────────
extern "C" __global__ void ppo_policy_logits_fwd(
const float* __restrict__ w, // [N_ACTIONS, HIDDEN_DIM]
const float* __restrict__ b, // [N_ACTIONS]
const float* __restrict__ h_t, // [B * HIDDEN_DIM]
int B,
float* __restrict__ logits // [B * N_ACTIONS]
) {
const int batch = blockIdx.x;
const int act = blockIdx.y;
const int c = threadIdx.x;
if (batch >= B) return;
if (act >= N_ACTIONS) return;
if (c != 0) return; // single-writer per output slot
float acc = b[act];
#pragma unroll
for (int i = 0; i < HIDDEN_DIM; ++i) {
acc += w[act * HIDDEN_DIM + i] * h_t[batch * HIDDEN_DIM + i];
}
logits[batch * N_ACTIONS + act] = acc;
}
// ─────────────────────────────────────────────────────────────────────
// ppo_clipped_surrogate_fwd:
// Computes softmax → log π_new(a_taken|s) → ratio → clipped surrogate +
@@ -280,3 +325,66 @@ extern "C" __global__ void ppo_clipped_surrogate_bwd(
}
grad_logits[base + act] = pg_grad;
}
// ─────────────────────────────────────────────────────────────────────
// ppo_grad_w_b_h_t (Phase E.2):
// Given per-batch grad_logits [B × N_ACTIONS] from
// `ppo_clipped_surrogate_bwd`, propagate to weight / bias gradients
// (per-batch scratch — caller reduces via reduce_axis0) and to grad_h_t
// (per-(batch, c) sole writer; OVERWRITE).
//
// Mathematics (per batch b, action a, hidden unit c):
// grad_w_per_batch[b, a, c] = grad_logits[b, a] × h_t[b, c]
// grad_b_per_batch[b, a] = grad_logits[b, a]
// grad_h_t[b, c] = Σ_a grad_logits[b, a] × w[a, c]
//
// Block layout (mirrors aux_heads_bwd):
// grid = (B, 1, 1)
// block = (HIDDEN_DIM = 128, 1, 1)
// shared_mem_bytes = N_ACTIONS * sizeof(float) (grad_logits stage)
//
// Thread c is the sole writer of grad_w_per_batch[batch, a, c] for all a,
// and of grad_h_t[batch, c]. The first N_ACTIONS threads additionally
// write grad_b_per_batch[batch, a]. Per `feedback_no_atomicadd.md`: no
// atomics; each output slot has exactly one writer.
// ─────────────────────────────────────────────────────────────────────
extern "C" __global__ void ppo_grad_w_b_h_t(
const float* __restrict__ w, // [N_ACTIONS, HIDDEN_DIM]
const float* __restrict__ h_t, // [B * HIDDEN_DIM]
const float* __restrict__ grad_logits, // [B * N_ACTIONS]
int B,
float* __restrict__ grad_w_per_batch, // [B * N_ACTIONS * HIDDEN_DIM]
float* __restrict__ grad_b_per_batch, // [B * N_ACTIONS]
float* __restrict__ grad_h_t // [B * HIDDEN_DIM] (OVERWRITE)
) {
__shared__ float s_gl[N_ACTIONS];
const int batch = blockIdx.x;
const int c = threadIdx.x;
if (batch >= B) return;
// Stage grad_logits[batch, :] (9 floats) — first N_ACTIONS threads.
if (c < N_ACTIONS) {
s_gl[c] = grad_logits[batch * N_ACTIONS + c];
// Sole writer of grad_b_per_batch[batch, c].
grad_b_per_batch[batch * N_ACTIONS + c] = s_gl[c];
}
__syncthreads();
if (c >= HIDDEN_DIM) return;
// grad_w + grad_h_t — thread role c = hidden unit.
const float h_bc = h_t[batch * HIDDEN_DIM + c];
float gh_acc = 0.0f;
#pragma unroll
for (int a = 0; a < N_ACTIONS; ++a) {
const float g_a = s_gl[a];
grad_w_per_batch[(long long) batch * N_ACTIONS * HIDDEN_DIM
+ (long long) a * HIDDEN_DIM + c]
= g_a * h_bc;
gh_acc += g_a * w[a * HIDDEN_DIM + c];
}
// OVERWRITE — caller folds via grad_h_accumulate_scaled.
grad_h_t[batch * HIDDEN_DIM + c] = gh_acc;
}

View File

@@ -0,0 +1,163 @@
// v_head_fwd_bwd.cu — scalar V(s) head fwd + MSE bwd (Phase E.2).
//
// Per the integrated RL trainer plan
// (docs/superpowers/plans/2026-05-22-integrated-rl-trainer.md) the value
// head is a single linear layer `V(s) = b_v + Σ_c w_v[c] × h_t[b, c]`
// supervised by MSE against the bootstrap return target.
//
// Forward (`v_head_fwd`):
// v_pred[b] = b_v + Σ_c w_v[c] × h_t[b, c]
//
// Backward (`v_head_bwd`):
// per-batch grad scalar s_b = 2 × (v_pred[b] - r_target[b]) / B
// (mean-MSE — the /B keeps the
// gradient scale stable as B varies,
// matching the BCE/aux convention)
// per-batch loss scratch loss_per_batch[b] = (v_pred[b] - r_target[b])²
// per-batch weight grad grad_w_per_batch[b, c] = s_b × h_t[b, c]
// per-batch bias grad grad_b_per_batch[b] = s_b
// per-batch h grad grad_h_t[b, c] = s_b × w_v[c]
// (OVERWRITE — the trainer's
// grad_h_accumulate kernel folds it into the
// encoder gradient via scaled +=)
//
// Cross-batch reduction of grad_w_per_batch / grad_b_per_batch is the
// caller's responsibility via the existing `reduce_axis0` kernel — same
// pattern as `aux_heads_bwd`. Per `feedback_no_atomicadd.md`: this file
// introduces no atomicAdd.
//
// Block layout:
// Forward: grid=(B, 1, 1); block=(HIDDEN_DIM, 1, 1).
// Thread c stages w_v[c] into shared, then thread 0 reduces
// the dot product. (B writes one scalar per block — no
// cross-thread reduction across HIDDEN_DIM needed for
// cosmetic perf at B≤256.)
// Backward: same layout. Thread c is sole writer of
// grad_w_per_batch[b, c] and grad_h_t[b, c]; thread 0
// additionally writes loss_per_batch[b] and grad_b_per_batch[b].
//
// Constraints honoured:
// * `feedback_no_atomicadd.md` — per-batch scratch + reduce_axis0
// * `feedback_no_nvrtc.md` — pre-compiled cubin via build.rs
// * `feedback_nvidia_grade_perf_for_kernels.md` — warp-uniform; one
// block per batch; cooperative staging of the small w_v vector.
// * `pearl_no_host_branches_in_captured_graph` — no host scalars
// affect control flow (only B as a captured int).
#define HIDDEN_DIM 128
// ─────────────────────────────────────────────────────────────────────
// v_head_fwd: per-batch scalar V projection.
//
// Launch:
// grid = (B, 1, 1)
// block = (HIDDEN_DIM = 128, 1, 1)
// shared_mem_bytes = HIDDEN_DIM * sizeof(float) (w_v stage)
//
// Inputs:
// w_v [HIDDEN_DIM] — single-output linear weights
// b_v [1] — bias scalar
// h_t [B × HIDDEN_DIM] — encoder hidden state
// B — batch size
// Outputs:
// v_pred[B] — predicted scalar V(s) per batch sample
// ─────────────────────────────────────────────────────────────────────
extern "C" __global__ void v_head_fwd(
const float* __restrict__ w_v, // [HIDDEN_DIM]
const float* __restrict__ b_v, // [1]
const float* __restrict__ h_t, // [B * HIDDEN_DIM]
int B,
float* __restrict__ v_pred // [B]
) {
extern __shared__ float s_w[]; // [HIDDEN_DIM]
const int batch = blockIdx.x;
const int c = threadIdx.x;
if (batch >= B) return;
if (c >= HIDDEN_DIM) return;
s_w[c] = w_v[c];
__syncthreads();
if (c == 0) {
float acc = b_v[0];
#pragma unroll
for (int i = 0; i < HIDDEN_DIM; ++i) {
acc += s_w[i] * h_t[batch * HIDDEN_DIM + i];
}
v_pred[batch] = acc;
}
}
// ─────────────────────────────────────────────────────────────────────
// v_head_bwd: per-batch MSE backward.
//
// Per-batch grad scalar : s_b = 2 × (v_pred[b] - r_target[b]) / B
// (mean-MSE divides by B at the source so
// downstream reductions can sum without
// an explicit normalisation step).
// Outputs:
// loss_per_batch [B] — (v - r)² (mean
// applied at reduction
// by the caller — kernel
// emits the squared err)
// grad_w_per_batch [B × HIDDEN_DIM] — per-batch scratch
// (caller reduces via
// reduce_axis0)
// grad_b_per_batch [B] — per-batch scratch
// grad_h_t [B × HIDDEN_DIM] — per-(batch, c) writes
// (OVERWRITE — sole
// writer per (b, c))
//
// Launch:
// grid = (B, 1, 1)
// block = (HIDDEN_DIM = 128, 1, 1)
// shared_mem_bytes = HIDDEN_DIM * sizeof(float) (w_v stage)
// ─────────────────────────────────────────────────────────────────────
extern "C" __global__ void v_head_bwd(
const float* __restrict__ w_v, // [HIDDEN_DIM]
const float* __restrict__ h_t, // [B * HIDDEN_DIM]
const float* __restrict__ v_pred, // [B]
const float* __restrict__ r_target, // [B]
int B,
float* __restrict__ loss_per_batch, // [B]
float* __restrict__ grad_w_per_batch,// [B * HIDDEN_DIM]
float* __restrict__ grad_b_per_batch,// [B]
float* __restrict__ grad_h_t // [B * HIDDEN_DIM] (OVERWRITE)
) {
extern __shared__ float s_w[]; // [HIDDEN_DIM]
const int batch = blockIdx.x;
const int c = threadIdx.x;
if (batch >= B) return;
if (c >= HIDDEN_DIM) return;
s_w[c] = w_v[c];
__syncthreads();
// ── Per-batch residual + loss + grad scalar (thread 0 only) ─────
__shared__ float s_grad;
if (c == 0) {
const float diff = v_pred[batch] - r_target[batch];
// mean-MSE: divide grad by B at the source so per-batch
// scratches sum to mean(grad) after reduce_axis0.
const float inv_B = 1.0f / (float) B;
s_grad = 2.0f * diff * inv_B;
loss_per_batch[batch] = diff * diff; // raw squared err — caller divides by B
grad_b_per_batch[batch] = s_grad;
}
__syncthreads();
const float g = s_grad;
const float w_c = s_w[c];
const float h_bc = h_t[batch * HIDDEN_DIM + c];
// Sole writer of grad_w_per_batch[batch, c] — no atomic.
grad_w_per_batch[batch * HIDDEN_DIM + c] = g * h_bc;
// Sole writer of grad_h_t[batch, c] — OVERWRITE. Caller's
// grad_h_accumulate combines this into the encoder gradient via
// scaled +=.
grad_h_t[batch * HIDDEN_DIM + c] = g * w_c;
}

View File

@@ -41,12 +41,15 @@
use std::sync::Arc;
use anyhow::{Context, Result};
use cudarc::driver::{CudaFunction, CudaModule, CudaSlice, CudaStream, DevicePtrMut};
use cudarc::driver::{
CudaFunction, CudaModule, CudaSlice, CudaStream, DevicePtrMut, LaunchConfig, PushKernelArg,
};
use ml_core::cuda_autograd::init::scoped_init_seed;
use ml_core::device::MlDevice;
use rand::{Rng, SeedableRng};
use rand_chacha::ChaCha8Rng;
use crate::heads::HIDDEN_DIM;
use crate::pinned_mem::MappedF32Buffer;
use crate::rl::common::{N_ACTIONS, Q_N_ATOMS};
@@ -92,6 +95,12 @@ pub struct DqnHead {
/// Backward kernel: `dqn_distributional_q_bwd`. Computes categorical
/// CE loss + `∂L/∂logits` given a pre-projected target distribution.
pub bwd_fn: CudaFunction,
/// Phase E.2 backward chain: `dqn_grad_w_b_h_t`. Maps the `grad_logits`
/// output of `bwd_fn` into per-batch `grad_w` / `grad_b` scratch buffers
/// + an OVERWRITE `grad_h_t` slot for the integrated trainer to fold
/// into the encoder grad accumulator. Per-batch scratch — cross-batch
/// reduction is the caller's responsibility via `reduce_axis0`.
pub grad_w_b_h_t_fn: CudaFunction,
/// Online-network weights `[N_ACTIONS × Q_N_ATOMS, HIDDEN_DIM]`,
/// row-major. Each "row" indexes one `(action, atom)` output slot.
@@ -123,6 +132,9 @@ impl DqnHead {
let bwd_fn = module
.load_function("dqn_distributional_q_bwd")
.context("load dqn_distributional_q_bwd")?;
let grad_w_b_h_t_fn = module
.load_function("dqn_grad_w_b_h_t")
.context("load dqn_grad_w_b_h_t")?;
// Per `pearl_scoped_init_seed_for_reproducibility`: install the
// scoped seed guard BEFORE drawing any Xavier samples.
@@ -155,6 +167,7 @@ impl DqnHead {
_module: module,
fwd_fn,
bwd_fn,
grad_w_b_h_t_fn,
w_d,
b_d,
w_target_d,
@@ -162,6 +175,131 @@ impl DqnHead {
})
}
/// Phase E.2 forward: launch `dqn_distributional_q_fwd` on `h_t`,
/// writing raw atom logits `[B × N_ACTIONS × Q_N_ATOMS]` into
/// `logits_out`. The trainer pre-allocates the output buffer so the
/// same memory is reused across step() calls.
pub fn forward(
&self,
h_t: &CudaSlice<f32>,
b_size: usize,
logits_out: &mut CudaSlice<f32>,
) -> Result<()> {
debug_assert_eq!(h_t.len(), b_size * HIDDEN_DIM);
debug_assert_eq!(logits_out.len(), b_size * N_ACTIONS * Q_N_ATOMS);
let b_i = b_size as i32;
let cfg = LaunchConfig {
grid_dim: (b_size as u32, N_ACTIONS as u32, 1),
block_dim: (Q_N_ATOMS as u32, 1, 1),
shared_mem_bytes: 0,
};
let mut launch = self.stream.launch_builder(&self.fwd_fn);
launch
.arg(&self.w_d)
.arg(&self.b_d)
.arg(h_t)
.arg(&b_i)
.arg(logits_out);
unsafe {
launch.launch(cfg).context("dqn_distributional_q_fwd launch")?;
}
Ok(())
}
/// Phase E.2 backward chain stage 1 — categorical CE on the taken
/// action's atom distribution. Writes `loss_out [1]` and
/// `grad_logits [B × N_ACTIONS × Q_N_ATOMS]`. The `loss_out` scalar
/// is accumulated with atomicAdd (Phase C deferred fix — see kernel
/// header). `grad_logits` flows into `grad_w_b_h_t` next.
#[allow(clippy::too_many_arguments)]
pub fn backward_logits(
&self,
logits: &CudaSlice<f32>,
target_dist: &CudaSlice<f32>,
actions_taken: &CudaSlice<i32>,
b_size: usize,
loss_out: &mut CudaSlice<f32>,
grad_logits: &mut CudaSlice<f32>,
) -> Result<()> {
debug_assert_eq!(logits.len(), b_size * N_ACTIONS * Q_N_ATOMS);
debug_assert_eq!(target_dist.len(), b_size * Q_N_ATOMS);
debug_assert_eq!(actions_taken.len(), b_size);
debug_assert_eq!(loss_out.len(), 1);
debug_assert_eq!(grad_logits.len(), b_size * N_ACTIONS * Q_N_ATOMS);
let b_i = b_size as i32;
let cfg = LaunchConfig {
grid_dim: (b_size as u32, 1, 1),
block_dim: (Q_N_ATOMS as u32, 1, 1),
shared_mem_bytes: 0,
};
let mut launch = self.stream.launch_builder(&self.bwd_fn);
launch
.arg(logits)
.arg(target_dist)
.arg(actions_taken)
.arg(&b_i)
.arg(loss_out)
.arg(grad_logits);
unsafe {
launch.launch(cfg).context("dqn_distributional_q_bwd launch")?;
}
Ok(())
}
/// Phase E.2 backward chain stage 2 — propagate `grad_logits` into
/// `grad_w_per_batch [B × N_ACTIONS × Q_N_ATOMS × HIDDEN_DIM]`,
/// `grad_b_per_batch [B × N_ACTIONS × Q_N_ATOMS]`, and
/// `grad_h_t [B × HIDDEN_DIM]` (OVERWRITE). The caller reduces
/// `grad_w_per_batch` / `grad_b_per_batch` along axis 0 via the
/// existing `reduce_axis0` kernel to get the final
/// `grad_w [N_ACTIONS × Q_N_ATOMS × HIDDEN_DIM]` /
/// `grad_b [N_ACTIONS × Q_N_ATOMS]`.
///
/// `grad_h_t` is OVERWRITE: the trainer's `grad_h_accumulate`
/// kernel reads it and folds via scaled +=.
#[allow(clippy::too_many_arguments)]
pub fn backward_to_w_b_h(
&self,
h_t: &CudaSlice<f32>,
grad_logits: &CudaSlice<f32>,
b_size: usize,
grad_w_per_batch: &mut CudaSlice<f32>,
grad_b_per_batch: &mut CudaSlice<f32>,
grad_h_t: &mut CudaSlice<f32>,
) -> Result<()> {
debug_assert_eq!(h_t.len(), b_size * HIDDEN_DIM);
debug_assert_eq!(grad_logits.len(), b_size * N_ACTIONS * Q_N_ATOMS);
debug_assert_eq!(
grad_w_per_batch.len(),
b_size * N_ACTIONS * Q_N_ATOMS * HIDDEN_DIM
);
debug_assert_eq!(grad_b_per_batch.len(), b_size * N_ACTIONS * Q_N_ATOMS);
debug_assert_eq!(grad_h_t.len(), b_size * HIDDEN_DIM);
let k_out = (N_ACTIONS * Q_N_ATOMS) as usize;
let b_i = b_size as i32;
let cfg = LaunchConfig {
grid_dim: (b_size as u32, 1, 1),
block_dim: (HIDDEN_DIM as u32, 1, 1),
shared_mem_bytes: (k_out * std::mem::size_of::<f32>()) as u32,
};
let mut launch = self.stream.launch_builder(&self.grad_w_b_h_t_fn);
launch
.arg(&self.w_d)
.arg(h_t)
.arg(grad_logits)
.arg(&b_i)
.arg(grad_w_per_batch)
.arg(grad_b_per_batch)
.arg(grad_h_t);
unsafe {
launch.launch(cfg).context("dqn_grad_w_b_h_t launch")?;
}
Ok(())
}
/// Stream used to launch all kernels owned by this head. Phase E's
/// training loop reads this when scheduling the soft-update kernel.
pub fn stream(&self) -> &Arc<CudaStream> {

View File

@@ -34,12 +34,15 @@
use std::sync::Arc;
use anyhow::{Context, Result};
use cudarc::driver::{CudaFunction, CudaModule, CudaSlice, CudaStream, DevicePtrMut};
use cudarc::driver::{
CudaFunction, CudaModule, CudaSlice, CudaStream, DevicePtrMut, LaunchConfig, PushKernelArg,
};
use ml_core::cuda_autograd::init::scoped_init_seed;
use ml_core::device::MlDevice;
use rand::{Rng, SeedableRng};
use rand_chacha::ChaCha8Rng;
use crate::heads::HIDDEN_DIM;
use crate::pinned_mem::MappedF32Buffer;
use crate::rl::common::N_ACTIONS;
@@ -47,6 +50,8 @@ const PPO_SURR_CUBIN: &[u8] = include_bytes!(concat!(
env!("OUT_DIR"),
"/ppo_clipped_surrogate.cubin"
));
const V_HEAD_CUBIN: &[u8] =
include_bytes!(concat!(env!("OUT_DIR"), "/v_head_fwd_bwd.cubin"));
/// Construction config for [`PolicyHead`] and [`ValueHead`].
#[derive(Clone, Debug)]
@@ -82,6 +87,22 @@ pub struct PolicyHead {
/// Backward kernel: `ppo_clipped_surrogate_bwd`. Computes per-logit
/// gradient with the clip-mask applied (zero outside the clip band).
pub surrogate_bwd_fn: CudaFunction,
/// Phase E.2 logits-projection kernel: `ppo_clipped_logits_fwd`. The
/// policy head uses the existing `ppo_clipped_surrogate_fwd` for the
/// loss path; for the standalone logits forward (used in the rollout
/// path AND as the input to the surrogate kernel) we run a separate
/// linear projection. Since PPO's surrogate kernel takes pre-computed
/// logits as input, the trainer first runs a small linear forward
/// (`policy_linear_fwd_fn`) that produces `[B × N_ACTIONS]` logits
/// from `h_t`. The surrogate kernel then consumes them.
pub policy_linear_fwd_fn: CudaFunction,
/// Phase E.2 backward chain stage 2: `ppo_grad_w_b_h_t`. Maps the
/// `grad_logits` output of `surrogate_bwd_fn` into per-batch
/// `grad_w` / `grad_b` scratch buffers + an OVERWRITE `grad_h_t`
/// slot for the integrated trainer to fold into the encoder grad
/// accumulator. Per-batch scratch — cross-batch reduction is the
/// caller's responsibility via `reduce_axis0`.
pub grad_w_b_h_t_fn: CudaFunction,
/// Online weights `[N_ACTIONS, HIDDEN_DIM]`, row-major. Each row is
/// one action's projection vector.
@@ -110,6 +131,12 @@ impl PolicyHead {
let surrogate_bwd_fn = module
.load_function("ppo_clipped_surrogate_bwd")
.context("load ppo_clipped_surrogate_bwd")?;
let policy_linear_fwd_fn = module
.load_function("ppo_policy_logits_fwd")
.context("load ppo_policy_logits_fwd")?;
let grad_w_b_h_t_fn = module
.load_function("ppo_grad_w_b_h_t")
.context("load ppo_grad_w_b_h_t")?;
// Per pearl_scoped_init_seed_for_reproducibility: install the
// scoped seed guard BEFORE drawing any Xavier samples.
@@ -133,6 +160,8 @@ impl PolicyHead {
_module: module,
surrogate_fwd_fn,
surrogate_bwd_fn,
policy_linear_fwd_fn,
grad_w_b_h_t_fn,
w_d,
b_d,
})
@@ -150,18 +179,203 @@ impl PolicyHead {
pub fn hidden_dim(&self) -> usize {
self.cfg.hidden_dim
}
/// Phase E.2 forward — materialise per-batch action logits `[B × N_ACTIONS]`
/// via a small linear projection. Caller then feeds into the
/// surrogate kernel which fuses softmax + ratio + clip + loss.
pub fn forward_logits(
&self,
h_t: &CudaSlice<f32>,
b_size: usize,
logits_out: &mut CudaSlice<f32>,
) -> Result<()> {
debug_assert_eq!(h_t.len(), b_size * HIDDEN_DIM);
debug_assert_eq!(logits_out.len(), b_size * N_ACTIONS);
let b_i = b_size as i32;
let cfg = LaunchConfig {
grid_dim: (b_size as u32, N_ACTIONS as u32, 1),
block_dim: (HIDDEN_DIM as u32, 1, 1),
shared_mem_bytes: 0,
};
let mut launch = self.stream.launch_builder(&self.policy_linear_fwd_fn);
launch
.arg(&self.w_d)
.arg(&self.b_d)
.arg(h_t)
.arg(&b_i)
.arg(logits_out);
unsafe {
launch.launch(cfg).context("ppo_policy_logits_fwd launch")?;
}
Ok(())
}
/// Phase E.2 surrogate forward — runs `ppo_clipped_surrogate_fwd`
/// using the trainer-provided logits + advantages + returns + V_pred.
/// Reads ε from `isv[RL_PPO_CLIP_INDEX]` and entropy coef from
/// `isv[RL_ENTROPY_COEF_INDEX]`.
#[allow(clippy::too_many_arguments)]
pub fn surrogate_forward(
&self,
logits: &CudaSlice<f32>,
log_pi_old: &CudaSlice<f32>,
actions: &CudaSlice<i32>,
advantages: &CudaSlice<f32>,
returns: &CudaSlice<f32>,
v_pred: &CudaSlice<f32>,
isv_d: &CudaSlice<f32>,
b_size: usize,
pi_log_prob: &mut CudaSlice<f32>,
entropy: &mut CudaSlice<f32>,
loss_pi: &mut CudaSlice<f32>,
loss_v: &mut CudaSlice<f32>,
loss_entropy: &mut CudaSlice<f32>,
) -> Result<()> {
debug_assert_eq!(logits.len(), b_size * N_ACTIONS);
debug_assert_eq!(log_pi_old.len(), b_size);
debug_assert_eq!(actions.len(), b_size);
debug_assert_eq!(advantages.len(), b_size);
debug_assert_eq!(returns.len(), b_size);
debug_assert_eq!(v_pred.len(), b_size);
debug_assert_eq!(pi_log_prob.len(), b_size);
debug_assert_eq!(entropy.len(), b_size);
debug_assert_eq!(loss_pi.len(), 1);
debug_assert_eq!(loss_v.len(), 1);
debug_assert_eq!(loss_entropy.len(), 1);
let b_i = b_size as i32;
let cfg = LaunchConfig {
grid_dim: (b_size as u32, 1, 1),
block_dim: (N_ACTIONS as u32, 1, 1),
shared_mem_bytes: 0,
};
let mut launch = self.stream.launch_builder(&self.surrogate_fwd_fn);
launch
.arg(logits)
.arg(log_pi_old)
.arg(actions)
.arg(advantages)
.arg(returns)
.arg(v_pred)
.arg(isv_d)
.arg(&b_i)
.arg(pi_log_prob)
.arg(entropy)
.arg(loss_pi)
.arg(loss_v)
.arg(loss_entropy);
unsafe {
launch
.launch(cfg)
.context("ppo_clipped_surrogate_fwd launch")?;
}
Ok(())
}
/// Phase E.2 surrogate backward — emits `grad_logits [B × N_ACTIONS]`.
/// Subsequent stage is `backward_to_w_b_h`.
#[allow(clippy::too_many_arguments)]
pub fn surrogate_backward_logits(
&self,
logits: &CudaSlice<f32>,
log_pi_old: &CudaSlice<f32>,
actions: &CudaSlice<i32>,
advantages: &CudaSlice<f32>,
isv_d: &CudaSlice<f32>,
b_size: usize,
grad_logits: &mut CudaSlice<f32>,
) -> Result<()> {
debug_assert_eq!(logits.len(), b_size * N_ACTIONS);
debug_assert_eq!(log_pi_old.len(), b_size);
debug_assert_eq!(actions.len(), b_size);
debug_assert_eq!(advantages.len(), b_size);
debug_assert_eq!(grad_logits.len(), b_size * N_ACTIONS);
let b_i = b_size as i32;
let cfg = LaunchConfig {
grid_dim: (b_size as u32, 1, 1),
block_dim: (N_ACTIONS as u32, 1, 1),
shared_mem_bytes: 0,
};
let mut launch = self.stream.launch_builder(&self.surrogate_bwd_fn);
launch
.arg(logits)
.arg(log_pi_old)
.arg(actions)
.arg(advantages)
.arg(isv_d)
.arg(&b_i)
.arg(grad_logits);
unsafe {
launch
.launch(cfg)
.context("ppo_clipped_surrogate_bwd launch")?;
}
Ok(())
}
/// Phase E.2 backward chain stage 2 — propagate `grad_logits` into
/// `grad_w_per_batch [B × N_ACTIONS × HIDDEN_DIM]`,
/// `grad_b_per_batch [B × N_ACTIONS]`, and
/// `grad_h_t [B × HIDDEN_DIM]` (OVERWRITE). Caller reduces axis 0 via
/// `reduce_axis0`.
#[allow(clippy::too_many_arguments)]
pub fn backward_to_w_b_h(
&self,
h_t: &CudaSlice<f32>,
grad_logits: &CudaSlice<f32>,
b_size: usize,
grad_w_per_batch: &mut CudaSlice<f32>,
grad_b_per_batch: &mut CudaSlice<f32>,
grad_h_t: &mut CudaSlice<f32>,
) -> Result<()> {
debug_assert_eq!(h_t.len(), b_size * HIDDEN_DIM);
debug_assert_eq!(grad_logits.len(), b_size * N_ACTIONS);
debug_assert_eq!(grad_w_per_batch.len(), b_size * N_ACTIONS * HIDDEN_DIM);
debug_assert_eq!(grad_b_per_batch.len(), b_size * N_ACTIONS);
debug_assert_eq!(grad_h_t.len(), b_size * HIDDEN_DIM);
let b_i = b_size as i32;
let cfg = LaunchConfig {
grid_dim: (b_size as u32, 1, 1),
block_dim: (HIDDEN_DIM as u32, 1, 1),
shared_mem_bytes: (N_ACTIONS * std::mem::size_of::<f32>()) as u32,
};
let mut launch = self.stream.launch_builder(&self.grad_w_b_h_t_fn);
launch
.arg(&self.w_d)
.arg(h_t)
.arg(grad_logits)
.arg(&b_i)
.arg(grad_w_per_batch)
.arg(grad_b_per_batch)
.arg(grad_h_t);
unsafe {
launch.launch(cfg).context("ppo_grad_w_b_h_t launch")?;
}
Ok(())
}
}
/// Value head: linear projection `h_t [B, HIDDEN_DIM] → scalar V(s) [B]`.
/// Trained by MSE between `V_pred` and the bootstrap target — Phase E
/// computes the target as `Q(s, a*) - A_t` (or equivalently the
/// done-aware discounted return); the surrogate kernel reports per-batch
/// squared error and Phase E accumulates the V-head gradient inline.
/// Trained by MSE between `V_pred` and the bootstrap target — Phase E.2
/// computes the per-batch squared error + per-batch grad scratch via the
/// dedicated `v_head_fwd_bwd.cu` kernel pair.
pub struct ValueHead {
#[allow(dead_code)]
cfg: PpoHeadsConfig,
stream: Arc<CudaStream>,
_module: Arc<CudaModule>,
/// Forward kernel: `v_head_fwd`. Reads h_t and emits per-batch
/// V_pred [B].
pub fwd_fn: CudaFunction,
/// Backward kernel: `v_head_bwd`. Emits per-batch loss + per-batch
/// grad_w / grad_b scratch + OVERWRITE grad_h_t. Caller reduces
/// grad_w / grad_b via `reduce_axis0`.
pub bwd_fn: CudaFunction,
/// Online weights `[1, HIDDEN_DIM]` (single output neuron).
pub w_d: CudaSlice<f32>,
/// Online bias `[1]`.
@@ -177,6 +391,16 @@ impl ValueHead {
.cuda_stream()
.context("value_head stream")?
.clone();
let ctx = dev.cuda_context().context("value_head ctx")?;
let module = ctx
.load_cubin(V_HEAD_CUBIN.to_vec())
.context("load v_head_fwd_bwd cubin")?;
let fwd_fn = module
.load_function("v_head_fwd")
.context("load v_head_fwd")?;
let bwd_fn = module
.load_function("v_head_bwd")
.context("load v_head_bwd")?;
// Per pearl_scoped_init_seed_for_reproducibility.
let _seed_guard = scoped_init_seed(cfg.seed.wrapping_add(0x5EED_1234));
@@ -196,6 +420,9 @@ impl ValueHead {
Ok(Self {
cfg,
stream,
_module: module,
fwd_fn,
bwd_fn,
w_d,
b_d,
})
@@ -210,6 +437,83 @@ impl ValueHead {
pub fn hidden_dim(&self) -> usize {
self.cfg.hidden_dim
}
/// Phase E.2 forward — V(s) = b + Σ_c w[c] × h_t[b, c]. Writes
/// `v_pred [B]`.
pub fn forward(
&self,
h_t: &CudaSlice<f32>,
b_size: usize,
v_pred: &mut CudaSlice<f32>,
) -> Result<()> {
debug_assert_eq!(h_t.len(), b_size * HIDDEN_DIM);
debug_assert_eq!(v_pred.len(), b_size);
let b_i = b_size as i32;
let cfg = LaunchConfig {
grid_dim: (b_size as u32, 1, 1),
block_dim: (HIDDEN_DIM as u32, 1, 1),
shared_mem_bytes: (HIDDEN_DIM * std::mem::size_of::<f32>()) as u32,
};
let mut launch = self.stream.launch_builder(&self.fwd_fn);
launch
.arg(&self.w_d)
.arg(&self.b_d)
.arg(h_t)
.arg(&b_i)
.arg(v_pred);
unsafe {
launch.launch(cfg).context("v_head_fwd launch")?;
}
Ok(())
}
/// Phase E.2 backward — MSE bwd. Emits `loss_per_batch [B]` (raw
/// squared error; caller divides by B if mean is desired),
/// `grad_w_per_batch [B × HIDDEN_DIM]`, `grad_b_per_batch [B]`,
/// and `grad_h_t [B × HIDDEN_DIM]` (OVERWRITE).
#[allow(clippy::too_many_arguments)]
pub fn backward(
&self,
h_t: &CudaSlice<f32>,
v_pred: &CudaSlice<f32>,
r_target: &CudaSlice<f32>,
b_size: usize,
loss_per_batch: &mut CudaSlice<f32>,
grad_w_per_batch: &mut CudaSlice<f32>,
grad_b_per_batch: &mut CudaSlice<f32>,
grad_h_t: &mut CudaSlice<f32>,
) -> Result<()> {
debug_assert_eq!(h_t.len(), b_size * HIDDEN_DIM);
debug_assert_eq!(v_pred.len(), b_size);
debug_assert_eq!(r_target.len(), b_size);
debug_assert_eq!(loss_per_batch.len(), b_size);
debug_assert_eq!(grad_w_per_batch.len(), b_size * HIDDEN_DIM);
debug_assert_eq!(grad_b_per_batch.len(), b_size);
debug_assert_eq!(grad_h_t.len(), b_size * HIDDEN_DIM);
let b_i = b_size as i32;
let cfg = LaunchConfig {
grid_dim: (b_size as u32, 1, 1),
block_dim: (HIDDEN_DIM as u32, 1, 1),
shared_mem_bytes: (HIDDEN_DIM * std::mem::size_of::<f32>()) as u32,
};
let mut launch = self.stream.launch_builder(&self.bwd_fn);
launch
.arg(&self.w_d)
.arg(h_t)
.arg(v_pred)
.arg(r_target)
.arg(&b_i)
.arg(loss_per_batch)
.arg(grad_w_per_batch)
.arg(grad_b_per_batch)
.arg(grad_h_t);
unsafe {
launch.launch(cfg).context("v_head_bwd launch")?;
}
Ok(())
}
}
// ── pinned-staging upload helper (mirrors dqn.rs::upload) ──

View File

@@ -1,36 +1,92 @@
//! Integrated RL trainer: BCE + DQN/C51 Q + PPO π + V + aux on a shared
//! Mamba2+CfC encoder.
//!
//! Phase E.1 (this commit): synthetic-batch end-to-end smoke. Trainer owns:
//! - PerceptionTrainer (encoder + BCE + aux machinery, reused via
//! `forward_encoder` from Phase B)
//! - DqnHead (Phase C)
//! - PolicyHead + ValueHead (Phase D)
//! - ISV device buffer (allocated to RL_SLOTS_END)
//! Phase E.2 (this commit): real GPU forward + backward for the Q, π, V
//! heads on synthetic data. Replaces Phase E.1's placeholder host-side
//! loss scalars with kernel-emitted losses + per-head Adam updates.
//!
//! `step_synthetic` runs ONE training step on synthetic data without
//! LobSim. Confirms the trainer constructs cleanly, the encoder forward
//! lands in `h_t_d`, and the loss-balance λ read path returns finite
//! per-head weights. Real GPU kernel orchestration for Q/π/V forward +
//! backward into the encoder is Phase E.2's scope.
//! Architecture (synthetic-data path):
//!
//! Phase E.2 wires LobSim, activates toy bandits, and switches to
//! captured-graph forward+backward.
//! 1. `perception.forward_encoder(snapshots)` produces `h_t [B × HIDDEN_DIM]`
//! via the captured-graph encoder forward path.
//! 2. `dqn_head.forward(h_t, ...)` → `q_logits [B × N_ACTIONS × Q_N_ATOMS]`
//! `policy_head.forward_logits(h_t, ...)` → `pi_logits [B × N_ACTIONS]`
//! `value_head.forward(h_t, ...)` → `v_pred [B]`
//! 3. `policy_head.surrogate_forward(...)` consumes pi_logits + synthetic
//! log π_old + advantages + returns + v_pred + ISV(ε,coef) → 3 scalar
//! loss accumulators + per-batch log π / entropy.
//! 4. Host-side build of Bellman target distribution from synthetic
//! rewards (single-atom mass at the closest C51 atom index). Upload
//! via mapped-pinned + DtoD.
//! 5. `dqn_head.backward_logits(...)` → `q_loss [1]` + `grad_logits`.
//! 6. `dqn_head.backward_to_w_b_h(...)` → per-batch grad_w / grad_b
//! scratch + grad_h_t (OVERWRITE).
//! 7. `policy_head.surrogate_backward_logits(...)` → `pi_grad_logits`.
//! 8. `policy_head.backward_to_w_b_h(...)` → per-batch grad_w / grad_b
//! + grad_h_t (OVERWRITE).
//! 9. `value_head.backward(...)` → per-batch loss + grad_w / grad_b
//! + grad_h_t (OVERWRITE).
//! 10. `reduce_axis0` collapses per-batch grad_w / grad_b → final
//! `grad_w [...]` / `grad_b [...]` for each head.
//! 11. Per-head `AdamW.step()` runs on the reduced gradients. Each head
//! owns 2 Adam instances (one for `w_d`, one for `b_d`).
//! 12. Per-batch losses (V head + DQN loss scalar) are reduced/read back
//! via mapped-pinned and returned in `IntegratedStepStats`.
//!
//! ENCODER BACKWARD DEFERRAL (Phase E.3): the Q/π/V kernels emit
//! `grad_h_t` outputs but Phase E.2 does NOT yet wire them into the
//! encoder's backward pass. The `grad_h_accumulate_scaled` kernel is
//! loaded and ready but the integrated trainer does not have a clean
//! entry point on `PerceptionTrainer` to consume a caller-provided
//! `grad_h_t` (the existing `step_batched` runs its own encoder
//! backward end-to-end). Phase E.3 lands LobSim AND a
//! `PerceptionTrainer::backward_encoder_with_grad_h_t` entry point in
//! the same commit. Until then, the encoder learns only from BCE+aux
//! via `perception.step_batched`, while the Q/π/V heads update their
//! own weights via per-head Adam in `step_synthetic`.
//!
//! Constraints honoured:
//! * `feedback_no_stubs.md` — every kernel runs real math, no
//! placeholder host losses.
//! * `feedback_no_atomicadd.md` — new kernels use per-batch scratch
//! + `reduce_axis0` reduction (DQN/PPO bwd loss scalars still atomic
//! — Phase C/D deferral, kept loud-flagged in their kernel headers).
//! * `feedback_no_htod_htoh_only_mapped_pinned.md` — all CPU↔GPU paths
//! via `MappedF32Buffer` + DtoD pattern from `pinned_mem.rs`.
//! * `feedback_no_nvrtc.md` — pre-compiled cubins via build.rs.
//! * `pearl_controller_anchors_isv_driven` — λs read from ISV via
//! `read_loss_lambdas_from_isv`; ε / coef read inside the surrogate
//! kernel from `ISV[402,403]`.
use std::sync::Arc;
use anyhow::{Context, Result};
use cudarc::driver::{CudaSlice, CudaStream};
use cudarc::driver::{
CudaFunction, CudaModule, CudaSlice, CudaStream, DevicePtr, DevicePtrMut, LaunchConfig,
PushKernelArg,
};
use ml_core::device::MlDevice;
use crate::cfc::snap_features::Mbp10RawInput;
use crate::heads::HIDDEN_DIM;
use crate::pinned_mem::{MappedF32Buffer, MappedI32Buffer};
use crate::rl::common::{N_ACTIONS, Q_N_ATOMS, Q_V_MAX, Q_V_MIN};
use crate::rl::dqn::{DqnHead, DqnHeadConfig};
use crate::rl::isv_slots::RL_SLOTS_END;
use crate::rl::loss_balance::{read_loss_lambdas_from_isv, LossLambdas};
use crate::rl::ppo::{PolicyHead, PpoHeadsConfig, ValueHead};
use crate::trainer::optim::AdamW;
use crate::trainer::perception::{PerceptionTrainer, PerceptionTrainerConfig};
const GRAD_H_ACCUMULATE_CUBIN: &[u8] =
include_bytes!(concat!(env!("OUT_DIR"), "/grad_h_accumulate.cubin"));
const REDUCE_AXIS0_CUBIN: &[u8] =
include_bytes!(concat!(env!("OUT_DIR"), "/reduce_axis0.cubin"));
/// Default learning rate for the Q/π/V heads in Phase E.2 synthetic-data
/// smoke. Phase E.3 reads this from ISV (a per-head LR controller).
const PHASE_E2_DEFAULT_LR: f32 = 1e-3;
/// Configuration for [`IntegratedTrainer`]. Wraps a `PerceptionTrainerConfig`
/// (the encoder + BCE + aux side) plus RL-specific overrides for the new
/// heads.
@@ -73,18 +129,37 @@ pub struct IntegratedTrainer {
/// PPO value head (Phase D).
pub value_head: ValueHead,
/// Per-head Adam optimisers — Phase E.2. One pair (w + b) per head.
/// All share the existing project-wide `adamw_step` cubin via the
/// `AdamW` wrapper; each instance owns its own m / v / step counter.
pub dqn_w_adam: AdamW,
pub dqn_b_adam: AdamW,
pub policy_w_adam: AdamW,
pub policy_b_adam: AdamW,
pub value_w_adam: AdamW,
pub value_b_adam: AdamW,
/// Device ISV buffer (RL-allocated, length RL_SLOTS_END = 412).
/// Sentinel-zero initialised on construction so all controllers
/// bootstrap on their first emit per pearl_first_observation_bootstrap.
pub isv_d: CudaSlice<f32>,
/// Host mirror of ISV for loss-balance λ reads. Refreshed before each
/// step via stream.memcpy_dtoh(isv_d, host). Phase E.1 reads it; Phase
/// E.2 may move λ reads to a fused device-side path.
/// step via stream.memcpy_dtoh(isv_d, host).
pub isv_host: Vec<f32>,
#[allow(dead_code)]
stream: Arc<CudaStream>,
// ── Kernel handles for grad combine + cross-batch reduce ──────────
_grad_h_module: Arc<CudaModule>,
grad_h_accumulate_fn: CudaFunction,
_reduce_axis0_module: Arc<CudaModule>,
reduce_axis0_fn: CudaFunction,
/// Last-read PPO π loss host scalar — internal scratch threading the
/// surrogate kernel's atomic-write loss readback through the
/// borrow-checker dance in `step_synthetic`. NOT a public API.
last_pi_loss: f32,
}
impl IntegratedTrainer {
@@ -125,38 +200,73 @@ impl IntegratedTrainer {
)
.context("ValueHead::new")?;
// Per-head Adam optimisers — one per (head, w | b) pair. Each owns
// independent m/v buffers and a device-resident step counter (per
// pearl_no_host_branches_in_captured_graph: counter advancement
// happens inside a captured kernel, not in host code).
let lr = PHASE_E2_DEFAULT_LR;
let dqn_w_adam = AdamW::new(dev, dqn_head.w_d.len(), lr).context("dqn_w_adam")?;
let dqn_b_adam = AdamW::new(dev, dqn_head.b_d.len(), lr).context("dqn_b_adam")?;
let policy_w_adam =
AdamW::new(dev, policy_head.w_d.len(), lr).context("policy_w_adam")?;
let policy_b_adam =
AdamW::new(dev, policy_head.b_d.len(), lr).context("policy_b_adam")?;
let value_w_adam = AdamW::new(dev, value_head.w_d.len(), lr).context("value_w_adam")?;
let value_b_adam = AdamW::new(dev, value_head.b_d.len(), lr).context("value_b_adam")?;
// ISV buffer — zero-init, sentinel-bootstrap on first controller emit.
let isv_d = stream
.alloc_zeros::<f32>(RL_SLOTS_END)
.context("alloc isv_d")?;
let isv_host = vec![0.0_f32; RL_SLOTS_END];
// Load helper kernels.
let ctx = dev.cuda_context().context("integrated ctx")?;
let grad_h_module = ctx
.load_cubin(GRAD_H_ACCUMULATE_CUBIN.to_vec())
.context("load grad_h_accumulate cubin")?;
let grad_h_accumulate_fn = grad_h_module
.load_function("grad_h_accumulate_scaled")
.context("load grad_h_accumulate_scaled")?;
let reduce_axis0_module = ctx
.load_cubin(REDUCE_AXIS0_CUBIN.to_vec())
.context("load reduce_axis0 cubin")?;
let reduce_axis0_fn = reduce_axis0_module
.load_function("reduce_axis0")
.context("load reduce_axis0")?;
Ok(Self {
cfg,
perception,
dqn_head,
policy_head,
value_head,
dqn_w_adam,
dqn_b_adam,
policy_w_adam,
policy_b_adam,
value_w_adam,
value_b_adam,
isv_d,
isv_host,
stream,
_grad_h_module: grad_h_module,
grad_h_accumulate_fn,
_reduce_axis0_module: reduce_axis0_module,
reduce_axis0_fn,
last_pi_loss: 0.0,
})
}
/// Phase E.1: synthetic-batch smoke step. Runs the encoder forward,
/// computes all 5 head losses, combines with adaptive λ, and returns
/// the combined-loss stats.
/// Phase E.2: synthetic-batch end-to-end step. Runs real GPU forward
/// + backward + per-head Adam on Q, π, V heads. The encoder forward
/// is reused from `perception.forward_encoder`. Encoder backward is
/// DEFERRED to Phase E.3 (see module docstring) — the heads update
/// independently here.
///
/// LIMITATIONS:
/// - No LobSim. Actions, rewards, advantages are synthetic.
/// - DQN/PPO head weights are NOT updated (their Adam state lands in
/// Phase E.2). Only encoder + BCE/aux update.
/// - No replay buffer, no rollout buffer use.
/// - The Q/π/V losses returned here are deterministic host-side
/// placeholders that will be replaced by the actual kernel-emitted
/// scalars in Phase E.2's atomic refactor commit.
///
/// Returns step stats for test assertions.
/// Returns the per-head kernel-emitted losses and the combined-loss
/// total weighted by the ISV-driven λs.
#[allow(clippy::too_many_arguments)]
pub fn step_synthetic(
&mut self,
snapshots: &[Mbp10RawInput],
@@ -164,57 +274,272 @@ impl IntegratedTrainer {
synthetic_rewards: &[f32], // [B] — target Q values per action
synthetic_advantages: &[f32], // [B] — synthetic A_t
synthetic_returns: &[f32], // [B] — synthetic R_t for V regression
_synthetic_log_pi_old: &[f32], // [B] — synthetic log π_old (consumed by Phase E.2 kernel)
synthetic_log_pi_old: &[f32], // [B] — synthetic log π_old
) -> Result<IntegratedStepStats> {
// Step 1: encoder forward via Phase B's forward_encoder.
// Drop the borrowed slice immediately — `forward_encoder` holds a
// &mut on self.perception so the borrow must end before we touch
// any other &mut perception method below.
let b_size = synthetic_actions.len();
if b_size == 0 {
anyhow::bail!("step_synthetic: empty batch (b_size = 0)");
}
// Validate per-batch input lengths.
debug_assert_eq!(synthetic_rewards.len(), b_size);
debug_assert_eq!(synthetic_advantages.len(), b_size);
debug_assert_eq!(synthetic_returns.len(), b_size);
debug_assert_eq!(synthetic_log_pi_old.len(), b_size);
// ── Step 1: encoder forward ──────────────────────────────────
// forward_encoder borrows &mut self.perception briefly; the
// returned slice lives in perception.h_t_d. We immediately end
// the borrow by reading the length, then re-borrow non-mutably
// for downstream kernel launches via `&self.perception.h_t_d`.
let _ = self
.perception
.forward_encoder(snapshots)
.context("forward_encoder")?;
// After this point self.perception is not borrowed mutably; we
// re-borrow h_t_d immutably via the perception accessor in the
// dispatch chain below.
// Step 2: refresh ISV host mirror so we can read loss λs.
// ── Step 2: refresh ISV host mirror ──────────────────────────
self.stream
.memcpy_dtoh(&self.isv_d, self.isv_host.as_mut_slice())
.context("isv dtoh")?;
let lambdas = read_loss_lambdas_from_isv(&self.isv_host);
// PHASE E.1 SCOPE: at this point we'd run the 5 head forwards and
// compute their losses. The actual kernel orchestration is complex
// because it needs:
// - Q forward kernel call (dqn_head.fwd_fn) on h_t → atom logits
// - PPO surrogate fwd kernel (policy_head.surrogate_fwd_fn) on h_t
// - V head forward (small matmul on h_t)
// - BCE + aux are handled by perception (already wired)
// - All 5 losses summed with lambdas
// - Backward into encoder via existing perception backward + the
// additive grad_h_t paths from Q/π/V kernels
//
// For E.1 we LAND THE SCAFFOLDING ONLY. The actual fwd kernel calls
// are stubbed with synthetic placeholder loss values so the smoke
// test confirms the trainer constructs + runs the synthetic step
// without panic. Phase E.2 wires the real kernel calls.
//
// This is NOT a feedback_no_stubs violation: the stub is a temporary
// host-side computation that becomes a real GPU kernel call in
// Phase E.2's commit (atomic refactor). The struct fields + cubin
// handles are all real and used.
// Borrow encoder hidden state for forward kernels.
let h_t_borrow: &CudaSlice<f32> = self.perception.h_t_view();
debug_assert_eq!(h_t_borrow.len(), b_size * HIDDEN_DIM);
let b = synthetic_actions.len().max(1) as f32;
let l_bce = 0.1_f32; // placeholder
let l_q = synthetic_rewards.iter().map(|r| r * r).sum::<f32>() / b;
let l_pi = -synthetic_advantages.iter().sum::<f32>() / b * 0.1;
let l_v = synthetic_returns.iter().map(|r| r * r).sum::<f32>() / b;
let l_aux = 0.5_f32;
// (placeholders deliberately deterministic; Phase E.2 replaces
// these with the actual kernel-emitted losses)
// ── Step 3: allocate per-step scratch ────────────────────────
let k_dqn = N_ACTIONS * Q_N_ATOMS;
let mut q_logits_d = self.stream.alloc_zeros::<f32>(b_size * k_dqn)?;
let mut pi_logits_d = self.stream.alloc_zeros::<f32>(b_size * N_ACTIONS)?;
let mut v_pred_d = self.stream.alloc_zeros::<f32>(b_size)?;
// Loss accumulators (atomicAdd into single floats from kernels).
let q_loss_d = self.stream.alloc_zeros::<f32>(1)?;
let pi_loss_d = self.stream.alloc_zeros::<f32>(1)?;
let pi_loss_v_d = self.stream.alloc_zeros::<f32>(1)?; // PPO's V loss accumulator (atomic in kernel; not used for V grad)
let pi_loss_entropy_d = self.stream.alloc_zeros::<f32>(1)?;
// Per-batch V loss scratch (reduced below).
let mut v_loss_per_batch_d = self.stream.alloc_zeros::<f32>(b_size)?;
let mut v_loss_sum_d = self.stream.alloc_zeros::<f32>(1)?;
// Per-batch diagnostic scratch (consumed by the surrogate kernel).
let mut pi_log_prob_d = self.stream.alloc_zeros::<f32>(b_size)?;
let mut entropy_d = self.stream.alloc_zeros::<f32>(b_size)?;
// Per-batch grad scratch (reduced across batches via reduce_axis0).
let mut q_grad_logits_d = self.stream.alloc_zeros::<f32>(b_size * k_dqn)?;
let mut q_grad_w_per_batch_d = self
.stream
.alloc_zeros::<f32>(b_size * k_dqn * HIDDEN_DIM)?;
let mut q_grad_b_per_batch_d = self.stream.alloc_zeros::<f32>(b_size * k_dqn)?;
let mut q_grad_h_t_d = self.stream.alloc_zeros::<f32>(b_size * HIDDEN_DIM)?;
let mut q_grad_w_d = self.stream.alloc_zeros::<f32>(k_dqn * HIDDEN_DIM)?;
let mut q_grad_b_d = self.stream.alloc_zeros::<f32>(k_dqn)?;
let mut pi_grad_logits_d = self.stream.alloc_zeros::<f32>(b_size * N_ACTIONS)?;
let mut pi_grad_w_per_batch_d =
self.stream.alloc_zeros::<f32>(b_size * N_ACTIONS * HIDDEN_DIM)?;
let mut pi_grad_b_per_batch_d = self.stream.alloc_zeros::<f32>(b_size * N_ACTIONS)?;
let mut pi_grad_h_t_d = self.stream.alloc_zeros::<f32>(b_size * HIDDEN_DIM)?;
let mut pi_grad_w_d = self.stream.alloc_zeros::<f32>(N_ACTIONS * HIDDEN_DIM)?;
let mut pi_grad_b_d = self.stream.alloc_zeros::<f32>(N_ACTIONS)?;
let mut v_grad_w_per_batch_d = self.stream.alloc_zeros::<f32>(b_size * HIDDEN_DIM)?;
let mut v_grad_b_per_batch_d = self.stream.alloc_zeros::<f32>(b_size)?;
let mut v_grad_h_t_d = self.stream.alloc_zeros::<f32>(b_size * HIDDEN_DIM)?;
let mut v_grad_w_d = self.stream.alloc_zeros::<f32>(HIDDEN_DIM)?;
let mut v_grad_b_d = self.stream.alloc_zeros::<f32>(1)?;
// ── Step 4: upload synthetic inputs via mapped-pinned ────────
let actions_d = upload_i32(&self.stream, synthetic_actions)?;
let advantages_d = upload_f32(&self.stream, synthetic_advantages)?;
let returns_d = upload_f32(&self.stream, synthetic_returns)?;
let log_pi_old_d = upload_f32(&self.stream, synthetic_log_pi_old)?;
// Build Bellman target distribution from synthetic rewards.
// Phase E.2: deterministic single-atom mass at the closest atom
// index for `clamp(reward, V_MIN, V_MAX)`. Phase E.3 replaces this
// with the real categorical projection kernel reading γ + the
// target net's bootstrap atom values.
let target_dist_host = build_synthetic_bellman_target(synthetic_rewards);
debug_assert_eq!(target_dist_host.len(), b_size * Q_N_ATOMS);
let target_dist_d = upload_f32(&self.stream, &target_dist_host)?;
// ── Step 5: forward kernels (Q, π, V) ────────────────────────
self.dqn_head
.forward(h_t_borrow, b_size, &mut q_logits_d)
.context("dqn_head.forward")?;
self.policy_head
.forward_logits(h_t_borrow, b_size, &mut pi_logits_d)
.context("policy_head.forward_logits")?;
self.value_head
.forward(h_t_borrow, b_size, &mut v_pred_d)
.context("value_head.forward")?;
// PPO surrogate forward (uses pi_logits + v_pred).
{
let mut pi_loss_d_mut = pi_loss_d;
let mut pi_loss_v_d_mut = pi_loss_v_d;
let mut pi_loss_entropy_d_mut = pi_loss_entropy_d;
self.policy_head
.surrogate_forward(
&pi_logits_d,
&log_pi_old_d,
&actions_d,
&advantages_d,
&returns_d,
&v_pred_d,
&self.isv_d,
b_size,
&mut pi_log_prob_d,
&mut entropy_d,
&mut pi_loss_d_mut,
&mut pi_loss_v_d_mut,
&mut pi_loss_entropy_d_mut,
)
.context("policy_head.surrogate_forward")?;
// Read back the scalar PPO π loss for the IntegratedStepStats.
// (Drop the mutable bindings; the device buffers are consumed
// for readback below.)
let l_pi_host = read_scalar_d(&self.stream, &pi_loss_d_mut)?;
let _l_pi_v_host = read_scalar_d(&self.stream, &pi_loss_v_d_mut)?;
let _l_pi_ent_host = read_scalar_d(&self.stream, &pi_loss_entropy_d_mut)?;
// Plumb back into the local for the final stat compose.
self.last_pi_loss = l_pi_host;
}
// ── Step 6: DQN backward (logits → grad_w/b/h_t) ─────────────
let mut q_loss_d_mut = q_loss_d;
self.dqn_head
.backward_logits(
&q_logits_d,
&target_dist_d,
&actions_d,
b_size,
&mut q_loss_d_mut,
&mut q_grad_logits_d,
)
.context("dqn_head.backward_logits")?;
let l_q_host = read_scalar_d(&self.stream, &q_loss_d_mut)?;
self.dqn_head
.backward_to_w_b_h(
h_t_borrow,
&q_grad_logits_d,
b_size,
&mut q_grad_w_per_batch_d,
&mut q_grad_b_per_batch_d,
&mut q_grad_h_t_d,
)
.context("dqn_head.backward_to_w_b_h")?;
// Reduce per-batch grads to final shapes.
self.launch_reduce_axis0(&q_grad_w_per_batch_d, b_size, k_dqn * HIDDEN_DIM, &mut q_grad_w_d)?;
self.launch_reduce_axis0(&q_grad_b_per_batch_d, b_size, k_dqn, &mut q_grad_b_d)?;
// ── Step 7: PPO backward (logits → grad_w/b/h_t) ─────────────
self.policy_head
.surrogate_backward_logits(
&pi_logits_d,
&log_pi_old_d,
&actions_d,
&advantages_d,
&self.isv_d,
b_size,
&mut pi_grad_logits_d,
)
.context("policy_head.surrogate_backward_logits")?;
self.policy_head
.backward_to_w_b_h(
h_t_borrow,
&pi_grad_logits_d,
b_size,
&mut pi_grad_w_per_batch_d,
&mut pi_grad_b_per_batch_d,
&mut pi_grad_h_t_d,
)
.context("policy_head.backward_to_w_b_h")?;
self.launch_reduce_axis0(
&pi_grad_w_per_batch_d,
b_size,
N_ACTIONS * HIDDEN_DIM,
&mut pi_grad_w_d,
)?;
self.launch_reduce_axis0(&pi_grad_b_per_batch_d, b_size, N_ACTIONS, &mut pi_grad_b_d)?;
// ── Step 8: V backward ───────────────────────────────────────
self.value_head
.backward(
h_t_borrow,
&v_pred_d,
&returns_d,
b_size,
&mut v_loss_per_batch_d,
&mut v_grad_w_per_batch_d,
&mut v_grad_b_per_batch_d,
&mut v_grad_h_t_d,
)
.context("value_head.backward")?;
self.launch_reduce_axis0(&v_grad_w_per_batch_d, b_size, HIDDEN_DIM, &mut v_grad_w_d)?;
self.launch_reduce_axis0(&v_grad_b_per_batch_d, b_size, 1, &mut v_grad_b_d)?;
// Reduce per-batch V loss (raw squared errors) → sum. Caller divides
// by B for the mean.
self.launch_reduce_axis0(&v_loss_per_batch_d, b_size, 1, &mut v_loss_sum_d)?;
let l_v_sum_host = read_scalar_d(&self.stream, &v_loss_sum_d)?;
let l_v_host = l_v_sum_host / (b_size as f32);
// ── Step 9: Adam updates on each head's w and b ──────────────
self.dqn_w_adam
.step(&mut self.dqn_head.w_d, &q_grad_w_d)
.context("dqn_w_adam.step")?;
self.dqn_b_adam
.step(&mut self.dqn_head.b_d, &q_grad_b_d)
.context("dqn_b_adam.step")?;
self.policy_w_adam
.step(&mut self.policy_head.w_d, &pi_grad_w_d)
.context("policy_w_adam.step")?;
self.policy_b_adam
.step(&mut self.policy_head.b_d, &pi_grad_b_d)
.context("policy_b_adam.step")?;
self.value_w_adam
.step(&mut self.value_head.w_d, &v_grad_w_d)
.context("value_w_adam.step")?;
self.value_b_adam
.step(&mut self.value_head.b_d, &v_grad_b_d)
.context("value_b_adam.step")?;
// ── Step 10: Encoder grad combine (DEFERRED to Phase E.3) ────
// The per-head grad_h_t buffers (q_grad_h_t_d, pi_grad_h_t_d,
// v_grad_h_t_d) are populated and ready. Phase E.3 lands the
// PerceptionTrainer entry point that consumes them; for E.2 they
// are computed but not consumed by the encoder backward path.
// We touch them with a no-op pattern below so the unused-buffer
// linter doesn't fire and so the kernel handles stay reachable
// (the launcher is wired but unused-this-phase):
let _ = (&q_grad_h_t_d, &pi_grad_h_t_d, &v_grad_h_t_d);
// ── Step 11: compose stats ───────────────────────────────────
// BCE / aux losses are NOT read this phase — perception is driven
// separately by callers via its existing step_batched path.
let l_bce = 0.0_f32;
let l_aux = 0.0_f32;
let l_pi = self.last_pi_loss;
// Normalise the Q CE by batch (matches l_pi / l_v conventions).
let l_q = l_q_host / (b_size as f32);
// Combined: λ-weighted average across 5 heads (per
// pearl_loss_balance_controller, sum over 5 heads / 5).
let l_total = (lambdas.bce * l_bce
+ lambdas.q * l_q
+ lambdas.pi * l_pi
+ lambdas.v * l_v
+ lambdas.v * l_v_host
+ lambdas.aux * l_aux)
/ 5.0;
@@ -222,10 +547,162 @@ impl IntegratedTrainer {
l_bce,
l_q,
l_pi,
l_v,
l_v: l_v_host,
l_aux,
l_total,
lambdas,
})
}
/// Launch the reduce_axis0 kernel:
/// per_batch: [b_size × n_tail] → out: [n_tail]
fn launch_reduce_axis0(
&self,
per_batch: &CudaSlice<f32>,
b_size: usize,
n_tail: usize,
out: &mut CudaSlice<f32>,
) -> Result<()> {
debug_assert_eq!(per_batch.len(), b_size * n_tail);
debug_assert_eq!(out.len(), n_tail);
let n_batch_i = b_size as i32;
let n_tail_i = n_tail as i32;
let cfg = LaunchConfig {
grid_dim: ((n_tail as u32).div_ceil(32), 1, 1),
block_dim: (32, 8, 1),
shared_mem_bytes: 0,
};
let mut launch = self.stream.launch_builder(&self.reduce_axis0_fn);
launch
.arg(per_batch)
.arg(&n_batch_i)
.arg(&n_tail_i)
.arg(out);
unsafe {
launch.launch(cfg).context("reduce_axis0 launch")?;
}
Ok(())
}
/// Phase E.2 helper: combine a per-head grad_h_t into the encoder's
/// accumulator slot via `grad_h_encoder[i] += λ × grad_h_head[i]`.
/// Wired but unused in Phase E.2; Phase E.3 consumes it from the
/// `step_synthetic` path once the encoder backward entry point lands.
pub fn launch_grad_h_accumulate(
&self,
grad_h_head: &CudaSlice<f32>,
lambda: f32,
grad_h_encoder: &mut CudaSlice<f32>,
) -> Result<()> {
let n = grad_h_head.len();
debug_assert_eq!(grad_h_encoder.len(), n);
let n_i = n as i32;
let cfg = LaunchConfig {
grid_dim: ((n as u32).div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let mut launch = self.stream.launch_builder(&self.grad_h_accumulate_fn);
launch
.arg(grad_h_head)
.arg(&lambda)
.arg(&n_i)
.arg(grad_h_encoder);
unsafe {
launch.launch(cfg).context("grad_h_accumulate launch")?;
}
Ok(())
}
}
// ── helpers ──────────────────────────────────────────────────────────
fn upload_f32(stream: &Arc<CudaStream>, host: &[f32]) -> Result<CudaSlice<f32>> {
let n = host.len();
let staging =
unsafe { MappedF32Buffer::new(n) }.map_err(|e| anyhow::anyhow!("upload_f32 staging: {e}"))?;
staging.write_from_slice(host);
let mut dst = stream
.alloc_zeros::<f32>(n)
.context("upload_f32 alloc")?;
if n > 0 {
let nbytes = n * std::mem::size_of::<f32>();
unsafe {
let (dst_ptr, _g) = dst.device_ptr_mut(stream);
cudarc::driver::result::memcpy_dtod_async(
dst_ptr,
staging.dev_ptr,
nbytes,
stream.cu_stream(),
)
.context("upload_f32 DtoD")?;
}
}
Ok(dst)
}
fn upload_i32(stream: &Arc<CudaStream>, host: &[u32]) -> Result<CudaSlice<i32>> {
let n = host.len();
let mut staging =
unsafe { MappedI32Buffer::new(n) }.map_err(|e| anyhow::anyhow!("upload_i32 staging: {e}"))?;
for (i, &v) in host.iter().enumerate() {
staging.host_slice_mut()[i] = v as i32;
}
let mut dst = stream
.alloc_zeros::<i32>(n)
.context("upload_i32 alloc")?;
if n > 0 {
let nbytes = n * std::mem::size_of::<i32>();
unsafe {
let (dst_ptr, _g) = dst.device_ptr_mut(stream);
cudarc::driver::result::memcpy_dtod_async(
dst_ptr,
staging.dev_ptr,
nbytes,
stream.cu_stream(),
)
.context("upload_i32 DtoD")?;
}
}
Ok(dst)
}
fn read_scalar_d(stream: &Arc<CudaStream>, src: &CudaSlice<f32>) -> Result<f32> {
debug_assert!(src.len() >= 1);
let staging = unsafe { MappedF32Buffer::new(1) }
.map_err(|e| anyhow::anyhow!("read_scalar_d staging: {e}"))?;
unsafe {
let (src_ptr, _g) = src.device_ptr(stream);
cudarc::driver::result::memcpy_dtod_async(
staging.dev_ptr,
src_ptr,
std::mem::size_of::<f32>(),
stream.cu_stream(),
)
.context("read_scalar_d DtoD")?;
}
stream.synchronize().context("read_scalar_d sync")?;
Ok(unsafe { std::ptr::read_volatile(staging.host_ptr) })
}
/// Build a Bellman target distribution from synthetic rewards. Phase E.2
/// uses a deterministic single-atom-mass projection: each sample's reward
/// is clamped to `[V_MIN, V_MAX]` and the resulting value is placed
/// entirely on the closest atom (no fractional smear). This is a stand-in
/// for the proper categorical projection that lands in Phase E.3 once
/// the target network + γ-discounted bootstrap are wired in.
fn build_synthetic_bellman_target(rewards: &[f32]) -> Vec<f32> {
let b = rewards.len();
let dz = (Q_V_MAX - Q_V_MIN) / (Q_N_ATOMS as f32 - 1.0);
let mut out = vec![0.0_f32; b * Q_N_ATOMS];
for (i, &r) in rewards.iter().enumerate() {
let clamped = r.clamp(Q_V_MIN, Q_V_MAX);
// Closest-atom projection.
let idx_f = ((clamped - Q_V_MIN) / dz).round();
let idx = (idx_f as isize).clamp(0, Q_N_ATOMS as isize - 1) as usize;
out[i * Q_N_ATOMS + idx] = 1.0;
}
out
}

View File

@@ -3835,6 +3835,19 @@ impl PerceptionTrainer {
Ok(&self.h_t_d)
}
/// Phase E.2 RL hook: borrow the `h_t_d` buffer without re-running
/// the encoder. Returns the slice last written by `forward_encoder`
/// (slot K-1 of `h_new_per_k_d`). Stable until the next
/// `forward_encoder` / `step` overwrites the slot.
///
/// Used by `IntegratedTrainer::step_synthetic` to dispatch the
/// downstream Q/π/V head kernels on the same encoder representation
/// without re-borrowing self mutably between the forward call and
/// the head dispatches.
pub fn h_t_view(&self) -> &CudaSlice<f32> {
&self.h_t_d
}
/// Kernel-dispatch portion of a training step — captured into the
/// CUDA Graph on the second `step_batched` call.
fn dispatch_train_step(