Files
nvfp4-megamoe-kernel/GETTING_CUDAGRAPH_READY.md
biondizzle ba68212fa7 Add CUDA graph readiness detector (Section A of GETTING_CUDAGRAPH_READY.md)
- Grep for Section B sync patterns in hot path files
- Method 1: run decode forward with torch.cuda.set_sync_debug_mode('error')
- Method 2: attempt CUDA graph capture of L0 decode step
- Full model load + prefill + warmup before detection
- Results saved to /tmp/cuda_graph_readiness_results.json
2026-06-03 16:34:15 +00:00

7.1 KiB

DSV4 → vLLM: CUDA-Graph Safety / GPU-Native Requirements (PART 2 companion)

Goal: the per-step decode forward must be fully GPU-native so vLLM can capture and replay it. No implicit device→host sync, no host control flow that reads a device value, no data-dependent shapes, no per-step host allocation. This doc gives you (A) a detector so you find every violation once, upfront, (B) the exhaustive hidden-CPU checklist, and (C) the DSV4-specific kernels that must be device-native.

The one rule that decides everything

Branching on a host-known integer (step number, position, batch size, dtype, static shape) is graph-compatible — you capture one graph per bucket and the scheduler picks by that integer. Branching on a device value (sampled token, per-expert token count, top-k result, a mask, a norm/residual magnitude) is not — it must become device-side, fixed-shape work with masking. Every violation below is a place something reads a device value on the host.

You do not need one monolithic graph. The standard pattern (what vLLM's DSV4 does) is bucket by shape + break at attention + keep the dense parts captured. Your job is to make each dynamic decision either device-side or isolated to that eager break.


SECTION A — The detector (build this FIRST, before porting anything)

Stop hunting syncs by hand. Make them fail at the exact line:

import torch
torch.cuda.set_sync_debug_mode("error")   # raises at any implicit device→host sync
# ... run one decode step of the forward ...
torch.cuda.set_sync_debug_mode("default")

And a capture-under-test (most illegal host ops error during capture):

g = torch.cuda.CUDAGraph()
# static input buffers allocated ONCE, outside capture:
with torch.cuda.graph(g):
    out = decode_step(static_inputs)     # capture fails loudly on .item(), sync, alloc, etc.
for _ in range(3):
    static_inputs.copy_(next_inputs);  g.replay()   # replay must reproduce eager output

Do this on the current single_shot forward first — it inventories every existing sync in one pass, so you get the whole hunt-list upfront instead of discovering them one at a time during vLLM bring-up. Then gate every commit on both checks in CI; the day someone adds a .item(), the build fails at that line.

Also useful: compute-sanitizer --tool synccheck, and nsys to eyeball CPU↔GPU stall gaps.


SECTION B — The hidden-CPU checklist (grep the hot path for these)

Explicit device→host transfers .item() · .cpu() · .tolist() · .numpy() · int(t) / float(t) / bool(t) · print(t) · f-strings/logging that interpolate a tensor · assert (device_condition) (e.g. assert (x>0).all()) · .to("cpu")

Host control flow on device values if t: · if mask.any(): · if x.sum() > thr: · while t > 0: · for i in range(n.item()) · convergence early-exit reading a device residual · choosing a kernel based on the sampled token

Data-dependent shapes (these both change shape AND sync) torch.nonzero · torch.where(cond) (one-arg form) · torch.unique · torch.bincount (when it drives a shape) · boolean/mask indexing x[mask], x[x>0] · masked_select · reshape(n.item(), ...) · any gather sized by a device-computed count

Per-step host allocation torch.empty/zeros/tensor([...]) created fresh inside the captured region · building a Python list then torch.tensor(list, device=...) · np.* anywhere on the path · any CPU tensor then .to(device) per step

Host RNG random.* / np.random.* / Python rng for sampling → use a device generator / captured philox state

Sync primitives & checks torch.cuda.synchronize() · stream.synchronize() · torch.isnan(x).any() / isinf(...).any() debug guards · pinned-copy syncs

Sneaky ones (the "didn't realize" category) sum(t) / min(t) / max(t) (Python builtins iterate → sync; use t.sum()) · a .cpu()/.item() hidden inside a logging, assert, or metrics helper · einops rearrange with a data-dependent dim · telemetry/progress hooks that read tensors · indexing a tensor with a Python int derived from .item()

What is FINE (no sync, don't waste time on these) .shape / .size() / .numel() / .dtype (host metadata, no sync) · branching on host-known ints (step/batch/static shape) · dtype/shape kernel dispatch · the stop-token check, detokenize, and your BF16 precision-floor dequant (all load-time or outside the captured graph — leave them on host, that's correct).


SECTION C — DSV4-specific kernels that must be GPU-native

# Hazard (current host/dynamic behavior) Requirement vLLM reference
1 Compressor returns None for 3/4 (CSA) or 127/128 (HCA) decode steps — periodic host branch Compress every step into a persistent partial-state/ring buffer; emit the compressed entry device-side on the boundary save_partial_states, fused_compress_quant_cache
2 KV grows each step → attention shape changes Paged KV (fixed blocks + block table) captured at fixed max-len with masking, or make attention the eager break breakable_cudagraph / eager_break_during_capture; AttentionCGSupport.ALWAYS
3 Indexer top-k → host reads selected count to size gather Always gather fixed k (padded), mask invalid; no host read of the count dequant_gather_k_cutedsl (fixed-shape gather)
4 MoE top-6 → per-expert token counts drive per-expert launches Routing permutation/offsets computed on device; grouped GEMM with device offsets and a fixed total launch prepare_megamoe
5 Next token / positions managed on host, fresh tensors per step Static I/O buffers allocated once; in-place copy_ of next token; positions via device-side increment (or per-shape bucketed graphs) vLLM persistent input buffers

Also confirm:

  • Sinkhorn runs a fixed 20 iterations with no host convergence check (a while not converged reading a device residual breaks capture). Fixed-iteration = safe.
  • Sampler is device-side; repetition_penalty reads from a fixed-size device recent-token buffer (not a growing Python list); the EOS/stop decision is a host step outside the graph (correct).

SECTION D — Integration order

  1. Build Section A's detector and run it on the current forward — get the full sync inventory in one pass.
  2. Fix Section C's five device-native kernels (these are the structural ones; the rest of Section B tends to be incidental .item()s once these are right).
  3. Re-run capture-under-test until it captures clean and replay matches eager bit-for-bit.
  4. Gate every commit on the capture test so violations can never silently return.

Guardrails

  • Keep the stop-check, detokenize, and load-time BF16 dequant on the host — they're outside the captured region by design; don't contort them to be "graph-safe."
  • Decide the attention model up front (paged-capturable vs eager-break) — retrofitting it later forces a KV-cache rewrite.
  • Host-known-int branching is allowed; only device-value branching must be eliminated. Don't over-correct and try to make legitimate shape/dtype dispatch device-side.