# GPU Performance Optimization Batch 1 — Design Spec **Date**: 2026-06-26 **Hardware target**: NVIDIA RTX 5060 Ti (Blackwell SM_120, 16 GB VRAM) **Research basis**: FlashAttention-3 (arXiv:2407.08608), MOSS FP8 (arXiv:2511.05811), SnapKV (arXiv:2404.14469), KeepKV (arXiv:2504.09936), PyTorch 2025 state survey --- ## Summary Four high-impact GPU performance optimizations, selected after cross-referencing the current rustytorch codebase against PyTorch 2025 and 2024–2025 arxiv papers. Items are ordered by implementation cost: CUDA Graphs is a half-day wiring task (infrastructure already exists), FP8 training is ~1 week, SnapKV/prefix caching is ~1 week, and FlashAttention-3 is the deepest CUDA kernel work at ~2 weeks. | Item | Expected gain | Complexity | Primary files | |------|--------------|-----------|---------------| | 1. CUDA Graphs | 5–20% throughput | Low (wiring only) | `training_loop.rs`, `cuda_graph.rs` | | 2. FP8 Training | 30–40% throughput, 50% memory | Medium | `fp8_cast.cu`, `fp8_gemm.rs`, `linear.rs` | | 3. FlashAttention-3 | 1.5–2× attention throughput | High | `flash_attention_v3_*.cu`, `backend_selector.rs` | | 4. SnapKV + Prefix Cache | 50–70% KV reduction | Medium | `eviction.rs`, `paged_kv_cache.rs` | --- ## Item 1: CUDA Graphs ### Problem Every GPU kernel launch carries 5–20 μs of CPU-side overhead (CUDA API call → driver → scheduler). A typical transformer training step launches hundreds of kernels. At small batch sizes common on consumer GPUs, this overhead is measurable. ### Solution CUDA Graphs capture the full sequence of GPU operations in a graph, then replay it in a single `cudaGraphLaunch` — eliminating per-kernel CPU overhead on all subsequent steps. ### Current state `CudaGraphManager` at `crates/core/rtx-runtime/src/cuda_graph.rs` (445 lines) is fully implemented with `begin_capture()`, `end_capture()`, `launch()`, and thread-safety via `RwLock`. It is **not connected to the training loop**. ### Design **`crates/training/rtx-transformers/src/training/config.rs`** — add two fields: ```rust pub enable_cuda_graphs: bool, // default: false pub cuda_graph_warmup_iters: usize, // default: 3 ``` **`crates/training/rtx-transformers/src/training/training_loop.rs`** (around line 167) — wrap the training step: - Iterations 0..warmup: execute normally (warms up allocator, avoids capturing allocations in graph) - Iteration `warmup`: `begin_capture(stream)` → `execute_step()` → `end_capture()` → store `graph_id` - Iterations `warmup+1..`: `graph_manager.launch(graph_id, stream)` — no per-kernel CPU calls **Constraint**: tensor shapes must be static across steps. Add a `requires_static_shapes()` check that disables graph mode and emits a warning when dynamic padding is detected (variable sequence lengths). **`TrainingLoop`** gets a new `graph_state: Option` field and a reference to `CudaGraphManager`. ### Verification ``` cargo test -p rtx-transformers -- test_cuda_graph_training_step ``` Test: capture a 2-layer transformer step, verify output tensors match non-graph execution within 1e-6, measure step time reduction. --- ## Item 2: FP8 Training (E4M3 forward / E5M2 gradients) ### Problem BF16 mixed-precision training uses 2 bytes per parameter activation. Blackwell's 5th-generation tensor cores support FP8 (1 byte), but the connection between rustytorch's existing FP8 type infrastructure and actual GEMM execution is missing. ### Current state - `DType::FP8E4M3` / `FP8E5M2` defined in `dtype.rs` with correct `size_bytes = 1` - `Fp8GradScaler` in `fp8_scaler.rs` — full amax-tracking, scale growth/backoff, overflow detection - `CudaDataType::R8F_E4M3 = 28` / `R8F_E5M2 = 29` in `cusparelt/types.rs` - MX block-wise quantization kernels in `microscaling.rs` - **Missing**: FP8 casting CUDA kernels, cuBLASLt GEMM dispatch, training loop wiring ### Design **Step 1 — Casting kernels** (new file: `crates/core/rtx-tensor/src/cuda_kernels/fp8_cast.cu`): ```cuda // Saturating cast: clamp to [-448, 448] before encoding E4M3 __global__ void cast_f32_to_fp8_e4m3( const float* __restrict__ src, uint8_t* __restrict__ dst, float scale, int n); __global__ void cast_bf16_to_fp8_e4m3( const __nv_bfloat16* __restrict__ src, uint8_t* __restrict__ dst, float scale, int n); __global__ void cast_fp8_e4m3_to_bf16( const uint8_t* __restrict__ src, __nv_bfloat16* __restrict__ dst, float inv_scale, int n); // E5M2 range ±57344, used for gradient tensors __global__ void cast_fp8_e5m2_to_bf16( const uint8_t* __restrict__ src, __nv_bfloat16* __restrict__ dst, float inv_scale, int n); ``` Follow `element_wise.cu` pattern for `build.rs` registration and Rust FFI wrapper. **Step 2 — cuBLASLt GEMM dispatcher** (new file: `crates/core/rtx-tensor/src/fp8_gemm.rs`): ```rust pub fn fp8_matmul( a: &CudaSlice, // E4M3, shape [M, K] b: &CudaSlice, // E4M3, shape [K, N] scale_a: f32, scale_b: f32, m: usize, n: usize, k: usize, stream: &CudaStream, ) -> Result> // BF16 output, shape [M, N] ``` Uses `cublasLtMatmul` with `CUBLAS_COMPUTE_32F` accumulator and `CUDA_R_8F_E4M3` A/B types. Add `cargo:rustc-link-lib=cublasLt` to `rtx-tensor/build.rs`. **Step 3 — Linear layer FP8 path** (`crates/training/rtx-transformers/src/layers/linear.rs`): - Add `fp8_mode: bool` to `LinearConfig` (default false) - In `forward()` when `fp8_mode`: 1. Cast weight matrix W from BF16 → E4M3 using `cast_bf16_to_fp8_e4m3` (can be pre-cached per step) 2. Cast activation X from BF16 → E4M3 3. Run `fp8_matmul(W_fp8, X_fp8, scale_w, scale_x, ...)` 4. Cast output Y from internal FP32 accumulator → BF16 - Master weights remain BF16/FP32 for optimizer stability **Step 4 — Training config** (`config.rs`): ```rust pub fp8_training: bool, // default: false pub fp8_format: Fp8Format, // default: E4M3Fwd_E5M2Grad ``` **Mixed precision recipe** (matching NVIDIA MOSS / Transformer Engine): - Forward pass: weights and activations cast to E4M3 - Gradient tensors: cast to E5M2 (wider range ±57344 handles gradient spikes) - Optimizer states: remain FP32 (master weights) - Scaling: `Fp8GradScaler` (existing) manages per-tensor amax and scale factors ### Verification ``` cargo test -p rtx-transformers -- test_fp8_forward_backward ``` Test: 2-layer linear model, FP8 forward output matches BF16 within 1e-2 tolerance, backward gradients are non-NaN and converge in 10 steps. --- ## Item 3: FlashAttention-3 Blackwell ### Problem Current kernels use WMMA (4th-generation tensor core API, 16×16×16 fragments) and `cp.async` for shared memory loads. Blackwell's SM_120 introduces Warpgroup MMA (WGMMA, 64×16×16 fragments — 4× compute density) and TMA (Tensor Memory Accelerator — asynchronous 2D tile loads with no warp involvement). FlashAttention-3 (Dao et al., arXiv:2407.08608) showed 1.5–2× speedup on H100 by exploiting these; the same techniques apply to SM_120. ### Current state - `flash_attention_forward.cu`: WMMA + tiled online softmax + cooperative groups - `utils.cu`: `wmma::fragment` with BF16, `cg::memcpy_async`, 163KB shared memory - `build.rs`: sm_120, PTX 8.0, `--use_fast_math --maxrregcount=255` - `FlashCudaKernels::new()` in `simple.rs`: drop-in kernel registration ### Design **New file: `src/kernels/flash_attention_v3_forward.cu`** Three algorithmic changes from FA2: **1. WGMMA (Warpgroup Matrix Multiply Accumulate)** Replace: ```cuda wmma::mma_sync(acc_frag, a_frag, b_frag, acc_frag); // 16×16×16 ``` With PTX inline asm for warpgroup-level operation (128 threads = 4 warps): ```cuda asm volatile( "wgmma.mma_async.sync.aligned.m64n128k16.f32.bf16.bf16 " "{%0,%1,...,%7}, [%8], {%9,%10,...}, 1, 1, 1, 0, 0;" : "+f"(d[0]), "+f"(d[1]), ... : "l"(desc_a), "r"(b[0]), "r"(b[1]), ... ); ``` Tile size: 64×128×16 per warpgroup vs 16×16×16 per warp. ~4× compute per instruction, amortizing instruction issue overhead. **2. TMA (Tensor Memory Accelerator)** Replace `cg::memcpy_async` with `cp.async.bulk.tensor.2d`: ```cuda // One instruction loads an entire 2D tile — no per-element loop, no warp threads consumed asm volatile( "cp.async.bulk.tensor.2d.shared::cluster.global.mbarrier::complete_tx::bytes" " [%0], [%1, {%2, %3}], [%4];" :: "r"(smem_ptr), "l"(gmem_base), "r"(row), "r"(col), "r"(mbar_ptr) ); ``` TMA runs in the memory subsystem — warp threads are free to compute while load is in flight. **3. Warp Specialization (Producer/Consumer split)** Split the 4 warpgroups per block: - **Producer warpgroup** (warpgroup 0): issues TMA loads for Q/K/V tiles, waits on TMA barrier, signals consumer via shared memory barrier - **Consumer warpgroups** (warpgroups 1–3): wait on producer signal, execute WGMMA accumulation, write output - Overlap: consumer processes tile `i` while producer loads tile `i+1` into the opposite shared-memory ping-pong buffer This eliminates the memory stall that currently serializes compute and memory in FA2. **Optional FP8 path** (gated by `use_fp8_attn: bool`): Before WGMMA, quantize each 16×16 Q/K block to E4M3 with a per-block scale factor. Enables ~2× additional throughput on FP8 tensor cores when combined with Item 2. **New file: `src/kernels/flash_attention_v3_backward.cu`** Same WGMMA + TMA + warp specialization for dQ/dK/dV, following the FA3 paper's recomputation approach (recompute softmax from stored LSE values to avoid storing the full N×N matrix). **`src/backend_selector.rs`** — runtime SM selection: ```rust fn select_flash_attention_kernel(sm_major: u32, sm_minor: u32) -> FlashKernelVariant { if sm_major >= 12 { FlashKernelVariant::V3Wgmma // Blackwell SM_120+ } else if sm_major >= 9 { FlashKernelVariant::V3Wgmma // Hopper SM_90 } else { FlashKernelVariant::V2Wmma // Ampere and below } } ``` Both kernel sets remain registered; no API change to callers. ### Verification ``` cargo test -p rtx-flash-attention -- test_v3_correctness cargo bench -p rtx-flash-attention ``` Correctness: FA3 output matches naive O(N²) attention within 1e-3 for sequence lengths 512, 1024, 4096. Performance: FA3 ≥ 1.5× throughput vs FA2 on SM_120 for seqlen ≥ 1024. --- ## Item 4: SnapKV + Prefix Caching ### Problem On 16 GB VRAM, KV cache is the primary constraint on batch size and context length during inference. Two complementary techniques address this: - **SnapKV** (Zhang et al., arXiv:2404.14469): prune low-attention keys during prefill, keeping 60% + recent window - **Prefix caching**: reuse KV pages for sequences sharing a common prefix (system prompt, few-shot examples) ### Current state - `PagedKvCache` in `paged_kv_cache.rs`: 3-tier storage (GPU/CPU/NVMe), UUID page IDs, copy-on-write, LRU eviction - `eviction.rs`: LRU/LFU/FIFO/Random policies — no attention-score policy - `config.rs`: `enable_prefix_sharing: bool = false` — no implementation behind the flag ### Design **SnapKV — attention-score eviction** (`eviction.rs`): Add new policy variant: ```rust pub enum EvictionPolicy { Lru, Lfu, Fifo, Random, AttentionScore { keep_ratio: f32, recent_window: usize }, // NEW } ``` Add sidecar to `PagedKvCache`: ```rust attention_scores: Arc>>, ``` During prefill, the flash attention kernel accumulates the sum of attention weights per key position into a host-accessible buffer. After prefill, scores are written to `attention_scores` via a kernel that reduces along the head dimension. In `select_victims()`: ```rust // Sort candidate pages by accumulated attention score ascending // Always retain last `recent_window` tokens (recency bias) // Evict bottom (1 - keep_ratio) fraction by score ``` This approach matches SnapKV's "observation pooling" (using cumulative attention across all query positions as importance proxy). **Prefix caching** (`paged_kv_cache.rs`): Add to `PagedKvCache`: ```rust prefix_index: HashMap>, ``` Hash function (Zobrist-style, collision probability ~1/2^64): ```rust fn compute_prefix_hash(tokens: &[u32]) -> u64 { // Precomputed random table: HASH_TABLE[pos % 1024][token % 65536] tokens.iter().enumerate().fold(0u64, |acc, (pos, &tok)| { acc ^ ZOBRIST_TABLE[pos % 1024][(tok % 65536) as usize] }) } ``` On each new request in the scheduler: 1. Compute prefix hash for first `N` tokens (where N is the longest aligned page boundary) 2. Look up `prefix_index[hash]` — if found, reuse those pages with copy-on-write 3. Allocate new pages only for tokens beyond the prefix On page eviction: - Remove matching entry from `prefix_index` - If CoW refcount > 1: decrement, do not free **Enable by default** (`InferenceConfig`): ```rust pub enable_prefix_sharing: bool = true, // was false pub snapkv_keep_ratio: f32 = 0.6, pub snapkv_recent_window: usize = 32, ``` ### Verification ``` cargo test -p rtx-inference -- test_prefix_cache_sharing cargo test -p rtx-inference -- test_snapkv_eviction ``` - Prefix test: two 2048-token requests sharing a 128-token system prompt must reuse identical pages (verified by page ID equality), saving ≥ 6% page allocations. - SnapKV test: after prefill on a 2048-token sequence with keep_ratio=0.6, active page count ≤ 40% of input length (allowing ≥ recent_window retained pages). --- ## Execution Order ``` Week 1: Item 1 (CUDA Graphs) — 0.5 day Item 2a–2b (FP8 casting kernels + cuBLASLt) — 3 days Week 2: Item 2c–2e (FP8 linear + training config + tests) — 3 days Item 4a (SnapKV eviction) — 1 day Week 3: Item 4b–4d (prefix caching + defaults + tests) — 2 days Item 3a (FA3 forward kernel: WGMMA + TMA + warp spec) — 3 days Week 4: Item 3b (FA3 backward kernel) — 3 days Item 3c–3d (SM selection + benchmarks) — 1 day Week 5: End-to-end integration benchmark: all 4 items together ``` --- ## Dependencies and Risks | Risk | Mitigation | |------|-----------| | cuBLASLt FP8 requires CUDA ≥ 12.0 | RTX 5060 Ti ships with CUDA 12.x; check `build.rs` version guard | | WGMMA PTX inline asm requires sm_90+ | Runtime SM check in `backend_selector.rs` falls back to FA2 | | CUDA Graphs incompatible with dynamic shapes | Static shape check + warning; disable gracefully | | SnapKV may hurt accuracy on long-context retrieval tasks | `keep_ratio` is configurable; default 0.6 leaves 60% of keys | | Prefix cache hash collisions | Zobrist hash with 64-bit space; false positive rate ~1/2^64 per request | --- ## Non-goals (Batch 1) - Metal FlashAttention (stub remains; deferred to Batch 2) - GaLore-2 optimizer (deferred; independent of these 4 items) - EAGLE-3 speculative decoding upgrade (spec exists; Batch 2) - Ring-AllReduce (requires multi-node setup; Batch 2) - MegaBlocks grouped GEMM for MoE (deferred to Batch 2)