diff --git a/docs/plans/2026-03-02-gpu-per-sumtree-design.md b/docs/plans/2026-03-02-gpu-per-sumtree-design.md new file mode 100644 index 000000000..8f21dbc73 --- /dev/null +++ b/docs/plans/2026-03-02-gpu-per-sumtree-design.md @@ -0,0 +1,170 @@ +# GPU-Resident Prioritized Experience Replay + +**Date**: 2026-03-02 +**Status**: Approved +**Goal**: Eliminate the last major CPU bottleneck in the DQN training loop by moving PER sampling, priority updates, and experience storage entirely to GPU. + +## Problem + +The DQN training loop achieves 31s/epoch on L40S after the GPU experience collector fix, but PER remains CPU-bound: + +- **Sampling**: CPU segment tree O(log n) per sample, sequential +- **Priority update**: TD errors transferred GPU→CPU via `to_vec1()`, then CPU scatter-update +- **Experience gather**: CPU builds `Vec`, converts to tensors, transfers CPU→GPU +- **Loss logging**: `loss.to_scalar()` forces GPU sync every training step + +These CPU round-trips add 3-5s/epoch on L40S and prevent full GPU saturation. + +## Approach + +**GPU-Resident Flat Priority Array + Parallel Prefix Sum.** Replace the CPU segment tree with GPU-native parallel primitives. All PER operations (sample, update, IS weights) execute on GPU with zero CPU involvement per training step. + +## Architecture + +### New struct: `GpuReplayBuffer` + +Owns all experience data and priorities on GPU: + +``` +GPU memory layout (contiguous, pre-allocated): + states: [capacity × state_dim] f32 + next_states: [capacity × state_dim] f32 + actions: [capacity] u32 + rewards: [capacity] f32 + dones: [capacity] u8 + priorities: [capacity] f32 + write_cursor: usize (CPU-side index) + size: usize (CPU-side, capped at capacity) +``` + +Fixed-capacity ring buffer. No allocations during training. + +### ReplayBufferType extension + +```rust +pub enum ReplayBufferType { + Uniform(Arc>), + Prioritized(Arc), + GpuPrioritized(Arc>), // NEW +} +``` + +### BatchSample extension + +```rust +pub struct BatchSample { + pub experiences: Vec, // empty when GPU path + pub weights: Vec, // empty when GPU path + pub indices: Vec, // empty when GPU path + pub gpu_batch: Option, // NEW: pre-built GPU tensors +} +``` + +`compute_gradients()` checks `gpu_batch` first. If present, uses GPU tensors directly. Falls back to CPU path when `None`. + +## Data Flow — Hot Path + +### Before (current) + +``` +CPU: PER sum-tree sample → Vec → build Tensor → memcpy to GPU → forward/backward +GPU: loss → to_scalar() → CPU: compute |td| → update_priorities on CPU sum-tree +``` + +### After + +``` +GPU: curand random values → prefix-sum → binary search → gather experiences → IS weights +GPU: forward → backward → TD errors → |td|^alpha → scatter-update priorities +GPU (async, secondary stream): loss_sum → host pinned buffer (non-blocking) +CPU: reads loss value at epoch boundary only (already available, no sync) +``` + +Zero CPU-GPU synchronization per training step. Only sync point: epoch boundary checkpoint save. + +## Sampling Kernels + +### Proportional (~1 CUDA launch) + +1. `curand_generate_uniform` fills B random values in `[0, total_sum)` +2. `cub::DeviceScan::InclusiveSum` builds cumulative priority array (or maintained incrementally) +3. Custom kernel: each thread binary-searches cumulative array → index +4. Gather kernel: `states[indices]`, `actions[indices]`, etc. → output batch tensors +5. IS weight kernel: `weight_i = (1/(N * p_i / total))^beta / max_weight` + +### Rank-based (~2 CUDA launches) + +1. `cub::DeviceRadixSort::SortPairs` sorts (priority, index) pairs descending +2. Custom kernel: rank probabilities `1/rank^alpha`, cumulative sum, sample, gather, IS weights + +### Priority update (piggybacks on backward pass) + +TD errors already on GPU after `compute_gradients()`: +- Scatter kernel: `priority[idx] = |td_error[i]|^alpha + epsilon` +- Incremental cumulative sum update: O(B log N) parallel, B = batch size + +## Experience Insertion + +GPU experience collector output (already on GPU) → `gpu_buffer.insert_batch()`: +- `cuMemcpyDtoD` (GPU→GPU, no CPU) +- New experiences get `max_priority` (ensures sampling) +- Ring buffer wraps: `write_cursor = (write_cursor + count) % capacity` + +When `use_per: false`, same GPU ring buffer with uniform sampling (`curand` indices into `[0, size)`). + +## Loss Logging — Async Readback + +Instead of `loss.to_scalar()` per step (forced GPU sync): +- Accumulate losses as GPU tensor across accumulation steps +- Single `mean()` at end of accumulation +- `cuMemcpyDtoHAsync` into pinned host buffer on secondary stream +- CUDA event signals completion +- CPU reads at epoch boundary — already available, no stall + +## VRAM Budget + +For 100K buffer, 51-dim states, 45 actions: + +| Component | Size | +|-----------|------| +| States + next_states | 38.8 MB | +| Actions + rewards + dones | 0.9 MB | +| Priorities | 0.4 MB | +| Prefix-sum + sort workspace | 6 MB | +| Sampling buffers | 0.5 MB | +| **Total** | **~47 MB** | + +<0.1% of L40S (48GB), <0.06% of H100 (80GB). At 1M buffer: ~470 MB, still trivial. + +## Integration Points + +Three call sites in `trainer.rs`: + +1. **Buffer creation** (init): `GpuPrioritized(GpuReplayBuffer::new(...))` when CUDA available +2. **Experience insertion** (per-epoch): GPU collector → `gpu_buffer.insert_batch(gpu_tensors)` (GPU-to-GPU) +3. **Priority update** (per batch): `gpu_buffer.update_priorities_gpu(indices_tensor, td_errors_tensor)` (no transfer) + +## Error Handling & Fallback + +- **OOM at init**: Fall back to `Prioritized` (CPU PER), log warning, training continues +- **CUDA kernel failure**: Same fallback — degrade to CPU PER for remainder +- **Ring buffer overflow**: Wraps. Old experiences overwritten. Priority implicitly replaced. +- **NaN/Inf priorities**: Clamp to `[epsilon, max_priority]` in scatter kernel + +## Testing + +- **Unit**: Insert/sample/update cycle, verify correct indices and IS weights +- **Integration**: Full training loop with `GpuPrioritized`, verify loss decreases +- **Correctness**: KS test — GPU vs CPU PER sampling distributions on 10K samples +- **VRAM**: Max-size buffer allocation on CI (L40S) +- **Fallback**: Force CUDA failure → verify graceful degradation to CPU PER + +## Expected Impact + +| Metric | L40S (current) | L40S (after) | H100 (after) | +|--------|---------------|-------------|-------------| +| PER sample time | 3-5s/epoch | <0.1s/epoch | <0.05s/epoch | +| Priority update | ~1s/epoch | 0 (piggybacked) | 0 | +| Loss sync stalls | ~0.5s/epoch | 0 (async) | 0 | +| **Epoch time** | **31s** | **~26s** | **~9-10s** | +| GPU idle % | ~15% | ~2% | ~1% |