GraphAR Capture-and-Replay with StaticKVCache in YuE: Implementation and Incompatible Configurations
The GraphAR capture-and-replay path requires StaticKVCache to maintain fixed memory addresses across CUDA graph replays, and fails with dynamic caches, disabled caching, or any configuration that reallocates KV tensors during generation.
The YuE library implements a high-performance graph-based auto-regressive (AR) generation mode that eliminates Python overhead by capturing the entire forward pass inside a CUDA graph. This mechanism depends entirely on the StaticKVCache implementation in src/yue2/modeling_yue2.py to ensure memory layouts remain constant during the capture-and-replay cycle.
How GraphAR Interacts with StaticKVCache
The graph-based AR implementation wraps the model's forward pass in a torch.cuda.graph to remove CPU overhead. Located in src/yue2/cuda_graph.py, the CUDAGraphAR class manages the capture and replay phases, while StaticKVCache provides the deterministic memory required for valid graph execution.
The Capture Phase
During capture, the model creates or receives a StaticKVCache instance and passes it to the backbone's forward call via the past_key_values argument. The CUDA graph records every kernel launch, including the cache's update() operations, because the cache guarantees no hidden memory allocations occur during this phase.
The Replay Phase
When replaying the captured graph, the same StaticKVCache instance is reused. Only new token embeddings and incremental cache slices are fed to the model, reproducing exactly the same memory pattern as the original capture. This enables zero-Python-overhead token generation where the graph re-executes identical GPU operations for each new token.
Why StaticKVCache Is Required
StaticKVCache satisfies four strict requirements for CUDA graph compatibility that dynamic caches cannot meet:
- Fixed-size storage: Pre-allocates tensors of shape
[batch, heads, max_seq_len, head_dim]for all layers at construction time. The CUDA graph cannot allocate new memory after capture, so the buffer must be static. - Append-only updates: The
update()method writes new key/value slices into pre-allocated buffers and advances an internalseen_tokenscounter without reshaping or reallocating. - Deterministic views: Returns the sub-tensor
[..., :end]representing the current cache prefix, guaranteeing identical stride and memory addresses on every replay. - No-copy semantics: Updates slices in-place rather than copying to new buffers, preserving the pointer graph captured in the CUDA graph.
Incompatible Model Configurations
The following configurations break the GraphAR capture-and-replay path in src/yue2/cuda_graph.py because they violate the static memory layout requirement:
use_cache=False: Disabling caching removes thepast_key_valuesargument, breaking the capture logic that expects a cache tensor.- DynamicCache: Expands or reallocates tensors during generation, invalidating the captured graph's memory layout.
- Insufficient
max_seq_len: If generation exceeds the pre-allocated buffer size,StaticKVCache.update()raisesValueErrorbecause the graph cannot grow the buffer dynamically. - Variable-length KV heads per layer:
StaticKVCacheassumes a singlenum_kv_headsvalue across all layers; changing head counts per layer breaks the static buffer layout. - Mid-generation dtype changes: Quantized models that switch cache dtype after capture corrupt the graph's expected memory layout.
- Beam-search or
reorder_cache: Operations that reorder cache entries viaindex_selectchange underlying memory addresses, defeating the no-copy guarantee required by the captured graph.
Implementation Example
The following code demonstrates correct usage of the GraphAR path with StaticKVCache:
import torch
from yue2.modeling_yue2 import YuE2Config, YuE2ForCausalLM, StaticKVCache
# 1. Build the model with caching enabled
config = YuE2Config(
vocab_size=50257,
hidden_size=1024,
num_hidden_layers=12,
num_key_value_heads=8,
head_dim=128,
max_position_embeddings=1024,
use_cache=True, # Required for GraphAR
)
model = YuE2ForCausalLM(config).cuda()
# 2. Create static KV cache matching model dimensions
cache = StaticKVCache(
num_layers=config.num_hidden_layers,
batch_size=1,
num_kv_heads=config.num_key_value_heads,
max_seq_len=config.max_position_embeddings,
head_dim=config.head_dim,
dtype=torch.float16,
device="cuda",
)
# 3. Capture the CUDA graph once
from yue2.cuda_graph import CUDAGraphAR
graph_ar = CUDAGraphAR(model, cache)
input_ids = torch.tensor([[config.bos_token_id]], device="cuda")
graph_ar.capture(input_ids) # Records the forward pass including cache updates
# 4. Replay for zero-overhead generation
generated = [config.bos_token_id]
for _ in range(20):
next_id = graph_ar.replay() # Fast GPU-only execution
generated.append(next_id.item())
print("Generated token IDs:", generated)
This example requires use_cache=True to ensure the forward pass accepts past_key_values, creates the StaticKVCache with exact model dimensions to prevent reallocation, and uses CUDAGraphAR to handle the graph lifecycle.
Key Source Files
src/yue2/modeling_yue2.py: Contains theStaticKVCacheclass providing fixed-size, append-only KV storage required for graph AR.src/yue2/cuda_graph.py: ImplementsCUDAGraphARwhich wrapstorch.cuda.grapharound the forward pass using the static cache.src/yue2/pipeline.py: High-level generation pipeline that may invokeCUDAGraphARfor optimized inference.tests/test_cuda_graph.py: Unit tests verifyingStaticKVCacheworks correctly with CUDA graphs and that incompatible configs raise errors.tests/test_model.py: Tests for directStaticKVCacheusage and size handling validation.
Summary
- GraphAR capture-and-replay in YuE eliminates Python overhead by recording the forward pass in a CUDA graph that includes cache updates via
past_key_values. - StaticKVCache is mandatory because it guarantees fixed memory addresses, append-only updates, and deterministic tensor views across all replay iterations.
- The system fails with
use_cache=False,DynamicCache, insufficientmax_seq_len, variable KV heads per layer, dtype changes during generation, and cache reordering operations like beam search. - Implementation requires strict adherence to the memory constraints enforced in
modeling_yue2.pyandcuda_graph.py.
Frequently Asked Questions
What happens if I use DynamicCache with GraphAR?
The CUDAGraphAR.capture() call will raise an error or the replay will produce incorrect results. DynamicCache reallocates tensors as sequence length increases in src/yue2/modeling_yue2.py, which invalidates the memory addresses recorded in the captured CUDA graph. Only StaticKVCache maintains the fixed buffer layout required for graph replay.
Why does StaticKVCache require a fixed max_seq_len?
CUDA graphs cannot allocate memory dynamically after capture. StaticKVCache pre-allocates buffers of shape [batch, heads, max_seq_len, head_dim] at construction time. If generation exceeds this length, the update() method raises ValueError because the captured graph in src/yue2/cuda_graph.py cannot expand the underlying storage.
Can I use beam search with GraphAR capture-and-replay?
No. Beam search requires reorder_cache operations that use index_select to rearrange cache entries, changing the underlying memory layout. This defeats the no-copy guarantee and fixed pointer graph required by the captured CUDA graph. Use greedy sampling or other static-order generation methods instead.
Is use_cache=True sufficient for GraphAR compatibility?
No, use_cache=True is necessary but not sufficient. You must explicitly use StaticKVCache rather than the default cache type, and ensure all other constraints (fixed max_seq_len, consistent dtype, no reordering) are met. The CUDAGraphAR class validates compatible cache types during initialization to prevent runtime graph corruption.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →