docs: GPU-resident PER sum-tree design

Approved design for moving Prioritized Experience Replay entirely to
GPU — flat priority array with parallel prefix-sum sampling, GPU-resident
ring buffer for experiences, async loss readback. Eliminates the last
major CPU bottleneck in the DQN training loop.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-02 23:40:57 +01:00
parent 2f97402d5d
commit 995c473ad9

View File

@@ -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<Experience>`, 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<Mutex<ExperienceReplayBuffer>>),
Prioritized(Arc<PrioritizedReplayBuffer>),
GpuPrioritized(Arc<Mutex<GpuReplayBuffer>>), // NEW
}
```
### BatchSample extension
```rust
pub struct BatchSample {
pub experiences: Vec<Experience>, // empty when GPU path
pub weights: Vec<f32>, // empty when GPU path
pub indices: Vec<usize>, // empty when GPU path
pub gpu_batch: Option<GpuBatch>, // 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<Experience> → 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% |