CUDA Graphs¶
Overview¶
A CUDA graph captures a deterministic sequence of kernel launches (with fixed memory addresses and shapes) into a graph object, then replays it with a single launch — eliminating per-kernel CPU launch overhead and inter-kernel synchronization. For small, launch-bound models this is often the single biggest win.
- Launch overhead: ~5-10 microseconds per kernel from CPU.
- A 300-kernel training step ⇒ up to ~2-3 ms saved just in launch cost.
- Capture requires static shapes, static memory, no dynamic allocation, no Python control flow in the capture region.
torch.cuda.CUDAGraph()+g.capture_begin()/capture_end(), theng.replay().
💡
torch.compile(mode="reduce-overhead")automatically uses CUDA graphs — but hand-rolled graphs remain the tool for specialized loops (e.g., RL environments, inference servers).
Warm-Up Is Mandatory¶
Capture cannot tolerate lazy allocations; you must run the workload once on a side stream to "warm" the caching allocator.
import torch
def warmup(fn, *args, n=10):
for _ in range(n):
fn(*args)
torch.cuda.synchronize()
Minimal Capture / Replay Pattern¶
import torch
class GraphModel:
def __init__(self, model, x):
self.model = model.cuda().eval()
self.x = x
# warm-up on a side stream
s = torch.cuda.Stream()
s.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(s):
for _ in range(3):
self.out = self.model(self.x)
torch.cuda.current_stream().wait_stream(s)
torch.cuda.synchronize()
# capture
self.g = torch.cuda.CUDAGraph()
with torch.cuda.graph(self.g):
self.out = self.model(self.x)
def __call__(self):
self.g.replay()
return self.out
model = torch.compile # placeholder; see full example below
Full working capture¶
import torch, torch.nn as nn
m = nn.Sequential(nn.Linear(256, 256), nn.ReLU(), nn.Linear(256, 16)).cuda().eval()
x = torch.randn(128, 256, device='cuda')
y = torch.empty(128, 16, device='cuda')
# warmup
s = torch.cuda.Stream()
s.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(s):
for _ in range(3):
y.copy_(m(x))
torch.cuda.current_stream().wait_stream(s)
torch.cuda.synchronize()
g = torch.cuda.CUDAGraph()
with torch.cuda.graph(g):
y.copy_(m(x)) # capture everything into graph
# replay: single launch, output lands in y
g.replay()
print(y.shape)
⚠️ Outputs
ymust be pre-allocated and captured via in-place copy — graph replays write into the same buffers.
What Breaks Capture¶
| Pitfall | Effect |
|---|---|
| New allocations inside capture | capture error / device sync crash |
| Data-dependent shapes | must be static |
| Python branching on tensor values | not allowed |
.item(), .numpy(), print |
syncs break capture |
| Changing input tensor content via new tensor | content changes OK; address must stay fixed |
Safe way to feed new inputs: copy new data into the captured input buffer.
buf = x.clone() # captured buffer
with torch.cuda.graph(g):
y.copy_(m(buf))
# later, for new sample:
buf.copy_(new_data) # memcpy into same address
g.replay()
CUDA Graphs + torch.compile¶
m_compiled = torch.compile(m, mode="reduce-overhead") # transparent cudagraphs
# same warm-up applies; the compiled graph is queried via .cudagraphs
If you're already using reduce-overhead, hand-rolled graphs buy little for inference; keep them for specialized control loops.
When to Use Hand-Rolled Graphs¶
- Inference servers — fixed latency targets, static shapes.
- RL rollout loops — small networks called thousands of times.
- Custom training loops — fixed-shape steps (no dynamic padding).
- kernel-bound small models (MLPs, small transformers).
Skip them when: dynamic batching required, shapes vary per request, or torch.compile(...reduce-overhead) already covers you.
Benchmarking Graphs Honestly¶
def bench(fn, n=100):
torch.cuda.synchronize()
s, e = torch.cuda.Event(True), torch.cuda.Event(True)
s.record()
for _ in range(n):
fn()
e.record(); torch.cuda.synchronize()
return s.elapsed_time(e) / n # ms/step
m = nn.Sequential(nn.Linear(256, 256), nn.ReLU(), nn.Linear(256, 16)).cuda().eval()
m_compiled = torch.compile(m, mode="reduce-overhead")
print("eager : %.3f ms" % bench(lambda: m(x)))
print("compiled: %.3f ms" % bench(lambda: m_compiled(x)))
Key Takeaways¶
- CUDA graphs remove launch overhead, not kernel time: best for launch-bound/small models.
- Capture is strict: pre-warm, static shapes, pre-allocated outputs, feed via buffer copies.
torch.compile(mode="reduce-overhead")gives you graphs for free — prefer it unless you need raw control.- Replays are deterministic and cheap; graphs compose with streams.
- Profile first: if GPU utilization is high, graphs won't save you.