fix: add out_buf to attn_paged_decode for CUDA graph capture compatibility
- Pre-allocate decode_out in InferenceWorkspace so attn_paged_decode does not call torch::empty inside graph capture - Wire decode_out through KVCache, PagePool.bind_tasks, and CudaBackend.fwd_decode - Run live forward before graph capture to get valid output (graph pool memory is zeroed after capture block exits) - Greedy generation with graph replay is bit-exact across all batch sizes - Decode speedups vs no-graph: B=1 2.09x, B=4 1.80x, B=8 1.94x, B=16 1.76x
This commit is contained in:
@@ -511,6 +511,7 @@ class CudaBackend(AttentionBackend):
|
|||||||
is_causal=True,
|
is_causal=True,
|
||||||
o_part_buf=kv_cache.decode_o_part,
|
o_part_buf=kv_cache.decode_o_part,
|
||||||
ml_part_buf=kv_cache.decode_ml_part,
|
ml_part_buf=kv_cache.decode_ml_part,
|
||||||
|
out_buf=kv_cache.decode_out,
|
||||||
)
|
)
|
||||||
return out.unsqueeze(1).flatten(2)
|
return out.unsqueeze(1).flatten(2)
|
||||||
|
|
||||||
|
|||||||
@@ -102,6 +102,7 @@ def attn_paged_decode(
|
|||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
o_part_buf: Optional[torch.Tensor] = None,
|
o_part_buf: Optional[torch.Tensor] = None,
|
||||||
ml_part_buf: Optional[torch.Tensor] = None,
|
ml_part_buf: Optional[torch.Tensor] = None,
|
||||||
|
out_buf: Optional[torch.Tensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""SGLang-style paged decode (q_len == 1, flat KV pool).
|
"""SGLang-style paged decode (q_len == 1, flat KV pool).
|
||||||
|
|
||||||
@@ -121,6 +122,7 @@ def attn_paged_decode(
|
|||||||
is_causal: apply causal mask
|
is_causal: apply causal mask
|
||||||
o_part_buf: pre-allocated split-KV o partial buffer (workflow bypass)
|
o_part_buf: pre-allocated split-KV o partial buffer (workflow bypass)
|
||||||
ml_part_buf: pre-allocated split-KV m/l buffer (workflow bypass)
|
ml_part_buf: pre-allocated split-KV m/l buffer (workflow bypass)
|
||||||
|
out_buf: pre-allocated output buffer [batch, n_heads, head_dim] (graph-safe)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
[batch, n_heads, head_dim] (bf16, 3D)
|
[batch, n_heads, head_dim] (bf16, 3D)
|
||||||
@@ -139,6 +141,7 @@ def attn_paged_decode(
|
|||||||
causal_offset=causal_offset,
|
causal_offset=causal_offset,
|
||||||
o_part_buf=o_part_buf,
|
o_part_buf=o_part_buf,
|
||||||
ml_part_buf=ml_part_buf,
|
ml_part_buf=ml_part_buf,
|
||||||
|
out_buf=out_buf,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -263,6 +263,7 @@ class KVCache:
|
|||||||
qo_indptr: [batch+1] int32 — prefill qo prefix-sum (None in decode)
|
qo_indptr: [batch+1] int32 — prefill qo prefix-sum (None in decode)
|
||||||
decode_o_part: split-KV o partial workspace (mirrors FlashInfer)
|
decode_o_part: split-KV o partial workspace (mirrors FlashInfer)
|
||||||
decode_ml_part: split-KV m/l partial workspace (mirrors FlashInfer)
|
decode_ml_part: split-KV m/l partial workspace (mirrors FlashInfer)
|
||||||
|
decode_out: pre-allocated decode output buffer (graph-safe)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
k_buffer: Tensor
|
k_buffer: Tensor
|
||||||
@@ -276,6 +277,7 @@ class KVCache:
|
|||||||
qo_indptr: Optional[Tensor] = None
|
qo_indptr: Optional[Tensor] = None
|
||||||
decode_o_part: Optional[Tensor] = None
|
decode_o_part: Optional[Tensor] = None
|
||||||
decode_ml_part: Optional[Tensor] = None
|
decode_ml_part: Optional[Tensor] = None
|
||||||
|
decode_out: Optional[Tensor] = None
|
||||||
|
|
||||||
|
|
||||||
class PagePool:
|
class PagePool:
|
||||||
@@ -571,6 +573,7 @@ class PagePool:
|
|||||||
)
|
)
|
||||||
qo_indptr = workspace.qo_indptr[: b + 1]
|
qo_indptr = workspace.qo_indptr[: b + 1]
|
||||||
decode_o_part, decode_ml_part = None, None
|
decode_o_part, decode_ml_part = None, None
|
||||||
|
decode_out = None
|
||||||
else:
|
else:
|
||||||
write_pos = seq_lens_t - 1
|
write_pos = seq_lens_t - 1
|
||||||
loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1)
|
loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1)
|
||||||
@@ -579,6 +582,7 @@ class PagePool:
|
|||||||
qo_indptr = None
|
qo_indptr = None
|
||||||
decode_o_part = getattr(workspace, "decode_o_part", None)
|
decode_o_part = getattr(workspace, "decode_o_part", None)
|
||||||
decode_ml_part = getattr(workspace, "decode_ml_part", None)
|
decode_ml_part = getattr(workspace, "decode_ml_part", None)
|
||||||
|
decode_out = getattr(workspace, "decode_out", None)
|
||||||
|
|
||||||
return KVCache(
|
return KVCache(
|
||||||
k_buffer=self._storage.k_buffer,
|
k_buffer=self._storage.k_buffer,
|
||||||
@@ -592,6 +596,7 @@ class PagePool:
|
|||||||
qo_indptr=qo_indptr,
|
qo_indptr=qo_indptr,
|
||||||
decode_o_part=decode_o_part,
|
decode_o_part=decode_o_part,
|
||||||
decode_ml_part=decode_ml_part,
|
decode_ml_part=decode_ml_part,
|
||||||
|
decode_out=decode_out,
|
||||||
)
|
)
|
||||||
|
|
||||||
# ---- internals ----
|
# ---- internals ----
|
||||||
|
|||||||
@@ -86,12 +86,14 @@ class CudaGraphContext:
|
|||||||
if key in self._graphs:
|
if key in self._graphs:
|
||||||
self._graphs[key].replay()
|
self._graphs[key].replay()
|
||||||
elif key in self._warmed:
|
elif key in self._warmed:
|
||||||
|
cap_output = model(**kwargs)
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
graph = torch.cuda.CUDAGraph()
|
graph = torch.cuda.CUDAGraph()
|
||||||
with torch.cuda.graph(graph):
|
with torch.cuda.graph(graph):
|
||||||
self._outputs[key] = model(**kwargs)
|
self._outputs[key] = model(**kwargs)
|
||||||
self._graphs[key] = graph
|
self._graphs[key] = graph
|
||||||
self._warmed.discard(key)
|
self._warmed.discard(key)
|
||||||
|
return cap_output
|
||||||
else:
|
else:
|
||||||
self._warmed.add(key)
|
self._warmed.add(key)
|
||||||
self._outputs[key] = model(**kwargs)
|
self._outputs[key] = model(**kwargs)
|
||||||
|
|||||||
@@ -108,6 +108,14 @@ class InferenceWorkspace:
|
|||||||
device=device,
|
device=device,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Decode output buffer (graph-safe pre-alloc). Shape matches the
|
||||||
|
# decode kernel's output: [batch, q_head, head_dim].
|
||||||
|
self.decode_out = torch.empty(
|
||||||
|
(max_batch_size, max_q_heads, head_dim),
|
||||||
|
dtype=dtype,
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
|
||||||
def decode_buffers(self, batch: int, q_heads: int):
|
def decode_buffers(self, batch: int, q_heads: int):
|
||||||
"""Return ``(o_part, ml_part)`` view sliced to live dimensions."""
|
"""Return ``(o_part, ml_part)`` view sliced to live dimensions."""
|
||||||
return (
|
return (
|
||||||
|
|||||||
@@ -13,7 +13,8 @@ torch::Tensor attn_paged_decode(
|
|||||||
int64_t causal_offset,
|
int64_t causal_offset,
|
||||||
double scale,
|
double scale,
|
||||||
c10::optional<torch::Tensor> o_part_buf,
|
c10::optional<torch::Tensor> o_part_buf,
|
||||||
c10::optional<torch::Tensor> ml_part_buf
|
c10::optional<torch::Tensor> ml_part_buf,
|
||||||
|
c10::optional<torch::Tensor> out_buf
|
||||||
) {
|
) {
|
||||||
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
|
||||||
auto stream = at::cuda::getCurrentCUDAStream();
|
auto stream = at::cuda::getCurrentCUDAStream();
|
||||||
@@ -23,7 +24,18 @@ torch::Tensor attn_paged_decode(
|
|||||||
req_to_token, req_pool_indices, kv_indptr,
|
req_to_token, req_pool_indices, kv_indptr,
|
||||||
max_seq_len, mask, causal_offset, scale, p);
|
max_seq_len, mask, causal_offset, scale, p);
|
||||||
|
|
||||||
auto O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
|
torch::Tensor O;
|
||||||
|
if (out_buf.has_value() && out_buf->defined()) {
|
||||||
|
TORCH_CHECK(out_buf->dtype() == q.dtype(), "out_buf dtype must match q");
|
||||||
|
TORCH_CHECK(out_buf->size(0) >= q.size(0), "out_buf batch too small");
|
||||||
|
TORCH_CHECK(out_buf->size(1) >= q.size(1), "out_buf heads too small");
|
||||||
|
TORCH_CHECK(out_buf->size(2) >= q.size(2), "out_buf head_dim too small");
|
||||||
|
O = out_buf.value().slice(0, 0, q.size(0))
|
||||||
|
.slice(1, 0, q.size(1))
|
||||||
|
.slice(2, 0, q.size(2));
|
||||||
|
} else {
|
||||||
|
O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
|
||||||
|
}
|
||||||
p.o = (bf16*)O.data_ptr();
|
p.o = (bf16*)O.data_ptr();
|
||||||
|
|
||||||
if (o_part_buf.has_value() && ml_part_buf.has_value()
|
if (o_part_buf.has_value() && ml_part_buf.has_value()
|
||||||
@@ -57,5 +69,6 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
|||||||
py::arg("scale") = 0.0,
|
py::arg("scale") = 0.0,
|
||||||
py::arg("o_part_buf") = py::none(),
|
py::arg("o_part_buf") = py::none(),
|
||||||
py::arg("ml_part_buf") = py::none(),
|
py::arg("ml_part_buf") = py::none(),
|
||||||
|
py::arg("out_buf") = py::none(),
|
||||||
"SGLang-style paged decode: flat KV pool + req_to_token + kv_indptr.");
|
"SGLang-style paged decode: flat KV pool + req_to_token + kv_indptr.");
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ from astrai import setup_logging
|
|||||||
from astrai.config import BaseModelConfig, ConfigFactory
|
from astrai.config import BaseModelConfig, ConfigFactory
|
||||||
from astrai.extension import ATTN_BACKEND, AttentionBackendFactory, attn_backend
|
from astrai.extension import ATTN_BACKEND, AttentionBackendFactory, attn_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.workspace import InferenceWorkspace
|
from astrai.inference.core.workspace import InferenceWorkspace
|
||||||
from astrai.model import AutoModel, AutoRegressiveLM
|
from astrai.model import AutoModel, AutoRegressiveLM
|
||||||
|
|
||||||
@@ -57,6 +58,7 @@ class GenerationBenchmark:
|
|||||||
dtype: torch.dtype = torch.bfloat16,
|
dtype: torch.dtype = torch.bfloat16,
|
||||||
cache_type: str = "contiguous",
|
cache_type: str = "contiguous",
|
||||||
backend: Union[str, ATTN_BACKEND] = ATTN_BACKEND.CUDA,
|
backend: Union[str, ATTN_BACKEND] = ATTN_BACKEND.CUDA,
|
||||||
|
cuda_graph: bool = False,
|
||||||
):
|
):
|
||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
@@ -64,6 +66,7 @@ class GenerationBenchmark:
|
|||||||
self.model = model
|
self.model = model
|
||||||
self.config = config
|
self.config = config
|
||||||
self.backend = backend
|
self.backend = backend
|
||||||
|
self.cuda_graph = cuda_graph
|
||||||
|
|
||||||
def _make_pool(self, batch_size: int, max_seq_len: int) -> PagePool:
|
def _make_pool(self, batch_size: int, max_seq_len: int) -> PagePool:
|
||||||
return PagePool(
|
return PagePool(
|
||||||
@@ -217,11 +220,97 @@ class GenerationBenchmark:
|
|||||||
prompt_length: int = 512,
|
prompt_length: int = 512,
|
||||||
gen_length: int = 128,
|
gen_length: int = 128,
|
||||||
num_trials: int = 5,
|
num_trials: int = 5,
|
||||||
|
) -> BenchmarkResult:
|
||||||
|
if self.cuda_graph and self.backend == "cuda":
|
||||||
|
return self._run_graph_decode_benchmark(
|
||||||
|
batch_size, prompt_length, gen_length, num_trials
|
||||||
|
)
|
||||||
|
return self._run_plain_decode_benchmark(
|
||||||
|
batch_size, prompt_length, gen_length, num_trials
|
||||||
|
)
|
||||||
|
|
||||||
|
def _run_graph_decode_benchmark(
|
||||||
|
self,
|
||||||
|
batch_size: int,
|
||||||
|
prompt_length: int,
|
||||||
|
gen_length: int,
|
||||||
|
num_trials: int,
|
||||||
|
) -> BenchmarkResult:
|
||||||
|
import time
|
||||||
|
|
||||||
|
max_seq_len = prompt_length + 5 + gen_length * num_trials
|
||||||
|
pool = self._make_pool(batch_size, max_seq_len)
|
||||||
|
workspace = self._make_workspace(pool, self.config)
|
||||||
|
task_ids = self._run_prefill(pool, batch_size, prompt_length, workspace)
|
||||||
|
|
||||||
|
b = batch_size
|
||||||
|
input_ids_buf = torch.zeros(b, 1, dtype=torch.long, device=self.device)
|
||||||
|
position_ids_buf = torch.zeros(b, dtype=torch.long, device=self.device)
|
||||||
|
arange = torch.arange(max_seq_len, device=self.device)
|
||||||
|
|
||||||
|
gctx = CudaGraphContext(enabled=True)
|
||||||
|
graph_key = (b,)
|
||||||
|
|
||||||
|
def _decode_graph_step(seq_len):
|
||||||
|
input_ids_buf.copy_(
|
||||||
|
torch.randint(0, self.config.vocab_size, (b, 1), device=self.device)
|
||||||
|
)
|
||||||
|
position_ids_buf[:] = seq_len
|
||||||
|
for tid in task_ids:
|
||||||
|
pool.task_extend(tid, seq_len)
|
||||||
|
kv_cache = pool.bind_tasks(task_ids, workspace, self.device)
|
||||||
|
|
||||||
|
input_mask = torch.ge(
|
||||||
|
position_ids_buf[:, None],
|
||||||
|
arange,
|
||||||
|
out=workspace.input_mask[:b, 0, :max_seq_len],
|
||||||
|
)
|
||||||
|
input_mask = input_mask.unsqueeze(1)
|
||||||
|
|
||||||
|
with torch.inference_mode(), attn_backend(self.backend):
|
||||||
|
return gctx.forward(
|
||||||
|
self.model,
|
||||||
|
key=graph_key,
|
||||||
|
input_ids=input_ids_buf,
|
||||||
|
input_mask=input_mask,
|
||||||
|
kv_cache=kv_cache,
|
||||||
|
position_ids=position_ids_buf.unsqueeze(1),
|
||||||
|
)
|
||||||
|
|
||||||
|
for i in range(5):
|
||||||
|
_decode_graph_step(prompt_length + i)
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
|
t0 = time.perf_counter()
|
||||||
|
for i in range(gen_length * num_trials):
|
||||||
|
_decode_graph_step(prompt_length + 5 + i)
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
elapsed = time.perf_counter() - t0
|
||||||
|
tokens = batch_size * gen_length * num_trials
|
||||||
|
tps = tokens / elapsed
|
||||||
|
return BenchmarkResult(
|
||||||
|
name="decode",
|
||||||
|
batch_size=batch_size,
|
||||||
|
seq_len=gen_length,
|
||||||
|
tokens_per_second=tps,
|
||||||
|
latency_ms=elapsed / (gen_length * num_trials) * 1000,
|
||||||
|
metadata={
|
||||||
|
"benchmark_type": "decode",
|
||||||
|
"num_trials": num_trials,
|
||||||
|
"prompt_length": prompt_length,
|
||||||
|
"cuda_graph": True,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
def _run_plain_decode_benchmark(
|
||||||
|
self,
|
||||||
|
batch_size: int,
|
||||||
|
prompt_length: int,
|
||||||
|
gen_length: int,
|
||||||
|
num_trials: int,
|
||||||
) -> BenchmarkResult:
|
) -> BenchmarkResult:
|
||||||
import time
|
import time
|
||||||
|
|
||||||
# Decode grows seq_len monotonically up to prompt + 5 + gen*num_trials
|
|
||||||
# (warmup 5 steps, then one step per trial), so size the pool to cover it.
|
|
||||||
max_seq_len = prompt_length + 5 + gen_length * num_trials
|
max_seq_len = prompt_length + 5 + gen_length * num_trials
|
||||||
pool = self._make_pool(batch_size, max_seq_len)
|
pool = self._make_pool(batch_size, max_seq_len)
|
||||||
workspace = self._make_workspace(pool, self.config)
|
workspace = self._make_workspace(pool, self.config)
|
||||||
@@ -288,6 +377,11 @@ def print_benchmark_result(result: BenchmarkResult) -> None:
|
|||||||
@click.option("--num_trials", type=int, default=5, help="Number of trials.")
|
@click.option("--num_trials", type=int, default=5, help="Number of trials.")
|
||||||
@click.option("--prefill_only", is_flag=True, help="Prefill benchmark only.")
|
@click.option("--prefill_only", is_flag=True, help="Prefill benchmark only.")
|
||||||
@click.option("--decode_only", is_flag=True, help="Decode benchmark only.")
|
@click.option("--decode_only", is_flag=True, help="Decode benchmark only.")
|
||||||
|
@click.option(
|
||||||
|
"--cuda-graph",
|
||||||
|
is_flag=True,
|
||||||
|
help="Enable CUDA graph capture for decode (cuda backend only).",
|
||||||
|
)
|
||||||
@click.option(
|
@click.option(
|
||||||
"--ckpt",
|
"--ckpt",
|
||||||
required=False,
|
required=False,
|
||||||
@@ -317,6 +411,7 @@ def benchmark_command(
|
|||||||
num_trials: int,
|
num_trials: int,
|
||||||
prefill_only: bool,
|
prefill_only: bool,
|
||||||
decode_only: bool,
|
decode_only: bool,
|
||||||
|
cuda_graph: bool,
|
||||||
ckpt: Optional[str],
|
ckpt: Optional[str],
|
||||||
config_path: Optional[Path],
|
config_path: Optional[Path],
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -357,6 +452,7 @@ def benchmark_command(
|
|||||||
dtype=dtype_map[dtype],
|
dtype=dtype_map[dtype],
|
||||||
cache_type=cache,
|
cache_type=cache,
|
||||||
backend=name,
|
backend=name,
|
||||||
|
cuda_graph=cuda_graph,
|
||||||
)
|
)
|
||||||
|
|
||||||
click.secho(
|
click.secho(
|
||||||
|
|||||||
Reference in New Issue
Block a user