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:
@@ -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,7 +45,9 @@ from astrai.extension.attention_ops import (
|
|||||||
attn_paged_prefill,
|
attn_paged_prefill,
|
||||||
)
|
)
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.inference.core.cache import KVCache
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from astrai.inference.core.cache import KVCache
|
||||||
|
|
||||||
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
|
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
|
||||||
"attn_backend"
|
"attn_backend"
|
||||||
@@ -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,
|
||||||
|
|||||||
@@ -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, :]
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user