Hybrid approach based on H100 measurements: - aux_child: grouped GEMM for IQL high+low, 2 streams for IQN+Attention - forward_child: backward branch parallelism (forward already multi-stream) - Target: 4.0s/step → 2.0s/step Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
136 lines
6.1 KiB
Markdown
136 lines
6.1 KiB
Markdown
# Phase 2: Multi-Stream Graph Parallelism
|
||
|
||
## Goal
|
||
|
||
Reduce per-step GPU time from 4.0s to ~2.0s by parallelizing independent operations
|
||
within the aux_child and forward_child graphs. Target: ~350s epochs (from ~718s).
|
||
|
||
## Measured Baseline (2026-04-18, H100 PCIe, batch=16384)
|
||
|
||
```
|
||
Per-step graph replay breakdown (last batch of epoch 1):
|
||
spectral = 0.2ms (<1%)
|
||
forward = 130.3ms (31%) ← branches already multi-stream, backward is NOT
|
||
ddqn = 1.0ms (<1%)
|
||
aux = 295.5ms (69%) ← 4 independent trainers running serial
|
||
adam = TBD (event fixed, awaiting measurement)
|
||
total = ~427ms measured + adam
|
||
```
|
||
|
||
Forward branches already use 4 streams with fork-join events (`batched_forward.rs:503-567`).
|
||
Backward branches are sequential (`batched_backward.rs` — no branch_stream infrastructure).
|
||
|
||
## Constraints (from NVIDIA research)
|
||
|
||
1. **SM saturation**: At B=16384, IQL GEMMs (M=128, N=16384, K=128) produce ~128 CTAs
|
||
filling 97% of 132 SMs. Multi-stream GEMM-vs-GEMM parallelism provides minimal benefit.
|
||
2. **Element-wise overlap**: Small kernels (bias_add, SiLU, reduce — 10-20µs) leave SMs idle.
|
||
A concurrent GEMM from another trainer CAN fill those gaps.
|
||
3. **Workspace isolation**: Each concurrent cublasLtMatmul stream needs its own workspace (32MB).
|
||
Same `cublasLtHandle_t` is safe (stateless per-call).
|
||
4. **Grouped GEMM**: `cublasLtMatmulGroupedBatched` packs multiple independent GEMMs with
|
||
different M/N/K into ONE kernel launch — eliminates inter-kernel gaps.
|
||
|
||
## Design
|
||
|
||
### 2A. Aux Child: Hybrid Parallelism
|
||
|
||
**IQL high + IQL low** have identical GEMM shapes (same architecture, different weights).
|
||
Use grouped GEMM to fuse both trainers' forward+backward into single kernel launches.
|
||
This halves IQL kernel count (16 GEMMs → 8 grouped GEMMs).
|
||
|
||
**IQN + Attention** run on 2 separate streams, forked from the main aux stream after
|
||
IQL completes (IQL must run first because IQL advantage weights are read by CQL later).
|
||
|
||
```
|
||
aux_child capture (RELAXED mode):
|
||
main_stream: HER relabel → EMA target → IQL grouped GEMM (fwd+bwd+adam) → [FORK]
|
||
│
|
||
iqn_stream: ─── IQN (fwd+bwd+adam) ──────────────────────────────────────────┤
|
||
attn_stream: ─── Attention (fwd+bwd+adam) ────────────────────────────────────┤
|
||
│
|
||
main_stream: ─── [JOIN] → CQL gradient → ensemble → recursive_conf_bwd → regime_scale
|
||
```
|
||
|
||
**Workspace**: 2 extra 32MB workspaces (IQN stream + Attention stream). Total: 64MB additional.
|
||
IQL grouped GEMM uses the existing shared workspace (sequential on main stream).
|
||
|
||
**Streams**: 2 new `CudaStream` objects, forked from the main stream context.
|
||
Pre-allocated events: `iql_done_event`, `iqn_done_event`, `attn_done_event`.
|
||
|
||
**Expected speedup on aux_child**:
|
||
- IQL: ~50% faster (grouped GEMM eliminates 8 kernel launch gaps)
|
||
- IQN + Attention: overlap with each other (~40% faster on combined time)
|
||
- CQL + ensemble + conf_bwd: unchanged (serial, depends on IQL/IQN/Attention outputs)
|
||
- Overall aux: ~295ms → ~180ms
|
||
|
||
### 2B. Forward Child: Backward Branch Parallelism
|
||
|
||
The forward pass already uses 4 branch streams. The backward pass does NOT.
|
||
Add the same fork-join pattern to the backward: after the trunk backward (serial),
|
||
fork 4 branch backward chains to their respective streams.
|
||
|
||
```
|
||
backward in forward_child:
|
||
main_stream: d_logits_clamp → [4 branch dW + dX GEMMs sequential] → trunk backward
|
||
|
||
↓ becomes:
|
||
|
||
main_stream: d_logits_clamp → [FORK]
|
||
│
|
||
branch_stream[0]: ── branch0 dW + dX ──┤
|
||
branch_stream[1]: ── branch1 dW + dX ──┤
|
||
branch_stream[2]: ── branch2 dW + dX ──┤
|
||
branch_stream[3]: ── branch3 dW + dX ──┤
|
||
│
|
||
main_stream: ──── [JOIN] → trunk backward (d_h_s2 reduction → d_h_s1 → d_states)
|
||
```
|
||
|
||
**Reuses existing infrastructure**: `branch_streams[4]`, `trunk_done_event`, `branch_done_events[4]`
|
||
from `batched_forward.rs`. The backward just needs to use them too.
|
||
|
||
**Expected speedup on forward_child**: ~130ms → ~100ms (backward branches ~30ms of the 130ms)
|
||
|
||
### 2C. Grouped GEMM for IQL
|
||
|
||
Use `cublasLtMatmulGroupedBatched` (cuBLAS 12.5+) to fuse IQL high-tau and IQL low-tau
|
||
GEMMs. Both trainers have identical shapes:
|
||
- Forward: 3 GEMMs (h1, h2, v_out) × 2 trainers = 6 GEMMs → 3 grouped GEMMs
|
||
- Backward: 5 GEMMs (dW3, dh2, dW2, dh1, dW1) × 2 trainers = 10 GEMMs → 5 grouped GEMMs
|
||
- Adam: 1 kernel × 2 trainers = 2 kernels → 1 grouped kernel (or keep separate)
|
||
|
||
Each grouped GEMM takes arrays of pointers (A_ptrs, B_ptrs, C_ptrs) and sizes.
|
||
The cuBLAS runtime dispatches both GEMMs in a single kernel launch.
|
||
|
||
**Fallback**: If `cublasLtMatmulGroupedBatched` is not available (requires cuBLAS 12.5+),
|
||
fall back to sequential execution (current behavior). Check at runtime.
|
||
|
||
## VRAM Budget
|
||
|
||
| Component | Current | Phase 2 |
|
||
|-----------|---------|---------|
|
||
| IQN workspace | shared (0 extra) | 32MB dedicated |
|
||
| Attention workspace | shared (0 extra) | 32MB dedicated |
|
||
| Branch streams | already allocated | 0 extra |
|
||
| Aux streams (2) | 0 | ~1KB (stream objects) |
|
||
| Aux events (3) | 0 | ~1KB |
|
||
| **Total extra** | | **64MB** |
|
||
|
||
On H100 with 80GB VRAM, 64MB is <0.1%.
|
||
|
||
## Files Modified
|
||
|
||
- `crates/ml/src/cuda_pipeline/batched_backward.rs` — add branch stream fork-join
|
||
- `crates/ml/src/cuda_pipeline/gpu_iql_trainer.rs` — grouped GEMM support
|
||
- `crates/ml/src/cuda_pipeline/shared_cublas_handle.rs` — aux workspace allocation
|
||
- `crates/ml/src/trainers/dqn/fused_training.rs` — aux_child multi-stream capture, RELAXED mode
|
||
|
||
## Success Criteria
|
||
|
||
- `GPU child graph breakdown` shows aux < 200ms (from 295ms)
|
||
- `GPU child graph breakdown` shows forward < 110ms (from 130ms)
|
||
- Epoch time < 400s (from ~718s)
|
||
- All 899 unit tests pass
|
||
- No cuStreamSynchronize in training loop
|
||
- Zero atomicAdd on gradient paths (maintained)
|