feat: add CUDA-graph capture for decode forward

- New CudaGraphContext class: warmup -> capture -> replay lifecycle
- One graph per batch_size key, all inputs at fixed workspace addresses
- Added position_ids buffer to InferenceWorkspace (required for graph capture)
- Graph only activates when CUDA backend is the current backend
- Default off (opt-in) due to slight numerical divergence in graph replay
- Sampling stays outside the graph (torch.multinomial uses mutable RNG)
- Resolved circular import: KVCache -> TYPE_CHECKING in attention_backend.py
This commit is contained in:
2026-08-06 19:57:12 +08:00
parent 6f67ba8942
commit 6054b8dbd4
4 changed files with 165 additions and 32 deletions
+15 -13
View File
@@ -34,7 +34,7 @@ import importlib
import threading import threading
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from contextlib import contextmanager from contextlib import contextmanager
from typing import Optional, Union from typing import TYPE_CHECKING, Optional, Union
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
@@ -45,6 +45,8 @@ from astrai.extension.attention_ops import (
attn_paged_prefill, attn_paged_prefill,
) )
from astrai.factory import BaseFactory from astrai.factory import BaseFactory
if TYPE_CHECKING:
from astrai.inference.core.cache import KVCache from astrai.inference.core.cache import KVCache
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar( _current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
@@ -199,7 +201,7 @@ def attention(
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache] = None, kv_cache: Optional["KVCache"] = None,
layer_id: int = 0, layer_id: int = 0,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
@@ -255,7 +257,7 @@ class AttentionBackend(ABC):
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional["KVCache"],
layer_id: int, layer_id: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
@@ -284,7 +286,7 @@ class AttentionBackend(ABC):
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional["KVCache"],
layer_id: int, layer_id: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
@@ -297,7 +299,7 @@ class AttentionBackend(ABC):
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional["KVCache"],
layer_id: int, layer_id: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
@@ -326,7 +328,7 @@ class TorchNativeBackend(AttentionBackend):
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional["KVCache"],
layer_id: int, layer_id: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
@@ -338,7 +340,7 @@ class TorchNativeBackend(AttentionBackend):
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional["KVCache"],
layer_id: int, layer_id: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
@@ -350,7 +352,7 @@ class TorchNativeBackend(AttentionBackend):
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional["KVCache"],
layer_id: int, layer_id: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
@@ -416,7 +418,7 @@ class CudaBackend(AttentionBackend):
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional["KVCache"],
layer_id: int, layer_id: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
@@ -451,7 +453,7 @@ class CudaBackend(AttentionBackend):
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional["KVCache"],
layer_id: int, layer_id: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
@@ -503,7 +505,7 @@ class FlashAttnBackend(AttentionBackend):
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional["KVCache"],
layer_id: int, layer_id: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
@@ -515,7 +517,7 @@ class FlashAttnBackend(AttentionBackend):
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional["KVCache"],
layer_id: int, layer_id: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
@@ -527,7 +529,7 @@ class FlashAttnBackend(AttentionBackend):
q: Tensor, q: Tensor,
k: Tensor, k: Tensor,
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional["KVCache"],
layer_id: int, layer_id: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
+43 -18
View File
@@ -5,7 +5,9 @@ from typing import List, Optional
import torch import torch
from torch import Tensor from torch import Tensor
from astrai.extension.attention_backend import CudaBackend, get_backend
from astrai.inference.core.cache import PagePool from astrai.inference.core.cache import PagePool
from astrai.inference.core.graph import CudaGraphContext
from astrai.inference.core.task import Task from astrai.inference.core.task import Task
from astrai.inference.core.workspace import InferenceWorkspace from astrai.inference.core.workspace import InferenceWorkspace
from astrai.inference.sample import sample from astrai.inference.sample import sample
@@ -90,6 +92,12 @@ class Executor:
dtype=self.dtype, dtype=self.dtype,
) )
# CUDA-graph capture: one graph per (batch_size, total_len) key.
# The graph captures model.forward() with fixed-address workspace
# inputs. Before each replay, input content is updated in-place so
# the graph sees fresh token IDs / positions / KV metadata.
self._graph_ctx = CudaGraphContext()
def _sample_logits( def _sample_logits(
self, self,
logits: Tensor, logits: Tensor,
@@ -204,42 +212,59 @@ class Executor:
if not tasks: if not tasks:
return [] return []
input_ids = self._workspace.fill_input_ids( b = len(tasks)
ws = self._workspace
# ---- pre-replay: update input buffers in-place ----
input_ids = ws.fill_input_ids(
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks] [t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks]
).unsqueeze(1) )
task_ids = [t.task_id for t in tasks] task_ids = [t.task_id for t in tasks]
cur_positions = [t.next_pos for t in tasks]
sig = tuple(task_ids) sig = tuple(task_ids)
cur_positions = [t.next_pos for t in tasks]
cached = self._decode_cache cached = self._decode_cache
if ( if (
cached is not None cached is not None
and cached[0] == sig and cached[0] == sig
and cur_positions == [p + 1 for p in cached[1]] and cur_positions == [p + 1 for p in cached[1]]
): ):
_, _, info, position_ids = cached info = cached[2]
position_ids += 1 ws.position_ids[:b] += 1
self._decode_cache = (sig, cur_positions, info, position_ids) self._decode_cache = (sig, cur_positions, info)
else: else:
info = _build_sampling_batch_info(tasks, self.device) info = _build_sampling_batch_info(tasks, self.device)
position_ids = torch.tensor( ws.position_ids[:b].copy_(
cur_positions, dtype=torch.long, device=self.device torch.tensor(cur_positions, dtype=torch.long, device=self.device)
) )
self._decode_cache = (sig, cur_positions, info, position_ids) self._decode_cache = (sig, cur_positions, info)
total_len = max(t.next_pos for t in tasks) + 1 total_len = max(cur_positions) + 1
input_mask = self._workspace.decode_mask(position_ids, total_len) input_mask = ws.decode_mask(ws.position_ids[:b], total_len)
kv_cache = self.kv_cache.bind_tasks(task_ids, ws)
# ---- forward (graph replay or live run + capture) ----
use_graph = (
self._graph_ctx.enabled
and "cuda" in str(self.device)
and isinstance(get_backend(), CudaBackend)
)
key = (b,)
if use_graph:
input_mask = ws.decode_mask(ws.position_ids[:b], ws.max_seq_len)
with torch.inference_mode(): with torch.inference_mode():
outputs = self.model( outputs = self._graph_ctx.forward(
input_ids, self.model,
key=key,
input_ids=input_ids.unsqueeze(1),
input_mask=input_mask, input_mask=input_mask,
kv_cache=self.kv_cache.bind_tasks( kv_cache=kv_cache,
task_ids, position_ids=ws.position_ids[:b].unsqueeze(1),
self._workspace,
),
position_ids=position_ids.unsqueeze(1),
) )
logits = outputs["logits"][:, -1, :] logits = outputs["logits"][:, -1, :]
+101
View File
@@ -0,0 +1,101 @@
"""CUDA-graph capture for the decode model-forward step.
Mirrors SGLang's cuda-graph manager: one graph per batch size. The graph
pair. The graph captures ``model.forward()`` with workspace-backed inputs
(all at fixed addresses). Before each replay the caller updates the input
buffer content in-place so the graph sees fresh data at the same tensor
addresses.
Only the model forward is captured — sampling runs outside the graph
(via ``torch.multinomial`` which consumes a mutable RNG state).
"""
import torch
from torch import Tensor
class CudaGraphContext:
"""CUDA-graph capture/replay for decode steps.
Parameters:
enabled: When ``False``, ``forward()`` always runs the live model
forward without capture/replay (graphs are cleared). Toggle at
runtime via the ``set_enabled()`` method.
Usage::
gctx = CudaGraphContext()
with torch.inference_mode():
outputs = gctx.forward(
model,
key=(batch_size,),
input_ids=workspace.input_ids[:b].unsqueeze(1),
input_mask=input_mask,
kv_cache=kv_cache,
position_ids=workspace.position_ids[:b].unsqueeze(1),
)
The first call at a given key runs *without* capture (warmup). The
second call captures the graph. Subsequent calls replay the captured
graph. A ``torch.cuda.synchronize()`` before capture drains in-flight
work so the graph trace is clean.
"""
def __init__(self, enabled: bool = False):
self._enabled = enabled
self._graphs: dict[tuple, torch.cuda.CUDAGraph] = {}
self._outputs: dict[tuple, dict[str, Tensor]] = {}
self._warmed: set[tuple] = set()
@property
def enabled(self) -> bool:
return self._enabled
def set_enabled(self, flag: bool):
"""Enable or disable CUDA-graph capture at runtime.
Disabling clears all captured graphs (frees GPU memory) and warmup
state. Re-enabling after disable starts fresh — graphs are
re-captured on the next warmup cycle.
"""
if flag == self._enabled:
return
self._enabled = flag
if not flag:
self._graphs.clear()
self._outputs.clear()
self._warmed.clear()
def forward(self, model, *, key, **kwargs) -> dict[str, Tensor]:
"""Run ``model(**kwargs)`` via graph replay or live forward.
Args:
model: callable, e.g. ``self.model.forward``.
key: ``(batch_size,)`` — the dispatch key (one graph per batch size).
**kwargs: arguments forwarded to ``model``. All tensor arguments
must reside at stable addresses (workspace buffers).
Returns:
The dict produced by ``model(**kwargs)``, e.g.
``{"logits": ..., "h0": ...}``.
"""
if not self._enabled:
self._outputs[key] = model(**kwargs)
return self._outputs[key]
if key in self._graphs:
self._graphs[key].replay()
elif key in self._warmed:
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
self._outputs[key] = model(**kwargs)
self._graphs[key] = graph
self._warmed.discard(key)
else:
self._warmed.add(key)
self._outputs[key] = model(**kwargs)
return self._outputs[key]
def has_graph(self, key: tuple) -> bool:
return key in self._graphs
+5
View File
@@ -88,6 +88,11 @@ class InferenceWorkspace:
(max_batch_size, 1), dtype=torch.long, device=device (max_batch_size, 1), dtype=torch.long, device=device
) )
# Per-step position IDs (must be at a fixed address for CUDA-graph capture).
self.position_ids = torch.empty(
(max_batch_size,), dtype=torch.long, device=device
)
# Split-KV partial-result buffers for decode (persistent, one global # Split-KV partial-result buffers for decode (persistent, one global
# alloc per process — mirrors FlashInfer's workspace pattern). # alloc per process — mirrors FlashInfer's workspace pattern).
# Shape: [max_batch_size, max_q_heads, _MAX_SPLITS, head_dim] (o_part) # Shape: [max_batch_size, max_q_heads, _MAX_SPLITS, head_dim] (o_part)