diff --git a/astrai/extension/attention_backend.py b/astrai/extension/attention_backend.py index a672303..ef7149e 100644 --- a/astrai/extension/attention_backend.py +++ b/astrai/extension/attention_backend.py @@ -511,6 +511,7 @@ class CudaBackend(AttentionBackend): is_causal=True, o_part_buf=kv_cache.decode_o_part, ml_part_buf=kv_cache.decode_ml_part, + out_buf=kv_cache.decode_out, ) return out.unsqueeze(1).flatten(2) diff --git a/astrai/extension/attention_ops.py b/astrai/extension/attention_ops.py index 590faa4..ca8cae8 100644 --- a/astrai/extension/attention_ops.py +++ b/astrai/extension/attention_ops.py @@ -102,6 +102,7 @@ def attn_paged_decode( is_causal: bool = False, o_part_buf: Optional[torch.Tensor] = None, ml_part_buf: Optional[torch.Tensor] = None, + out_buf: Optional[torch.Tensor] = None, ) -> torch.Tensor: """SGLang-style paged decode (q_len == 1, flat KV pool). @@ -121,6 +122,7 @@ def attn_paged_decode( is_causal: apply causal mask o_part_buf: pre-allocated split-KV o partial 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: [batch, n_heads, head_dim] (bf16, 3D) @@ -139,6 +141,7 @@ def attn_paged_decode( causal_offset=causal_offset, o_part_buf=o_part_buf, ml_part_buf=ml_part_buf, + out_buf=out_buf, ) diff --git a/astrai/inference/core/cache.py b/astrai/inference/core/cache.py index cc4b5e7..86caa70 100644 --- a/astrai/inference/core/cache.py +++ b/astrai/inference/core/cache.py @@ -263,6 +263,7 @@ class KVCache: qo_indptr: [batch+1] int32 — prefill qo prefix-sum (None in decode) decode_o_part: split-KV o 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 @@ -276,6 +277,7 @@ class KVCache: qo_indptr: Optional[Tensor] = None decode_o_part: Optional[Tensor] = None decode_ml_part: Optional[Tensor] = None + decode_out: Optional[Tensor] = None class PagePool: @@ -571,6 +573,7 @@ class PagePool: ) qo_indptr = workspace.qo_indptr[: b + 1] decode_o_part, decode_ml_part = None, None + decode_out = None else: write_pos = seq_lens_t - 1 loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1) @@ -579,6 +582,7 @@ class PagePool: qo_indptr = None decode_o_part = getattr(workspace, "decode_o_part", None) decode_ml_part = getattr(workspace, "decode_ml_part", None) + decode_out = getattr(workspace, "decode_out", None) return KVCache( k_buffer=self._storage.k_buffer, @@ -592,6 +596,7 @@ class PagePool: qo_indptr=qo_indptr, decode_o_part=decode_o_part, decode_ml_part=decode_ml_part, + decode_out=decode_out, ) # ---- internals ---- diff --git a/astrai/inference/core/graph.py b/astrai/inference/core/graph.py index 8692ca1..85ad18b 100644 --- a/astrai/inference/core/graph.py +++ b/astrai/inference/core/graph.py @@ -86,12 +86,14 @@ class CudaGraphContext: if key in self._graphs: self._graphs[key].replay() elif key in self._warmed: + cap_output = model(**kwargs) 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) + return cap_output else: self._warmed.add(key) self._outputs[key] = model(**kwargs) diff --git a/astrai/inference/core/workspace.py b/astrai/inference/core/workspace.py index 4bca49f..7fc08e7 100644 --- a/astrai/inference/core/workspace.py +++ b/astrai/inference/core/workspace.py @@ -108,6 +108,14 @@ class InferenceWorkspace: 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): """Return ``(o_part, ml_part)`` view sliced to live dimensions.""" return ( diff --git a/csrc/kernels/attn_paged_decode.cu b/csrc/kernels/attn_paged_decode.cu index 682ebf3..ba6c963 100644 --- a/csrc/kernels/attn_paged_decode.cu +++ b/csrc/kernels/attn_paged_decode.cu @@ -13,7 +13,8 @@ torch::Tensor attn_paged_decode( int64_t causal_offset, double scale, c10::optional o_part_buf, - c10::optional ml_part_buf + c10::optional ml_part_buf, + c10::optional out_buf ) { const at::cuda::OptionalCUDAGuard device_guard(device_of(q)); auto stream = at::cuda::getCurrentCUDAStream(); @@ -23,7 +24,18 @@ torch::Tensor attn_paged_decode( req_to_token, req_pool_indices, kv_indptr, 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(); 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("o_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."); } diff --git a/scripts/tools/benchmark.py b/scripts/tools/benchmark.py index ca1307c..7623754 100644 --- a/scripts/tools/benchmark.py +++ b/scripts/tools/benchmark.py @@ -9,6 +9,7 @@ from astrai import setup_logging from astrai.config import BaseModelConfig, ConfigFactory from astrai.extension import ATTN_BACKEND, AttentionBackendFactory, attn_backend from astrai.inference.core.cache import PagePool +from astrai.inference.core.graph import CudaGraphContext from astrai.inference.core.workspace import InferenceWorkspace from astrai.model import AutoModel, AutoRegressiveLM @@ -57,6 +58,7 @@ class GenerationBenchmark: dtype: torch.dtype = torch.bfloat16, cache_type: str = "contiguous", backend: Union[str, ATTN_BACKEND] = ATTN_BACKEND.CUDA, + cuda_graph: bool = False, ): self.device = device self.dtype = dtype @@ -64,6 +66,7 @@ class GenerationBenchmark: self.model = model self.config = config self.backend = backend + self.cuda_graph = cuda_graph def _make_pool(self, batch_size: int, max_seq_len: int) -> PagePool: return PagePool( @@ -217,11 +220,97 @@ class GenerationBenchmark: prompt_length: int = 512, gen_length: int = 128, 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: 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 pool = self._make_pool(batch_size, max_seq_len) 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("--prefill_only", is_flag=True, help="Prefill 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( "--ckpt", required=False, @@ -317,6 +411,7 @@ def benchmark_command( num_trials: int, prefill_only: bool, decode_only: bool, + cuda_graph: bool, ckpt: Optional[str], config_path: Optional[Path], ) -> None: @@ -357,6 +452,7 @@ def benchmark_command( dtype=dtype_map[dtype], cache_type=cache, backend=name, + cuda_graph=cuda_graph, ) click.secho(