From 766adbc718eeaf6aa70fdd2c565b88d3e3062ec8 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Tue, 26 May 2026 00:04:03 +0200 Subject: [PATCH] =?UTF-8?q?plan:=20GPU-resident=20PER=20+=20async=20diag?= =?UTF-8?q?=20=E2=80=94=207=20tasks,=20greenfield?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude Opus 4.7 --- .../plans/2026-05-26-gpu-per-async-diag.md | 530 ++++++++++++++++++ 1 file changed, 530 insertions(+) create mode 100644 docs/superpowers/plans/2026-05-26-gpu-per-async-diag.md diff --git a/docs/superpowers/plans/2026-05-26-gpu-per-async-diag.md b/docs/superpowers/plans/2026-05-26-gpu-per-async-diag.md new file mode 100644 index 000000000..3dfcf2dc3 --- /dev/null +++ b/docs/superpowers/plans/2026-05-26-gpu-per-async-diag.md @@ -0,0 +1,530 @@ +# GPU-Resident PER + Async Diag Implementation Plan + +> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** Replace host-side PER + blocking diag readback with all-GPU replay buffer and async non-blocking diagnostics. Target 100+ sps @ b=256 on L40S. + +**Architecture:** GPU-resident circular replay buffer with device-side sum-tree for priority sampling. Separate CUDA stream with double-buffered mapped-pinned memory for zero-stall diag. Old host ReplayBuffer, push_to_replay, sample_and_gather, NStepEntry, n_step_buffer, Transition struct, and --diag-every flag all deleted. + +**Tech Stack:** Rust 1.85, cudarc 0.19, CUDA 12.4, pre-compiled cubins + +--- + +## File Map + +| File | Responsibility | +|------|---------------| +| `crates/ml-alpha/src/rl/gpu_replay.rs` | NEW: GpuReplayBuffer struct, alloc, field accessors | +| `crates/ml-alpha/src/rl/mod.rs` | Add `pub mod gpu_replay;`, remove `pub mod replay;` | +| `crates/ml-alpha/src/rl/replay.rs` | DELETE entirely | +| `crates/ml-alpha/src/rl/common.rs` | DELETE `Transition` struct + `NStepEntry` | +| `crates/ml-alpha/src/trainer/diag_staging.rs` | NEW: DiagStaging double-buffer + async snapshot | +| `crates/ml-alpha/src/trainer/mod.rs` | Add `pub mod diag_staging;` | +| `crates/ml-alpha/src/trainer/integrated.rs` | Remove old PER code, wire new GPU PER + diag | +| `crates/ml-alpha/cuda/rl_per_push.cu` | NEW: push kernel (n-step + circular write) | +| `crates/ml-alpha/cuda/rl_per_sample.cu` | NEW: priority-weighted sample + gather | +| `crates/ml-alpha/cuda/rl_per_update_priority.cu` | NEW: write TD-errors to tree leaves | +| `crates/ml-alpha/cuda/rl_per_tree_rebuild.cu` | NEW: bottom-up parallel sum-tree rebuild | +| `crates/ml-alpha/cuda/rl_diag_snapshot.cu` | NEW: batch-copy diag fields to staging | +| `crates/ml-alpha/build.rs` | Register 5 new kernels | +| `crates/ml-alpha/examples/alpha_rl_train.rs` | Remove blocking readback, use DiagStaging | +| `crates/ml-alpha/tests/gpu_per_oracle.rs` | NEW: GPU oracle tests for PER kernels | +| `infra/k8s/argo/alpha-rl-template.yaml` | Remove `--diag-every` param | + +--- + +## Task 1: Delete old host-side PER + diag-every band-aid + +**Files:** +- Delete: `crates/ml-alpha/src/rl/replay.rs` +- Modify: `crates/ml-alpha/src/rl/mod.rs` (remove `pub mod replay;`) +- Modify: `crates/ml-alpha/src/rl/common.rs` (delete `Transition`, `NStepEntry`) +- Modify: `crates/ml-alpha/src/trainer/integrated.rs` (remove `replay` field, `push_to_replay`, `sample_and_gather`, `n_step_buffer`, `NStepEntry` usage) +- Modify: `crates/ml-alpha/examples/alpha_rl_train.rs` (remove `--diag-every`, remove `diag_every` gate, remove blocking readback loop) +- Modify: `infra/k8s/argo/alpha-rl-template.yaml` (remove `diag-every` param) + +- [ ] **Step 1: Delete `crates/ml-alpha/src/rl/replay.rs`** + +- [ ] **Step 2: Remove `pub mod replay;` from `crates/ml-alpha/src/rl/mod.rs`** + +- [ ] **Step 3: Delete `Transition` struct from `crates/ml-alpha/src/rl/common.rs`** + +The struct is at line 115. Delete it and any associated imports. + +- [ ] **Step 4: Remove from `integrated.rs`:** + +Delete these fields from `IntegratedTrainer`: +- `pub replay: crate::rl::replay::ReplayBuffer` (line ~761) +- `n_step_buffer: Vec>` (line ~762) + +Delete the `NStepEntry` struct (line ~349). + +Delete the constructor init for `replay` and `n_step_buffer`. + +Delete the entire `push_to_replay` method (~line 5730, ~110 lines). + +Delete the entire `sample_and_gather` method (~line 5861, ~80 lines). + +In `step_with_lobsim`, stub out the PER call sites with `todo!("GPU PER")` temporarily (will be replaced in Task 4). + +- [ ] **Step 5: Remove `--diag-every` from `alpha_rl_train.rs`** + +Delete the `diag_every` CLI field, the `if step % cli.diag_every` gate, and the closing brace. + +- [ ] **Step 6: Remove `diag-every` from Argo template** + +Delete the parameter definition and the `--diag-every {{...}}` arg line. + +- [ ] **Step 7: Verify compilation** + +```bash +SQLX_OFFLINE=true cargo check -p ml-alpha +``` + +Expected: compiles (with `todo!` in the PER call sites). + +- [ ] **Step 8: Commit** + +```bash +git commit -m "refactor(rl): delete host-side PER + diag-every band-aid (greenfield prep) + +Remove: ReplayBuffer, Transition, NStepEntry, push_to_replay, +sample_and_gather, n_step_buffer, --diag-every flag. +PER call sites stubbed with todo!() — replaced in next commits." +``` + +--- + +## Task 2: GpuReplayBuffer struct + device allocations + +**Files:** +- Create: `crates/ml-alpha/src/rl/gpu_replay.rs` +- Modify: `crates/ml-alpha/src/rl/mod.rs` + +- [ ] **Step 1: Create `gpu_replay.rs` with the struct and constructor** + +```rust +// crates/ml-alpha/src/rl/gpu_replay.rs +use std::sync::Arc; +use anyhow::{Context, Result}; +use cudarc::driver::{CudaSlice, CudaStream}; +use crate::heads::HIDDEN_DIM; + +pub struct GpuReplayBuffer { + pub capacity: usize, + pub b_size: usize, + + // Circular replay storage + pub h_t_d: CudaSlice, // [capacity × HIDDEN_DIM] + pub h_tp1_d: CudaSlice, // [capacity × HIDDEN_DIM] + pub scalars_d: CudaSlice, // [capacity × 7] (action, scaled_reward, raw_reward, done, log_pi_old, n_step_gamma, n_step_return_raw) + pub write_head_d: CudaSlice, // [1] + pub replay_len_d: CudaSlice, // [1] + pub max_priority_d: CudaSlice, // [1] + + // Priority sum-tree (leaves at [capacity..2×capacity)) + pub priority_tree_d: CudaSlice, // [2 × capacity] + + // N-step ring (per-batch) + pub nstep_scalars_d: CudaSlice, // [B × N_STEP_MAX × 3] (scaled_reward, raw_reward, done) + pub nstep_h_t_d: CudaSlice, // [B × N_STEP_MAX × HIDDEN_DIM] + pub nstep_write_idx_d: CudaSlice, // [B] + pub nstep_count_d: CudaSlice, // [B] + + // Sample output indices (for priority update after training) + pub sample_indices_d: CudaSlice, // [B] + + // PRNG state for sampling + pub sample_prng_d: CudaSlice, // [B] +} + +pub const N_STEP_MAX: usize = 16; +pub const SCALARS_PER_TRANSITION: usize = 7; + +impl GpuReplayBuffer { + pub fn new(stream: &Arc, capacity: usize, b_size: usize) -> Result { + Ok(Self { + capacity, + b_size, + h_t_d: stream.alloc_zeros(capacity * HIDDEN_DIM).context("replay h_t")?, + h_tp1_d: stream.alloc_zeros(capacity * HIDDEN_DIM).context("replay h_tp1")?, + scalars_d: stream.alloc_zeros(capacity * SCALARS_PER_TRANSITION).context("replay scalars")?, + write_head_d: stream.alloc_zeros(1).context("replay write_head")?, + replay_len_d: stream.alloc_zeros(1).context("replay len")?, + max_priority_d: stream.alloc_zeros(1).context("replay max_priority")?, + priority_tree_d: stream.alloc_zeros(2 * capacity).context("replay tree")?, + nstep_scalars_d: stream.alloc_zeros(b_size * N_STEP_MAX * 3).context("nstep scalars")?, + nstep_h_t_d: stream.alloc_zeros(b_size * N_STEP_MAX * HIDDEN_DIM).context("nstep h_t")?, + nstep_write_idx_d: stream.alloc_zeros(b_size).context("nstep write_idx")?, + nstep_count_d: stream.alloc_zeros(b_size).context("nstep count")?, + sample_indices_d: stream.alloc_zeros(b_size).context("sample indices")?, + sample_prng_d: stream.alloc_zeros(b_size).context("sample prng")?, + }) + } + + pub fn len_host_approx(&self) -> usize { + // For the log line only — not in hot path. + // Actual len is on device; this reads a stale cached value. + self.capacity // placeholder until wired + } +} +``` + +- [ ] **Step 2: Add `pub mod gpu_replay;` to `crates/ml-alpha/src/rl/mod.rs`** + +- [ ] **Step 3: Verify compilation** + +```bash +SQLX_OFFLINE=true cargo check -p ml-alpha +``` + +- [ ] **Step 4: Commit** + +```bash +git commit -m "feat(rl): GpuReplayBuffer struct — device-resident circular replay + sum-tree" +``` + +--- + +## Task 3: Write the 4 PER CUDA kernels + +**Files:** +- Create: `crates/ml-alpha/cuda/rl_per_push.cu` +- Create: `crates/ml-alpha/cuda/rl_per_sample.cu` +- Create: `crates/ml-alpha/cuda/rl_per_update_priority.cu` +- Create: `crates/ml-alpha/cuda/rl_per_tree_rebuild.cu` +- Modify: `crates/ml-alpha/build.rs` + +- [ ] **Step 1: Write `rl_per_push.cu`** + +Grid=(b_size), Block=(1). Per-thread: +1. Write scalars to nstep ring +2. DtoD h_t slice to nstep_h_t ring +3. If flush condition (count >= n_step OR done): + - Compute n-step return from ring + - Write oldest h_t → replay_h_t[write_head] + - Copy current h_tp1 → replay_h_tp1[write_head] + - Write scalars to replay_scalars[write_head] + - Set priority leaf = max_priority + - Advance write_head (mod capacity), bump replay_len (clamped) + - Reset ring for this batch + +Kernel signature: +```cuda +extern "C" __global__ void rl_per_push( + const float* __restrict__ h_t_current, // [B × HIDDEN_DIM] + const float* __restrict__ h_tp1_current, // [B × HIDDEN_DIM] + const float* __restrict__ rewards_scaled, // [B] + const float* __restrict__ rewards_raw, // [B] + const float* __restrict__ dones, // [B] + const float* __restrict__ log_pi_old, // [B] + const float* __restrict__ isv, + // Replay storage (circular) + float* __restrict__ replay_h_t, // [capacity × HIDDEN_DIM] + float* __restrict__ replay_h_tp1, // [capacity × HIDDEN_DIM] + float* __restrict__ replay_scalars, // [capacity × 7] + float* __restrict__ priority_tree, // [2 × capacity] + unsigned int* __restrict__ write_head, // [1] + unsigned int* __restrict__ replay_len, // [1] + float* __restrict__ max_priority, // [1] + // N-step ring (per-batch) + float* __restrict__ nstep_scalars, // [B × N_STEP_MAX × 3] + float* __restrict__ nstep_h_t, // [B × N_STEP_MAX × HIDDEN_DIM] + unsigned int* __restrict__ nstep_write_idx,// [B] + unsigned int* __restrict__ nstep_count, // [B] + int b_size, + int capacity, + int n_step_max +) +``` + +- [ ] **Step 2: Write `rl_per_sample.cu`** + +Grid=(b_size), Block=(1). Per-thread: +1. Segment = total_priority / b_size; draw u ~ U(b*seg, (b+1)*seg) +2. Walk sum-tree top-down to find leaf index +3. Write sample_indices[b] = leaf +4. Gather: copy replay data to sampled buffers +5. Compute IS weight + +```cuda +extern "C" __global__ void rl_per_sample( + const float* __restrict__ priority_tree, // [2 × capacity] + const float* __restrict__ replay_h_t, // [capacity × HIDDEN_DIM] + const float* __restrict__ replay_h_tp1, // [capacity × HIDDEN_DIM] + const float* __restrict__ replay_scalars, // [capacity × 7] + const unsigned int* __restrict__ replay_len,// [1] + unsigned int* __restrict__ prng_state, // [B] + // Outputs + float* __restrict__ sampled_h_t, // [B × HIDDEN_DIM] + float* __restrict__ sampled_h_tp1, // [B × HIDDEN_DIM] + float* __restrict__ sampled_rewards, // [B] + float* __restrict__ sampled_raw_rewards, // [B] + float* __restrict__ sampled_dones, // [B] + float* __restrict__ sampled_log_pi_old, // [B] + float* __restrict__ sampled_n_step_gammas, // [B] + int* __restrict__ sampled_actions, // [B] + unsigned int* __restrict__ sample_indices, // [B] + float* __restrict__ is_weights, // [B] + const float* __restrict__ isv, + int b_size, + int capacity +) +``` + +- [ ] **Step 3: Write `rl_per_update_priority.cu`** + +Grid=(b_size), Block=(1). Per-thread: +1. Read sample_indices[b], compute new priority |td_error[b]|^α + ε +2. Write to priority_tree[capacity + leaf] +3. Thread 0: reduce max priority across batch, update max_priority[0] + +```cuda +extern "C" __global__ void rl_per_update_priority( + const unsigned int* __restrict__ sample_indices, // [B] + const float* __restrict__ td_errors, // [B] + float* __restrict__ priority_tree, // [2 × capacity] + float* __restrict__ max_priority, // [1] + const float* __restrict__ isv, + int b_size, + int capacity +) +``` + +- [ ] **Step 4: Write `rl_per_tree_rebuild.cu`** + +Bottom-up rebuild, no atomics. Single kernel, multiple passes. +Grid=(min(capacity/2, 256)), Block=(256). For each level from bottom to root: +``` +for level in (0..log2(capacity)).rev(): + nodes_at_level = capacity >> (level + 1) + for i in tid..nodes_at_level step gridDim*blockDim: + node = (1 << level) + i + tree[node] = tree[2*node] + tree[2*node+1] + __threadfence() // device-wide visibility before next level +``` + +```cuda +extern "C" __global__ void rl_per_tree_rebuild( + float* __restrict__ priority_tree, // [2 × capacity] + int capacity +) +``` + +- [ ] **Step 5: Register all 4 in `build.rs`** + +```rust +"rl_per_push", +"rl_per_sample", +"rl_per_update_priority", +"rl_per_tree_rebuild", +``` + +- [ ] **Step 6: Verify CUDA compilation** + +```bash +for f in rl_per_push rl_per_sample rl_per_update_priority rl_per_tree_rebuild; do + nvcc -cubin -arch sm_86 -O3 crates/ml-alpha/cuda/${f}.cu -o /tmp/${f}.cubin && echo "$f OK" || echo "$f FAIL" +done +``` + +- [ ] **Step 7: Verify Rust compilation** + +```bash +SQLX_OFFLINE=true cargo check -p ml-alpha +``` + +- [ ] **Step 8: Commit** + +```bash +git commit -m "feat(cuda): GPU-resident PER kernels — push, sample, update_priority, tree_rebuild" +``` + +--- + +## Task 4: Wire GPU PER into IntegratedTrainer + +**Files:** +- Modify: `crates/ml-alpha/src/trainer/integrated.rs` + +- [ ] **Step 1: Add GpuReplayBuffer field + kernel handles** + +Replace old `replay` field with `gpu_replay: GpuReplayBuffer`. Add cubin consts, module/fn fields for all 4 PER kernels. + +- [ ] **Step 2: Init in constructor** + +```rust +let gpu_replay = GpuReplayBuffer::new(&stream, per_capacity, b_size)?; +``` + +Load all 4 cubins and functions. + +- [ ] **Step 3: Replace push_to_replay call site** + +In `step_with_lobsim`, replace the `todo!()` for push with a launch of `rl_per_push`: + +```rust +// Launch rl_per_push — replaces host push_to_replay entirely. +{ + let cfg = LaunchConfig { grid_dim: (b_size as u32, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 }; + let mut launch = self.stream.launch_builder(&self.rl_per_push_fn); + launch + .arg(self.perception.h_t_view()) + .arg(&self.h_tp1_d) + .arg(&self.rewards_d) + .arg(&self.raw_rewards_d) + .arg(&self.dones_d) + .arg(&self.log_pi_old_d) + .arg(&self.isv_d) + .arg(&mut self.gpu_replay.h_t_d) + // ... all other args from kernel signature + ; + unsafe { launch.launch(cfg).context("rl_per_push")?; } +} +``` + +- [ ] **Step 4: Replace sample_and_gather call site** + +In the K-loop, replace `todo!()` with launch of `rl_per_sample`: + +```rust +{ + let cfg = LaunchConfig { grid_dim: (b_size as u32, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 }; + let mut launch = self.stream.launch_builder(&self.rl_per_sample_fn); + launch + .arg(&self.gpu_replay.priority_tree_d) + .arg(&self.gpu_replay.h_t_d) + // ... writes directly into self.sampled_h_t_d etc. + ; + unsafe { launch.launch(cfg).context("rl_per_sample")?; } +} +``` + +- [ ] **Step 5: Add priority update after training step** + +After `step_synthetic`, launch `rl_per_update_priority` then `rl_per_tree_rebuild`: + +```rust +// Update priorities from TD errors computed in step_synthetic +{ + let cfg = LaunchConfig { grid_dim: (b_size as u32, 1, 1), block_dim: (1, 1, 1), shared_mem_bytes: 0 }; + // ... rl_per_update_priority launch +} +{ + let cfg = LaunchConfig { grid_dim: (128, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 }; + // ... rl_per_tree_rebuild launch +} +``` + +- [ ] **Step 6: Verify compilation + smoke test** + +```bash +SQLX_OFFLINE=true cargo check -p ml-alpha +SQLX_OFFLINE=true FOXHUNT_TEST_DATA=test_data/futures-baseline cargo test -p ml-alpha --test integrated_trainer_smoke -- --ignored --nocapture +``` + +- [ ] **Step 7: Commit** + +```bash +git commit -m "feat(rl): wire GPU PER — push/sample/update/rebuild all device-side" +``` + +--- + +## Task 5: Async DiagStaging (separate stream + double-buffer) + +**Files:** +- Create: `crates/ml-alpha/src/trainer/diag_staging.rs` +- Modify: `crates/ml-alpha/src/trainer/mod.rs` +- Modify: `crates/ml-alpha/examples/alpha_rl_train.rs` + +- [ ] **Step 1: Create DiagStaging struct** + +```rust +// crates/ml-alpha/src/trainer/diag_staging.rs +use std::sync::Arc; +use anyhow::{Context, Result}; +use cudarc::driver::{CudaStream, DevicePtr, DevicePtrMut}; +use crate::pinned_mem::MappedF32Buffer; +use crate::rl::isv_slots::RL_SLOTS_END; + +pub struct DiagStaging { + stream: Arc, // separate diag stream + bufs: [MappedF32Buffer; 2], // double-buffer + current: usize, // 0 or 1 + n_floats: usize, + pub isv_offset: usize, + pub rewards_offset: usize, + pub dones_offset: usize, + pub actions_offset: usize, +} + +impl DiagStaging { + pub fn new(dev: &MlDevice, b_size: usize) -> Result { ... } + pub fn snapshot_async(&mut self, isv_d, rewards_d, dones_d, actions_d, train_stream) -> Result<()> { ... } + pub fn read_previous(&self) -> &[f32] { ... } // zero-cost host read of previous step +} +``` + +- [ ] **Step 2: Wire into alpha_rl_train.rs** + +Replace all `read_slice_d` calls with reads from `diag_staging.read_previous()`. +Call `diag_staging.snapshot_async(...)` once per step (non-blocking). + +- [ ] **Step 3: Remove the `if step % diag_every` gate entirely** + +Every step gets diag at zero cost. + +- [ ] **Step 4: Verify compilation + smoke** + +- [ ] **Step 5: Commit** + +```bash +git commit -m "feat(diag): async non-blocking diag via separate stream + double-buffered staging" +``` + +--- + +## Task 6: GPU oracle tests for PER kernels + +**Files:** +- Create: `crates/ml-alpha/tests/gpu_per_oracle.rs` + +- [ ] **Step 1: Write 4 oracle tests** + +1. `per_push_writes_to_replay` — push b=2 transitions, read back replay_scalars, verify written correctly +2. `per_sample_uniform_when_equal_priority` — set all leaves to 1.0, rebuild tree, sample → indices should cover the buffer uniformly +3. `per_tree_rebuild_correct_sums` — set known leaf values, rebuild, verify tree[1] = sum of all leaves +4. `per_update_priority_sets_leaves` — sample, set TD errors, verify tree leaves updated + +- [ ] **Step 2: Run on local GPU** + +```bash +SQLX_OFFLINE=true cargo test -p ml-alpha --test gpu_per_oracle -- --ignored --nocapture +``` + +- [ ] **Step 3: Commit** + +```bash +git commit -m "test(rl): GPU oracle tests for PER push/sample/rebuild/update" +``` + +--- + +## Task 7: Benchmark b=256 + validate throughput + +- [ ] **Step 1: Push all changes** +- [ ] **Step 2: Submit 5k-step benchmark at b=256** +- [ ] **Step 3: Verify sps ≥ 80 (target 100)** +- [ ] **Step 4: Submit 500k production run if pass** + +--- + +## Kill Criteria + +- Task 1: compiles with `todo!()` stubs, no old PER code remains +- Task 3: all 4 kernels compile with nvcc sm_86 +- Task 4: smoke test passes (single step exercises push + sample) +- Task 5: sps unchanged with/without JSONL writing +- Task 6: all oracle tests pass on local RTX +- Task 7: ≥80 sps @ b=256 on L40S