From 6f67ba8942ff40d1a00827cd2aa68e4d62e56405 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Thu, 6 Aug 2026 19:07:26 +0800 Subject: [PATCH] perf: move decode split partials to InferenceWorkspace - Replace per-.cu-file static cached tensors with workspace-managed pre-allocated buffers - InferenceWorkspace now owns decode_o_part / decode_ml_part (mirrors FlashInfer's workspace pattern) - KVCache carries the buffers through the backend -> C++ kernel chain - C++ kernels accept optional pre-allocated buffers; fallback to alloc_split_partials for backward compat - Pre-allocates once at Executor init, zero allocation in the decode hot loop - Prerequisite for CUDA-graph capture (all kernel addresses are stable) --- astrai/extension/attention_backend.py | 2 + astrai/extension/attention_ops.py | 6 +++ astrai/inference/core/cache.py | 10 +++++ astrai/inference/core/executor.py | 5 +++ astrai/inference/core/workspace.py | 47 ++++++++++++++++----- csrc/kernels/attn_decode.cu | 24 ++++++----- csrc/kernels/attn_paged_decode.cu | 24 ++++++----- scripts/tools/benchmark.py | 13 ++++-- tests/extension/test_backend_equivalence.py | 7 ++- tests/inference/test_cache.py | 7 ++- 10 files changed, 109 insertions(+), 36 deletions(-) diff --git a/astrai/extension/attention_backend.py b/astrai/extension/attention_backend.py index 133afb9..4448a0c 100644 --- a/astrai/extension/attention_backend.py +++ b/astrai/extension/attention_backend.py @@ -441,6 +441,8 @@ class CudaBackend(AttentionBackend): kv_indptr, kv_cache.max_len, is_causal=True, + o_part_buf=kv_cache.decode_o_part, + ml_part_buf=kv_cache.decode_ml_part, ) return out.unsqueeze(1).flatten(2) diff --git a/astrai/extension/attention_ops.py b/astrai/extension/attention_ops.py index 60fab6c..590faa4 100644 --- a/astrai/extension/attention_ops.py +++ b/astrai/extension/attention_ops.py @@ -100,6 +100,8 @@ def attn_paged_decode( max_seq_len: int, mask: Optional[torch.Tensor] = None, is_causal: bool = False, + o_part_buf: Optional[torch.Tensor] = None, + ml_part_buf: Optional[torch.Tensor] = None, ) -> torch.Tensor: """SGLang-style paged decode (q_len == 1, flat KV pool). @@ -117,6 +119,8 @@ def attn_paged_decode( max_seq_len: max per-request seq_len (Python int, for split computation) mask: 2D [batch, max_seq_len] (bool, True=keep) or None 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) Returns: [batch, n_heads, head_dim] (bf16, 3D) @@ -133,6 +137,8 @@ def attn_paged_decode( max_seq_len, mask=mask, causal_offset=causal_offset, + o_part_buf=o_part_buf, + ml_part_buf=ml_part_buf, ) diff --git a/astrai/inference/core/cache.py b/astrai/inference/core/cache.py index bd6fd9d..cc4b5e7 100644 --- a/astrai/inference/core/cache.py +++ b/astrai/inference/core/cache.py @@ -260,6 +260,9 @@ class KVCache: max_len: max(seq_lens) as Python int — avoids GPU sync in decode kv_indptr: [batch+1] int32 — prefix sum of seq_lens, precomputed once per step so the attention backend avoids rebuilding it per layer. + 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) """ k_buffer: Tensor @@ -271,6 +274,8 @@ class KVCache: max_len: int = 0 kv_indptr: Optional[Tensor] = None qo_indptr: Optional[Tensor] = None + decode_o_part: Optional[Tensor] = None + decode_ml_part: Optional[Tensor] = None class PagePool: @@ -565,12 +570,15 @@ class PagePool: torch.arange(b + 1, dtype=torch.int32, device=device) * q_len ) qo_indptr = workspace.qo_indptr[: b + 1] + decode_o_part, decode_ml_part = None, None else: write_pos = seq_lens_t - 1 loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1) ocl_buf[:b].copy_(loc) out_cache_loc = ocl_buf[:b] qo_indptr = None + decode_o_part = getattr(workspace, "decode_o_part", None) + decode_ml_part = getattr(workspace, "decode_ml_part", None) return KVCache( k_buffer=self._storage.k_buffer, @@ -582,6 +590,8 @@ class PagePool: max_len=max(seq_lens), kv_indptr=kv_indptr, qo_indptr=qo_indptr, + decode_o_part=decode_o_part, + decode_ml_part=decode_ml_part, ) # ---- internals ---- diff --git a/astrai/inference/core/executor.py b/astrai/inference/core/executor.py index b436c43..5f3bd3c 100644 --- a/astrai/inference/core/executor.py +++ b/astrai/inference/core/executor.py @@ -78,9 +78,14 @@ class Executor: # (input_ids, decode mask, KV bind metadata). Eagerly sized at init # so the workspace is CUDA-graph-capture friendly — no allocation # during capture. + config = model.config + max_q_heads = config.num_attention_heads + head_dim = config.hidden_size // config.num_attention_heads self._workspace = InferenceWorkspace( max_batch_size=kv_cache.max_batch_size, max_seq_len=kv_cache.max_seq_len, + max_q_heads=max_q_heads, + head_dim=head_dim, device=self.device, dtype=self.dtype, ) diff --git a/astrai/inference/core/workspace.py b/astrai/inference/core/workspace.py index edd3959..4eb93a6 100644 --- a/astrai/inference/core/workspace.py +++ b/astrai/inference/core/workspace.py @@ -1,20 +1,16 @@ """Pre-allocated buffers for the inference decode hot path. -Mirrors SGLang's pre-allocated input buffers (``input_buffers.py``): tensors -are sized once to the server's maximum dimensions and sliced to the live -batch each step, so the per-token decode loop never calls -``torch.empty``/``torch.zeros``/``torch.arange`` for the hot shapes. Fills -go through ``out=`` variants (``torch.ge``) which write into the stable -buffers instead of allocating fresh results. - -All buffers are allocated eagerly at init (nothing is lazy), so the -workspace is CUDA-graph-capture friendly: the decode step reads/writes -fixed-address tensors with no allocation during capture. +Mirrors FlashInfer / SGLang's global workspace pattern: all per-step tensors +are allocated eagerly at init (nothing is lazy), so the decode step +reads/writes fixed-address tensors with zero ``torch.empty`` calls during +the hot loop — a prerequisite for CUDA-graph capture. """ import torch from torch import Tensor +_MAX_SPLITS = 32 + class InferenceWorkspace: """Reusable fixed-shape per-step buffers for decode. @@ -30,6 +26,11 @@ class InferenceWorkspace: - KV-cache bind metadata (``req_pool_indices``, ``seq_lens``, ``kv_indptr``, ``inc``, ``out_cache_loc``), written by ``PagePool.bind_tasks`` when the Executor passes this workspace. + - ``decode_o_part`` / ``decode_ml_part``: split-KV partial result buffers + (mirrors FlashInfer's workspace). One global alloc, reused by every + decode step across all layers. Sliced views are passed to the CUDA + attention kernel so its internal ``torch.empty`` hot-path alloc goes + through a stable address (CUDA-graph capturable). No re-allocation while the server's bounds are respected. """ @@ -38,11 +39,15 @@ class InferenceWorkspace: self, max_batch_size: int, max_seq_len: int, + max_q_heads: int, + head_dim: int, device: torch.device, dtype: torch.dtype, ): self.max_batch_size = max_batch_size self.max_seq_len = max_seq_len + self.max_q_heads = max_q_heads + self.head_dim = head_dim self.device = device self.dtype = dtype @@ -83,6 +88,28 @@ class InferenceWorkspace: (max_batch_size, 1), dtype=torch.long, device=device ) + # Split-KV partial-result buffers for decode (persistent, one global + # alloc per process — mirrors FlashInfer's workspace pattern). + # Shape: [max_batch_size, max_q_heads, _MAX_SPLITS, head_dim] (o_part) + # [max_batch_size, max_q_heads, _MAX_SPLITS, 2] (ml_part) + self.decode_o_part = torch.empty( + (max_batch_size, max_q_heads, _MAX_SPLITS, head_dim), + dtype=torch.float32, + device=device, + ) + self.decode_ml_part = torch.empty( + (max_batch_size, max_q_heads, _MAX_SPLITS, 2), + dtype=torch.float32, + device=device, + ) + + def decode_buffers(self, batch: int, q_heads: int): + """Return ``(o_part, ml_part)`` view sliced to live dimensions.""" + return ( + self.decode_o_part[:batch, :q_heads], + self.decode_ml_part[:batch, :q_heads], + ) + def fill_input_ids(self, ids: "list[int]") -> Tensor: """Write ``ids`` into the device buffer and return ``[B]``. diff --git a/csrc/kernels/attn_decode.cu b/csrc/kernels/attn_decode.cu index 7023ced..d85844e 100644 --- a/csrc/kernels/attn_decode.cu +++ b/csrc/kernels/attn_decode.cu @@ -8,7 +8,9 @@ torch::Tensor attn_decode( c10::optional mask, int64_t causal_offset, double scale, - int64_t layout + int64_t layout, + torch::Tensor o_part_buf, + torch::Tensor ml_part_buf ) { const at::cuda::OptionalCUDAGuard device_guard(device_of(q)); auto stream = at::cuda::getCurrentCUDAStream(); @@ -22,16 +24,16 @@ torch::Tensor attn_decode( auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O; p.o = (bf16*)O_view.data_ptr(); - { - static torch::Tensor s_o_part, s_ml_part; + if (o_part_buf.defined() && ml_part_buf.defined()) { + TORCH_CHECK(o_part_buf.scalar_type() == torch::kFloat32, "o_part_buf must be f32"); + TORCH_CHECK(ml_part_buf.scalar_type() == torch::kFloat32, "ml_part_buf must be f32"); int64_t o_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * p.head_dim; - auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); - if (!s_o_part.defined() || s_o_part.numel() < o_needed) { - s_o_part = torch::empty({p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt); - s_ml_part = torch::empty({p.batch, p.q_head, MAX_SPLITS, 2}, fopt); - } - p.o_part = (float*)s_o_part.data_ptr(); - p.ml_part = (float*)s_ml_part.data_ptr(); + TORCH_CHECK(o_part_buf.numel() >= o_needed, + "o_part_buf too small: need ", o_needed, " got ", o_part_buf.numel()); + p.o_part = (float*)o_part_buf.data_ptr(); + p.ml_part = (float*)ml_part_buf.data_ptr(); + } else { + alloc_split_partials(p); } DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p, stream); C10_CUDA_CHECK(cudaGetLastError()); @@ -47,5 +49,7 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { py::arg("causal_offset") = -1, py::arg("scale") = 0.0, py::arg("layout") = (int64_t)BHLD, + py::arg("o_part_buf") = py::none(), + py::arg("ml_part_buf") = py::none(), "GQA decode (tensor-core head-packing on sm_80+, scalar fallback)"); } diff --git a/csrc/kernels/attn_paged_decode.cu b/csrc/kernels/attn_paged_decode.cu index 4e54f38..d6baed4 100644 --- a/csrc/kernels/attn_paged_decode.cu +++ b/csrc/kernels/attn_paged_decode.cu @@ -11,7 +11,9 @@ torch::Tensor attn_paged_decode( int64_t max_seq_len, c10::optional mask, int64_t causal_offset, - double scale + double scale, + torch::Tensor o_part_buf, + torch::Tensor ml_part_buf ) { const at::cuda::OptionalCUDAGuard device_guard(device_of(q)); auto stream = at::cuda::getCurrentCUDAStream(); @@ -24,16 +26,16 @@ torch::Tensor attn_paged_decode( auto O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options()); p.o = (bf16*)O.data_ptr(); - { - static torch::Tensor s_o_part, s_ml_part; + if (o_part_buf.defined() && ml_part_buf.defined()) { + TORCH_CHECK(o_part_buf.scalar_type() == torch::kFloat32, "o_part_buf must be f32"); + TORCH_CHECK(ml_part_buf.scalar_type() == torch::kFloat32, "ml_part_buf must be f32"); int64_t o_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * p.head_dim; - auto fopt = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA); - if (!s_o_part.defined() || s_o_part.numel() < o_needed) { - s_o_part = torch::empty({p.batch, p.q_head, MAX_SPLITS, p.head_dim}, fopt); - s_ml_part = torch::empty({p.batch, p.q_head, MAX_SPLITS, 2}, fopt); - } - p.o_part = (float*)s_o_part.data_ptr(); - p.ml_part = (float*)s_ml_part.data_ptr(); + TORCH_CHECK(o_part_buf.numel() >= o_needed, + "o_part_buf too small: need ", o_needed, " got ", o_part_buf.numel()); + p.o_part = (float*)o_part_buf.data_ptr(); + p.ml_part = (float*)ml_part_buf.data_ptr(); + } else { + alloc_split_partials(p); } DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p, stream); C10_CUDA_CHECK(cudaGetLastError()); @@ -52,5 +54,7 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { py::arg("mask") = py::none(), py::arg("causal_offset") = -1, py::arg("scale") = 0.0, + py::arg("o_part_buf") = py::none(), + py::arg("ml_part_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 eb35aa8..ca1307c 100644 --- a/scripts/tools/benchmark.py +++ b/scripts/tools/benchmark.py @@ -79,9 +79,14 @@ class GenerationBenchmark: ) @staticmethod - def _make_workspace(pool: PagePool) -> InferenceWorkspace: + def _make_workspace(pool: PagePool, config: BaseModelConfig) -> InferenceWorkspace: return InferenceWorkspace( - pool.max_batch_size, pool.max_seq_len, pool.device, pool.dtype + pool.max_batch_size, + pool.max_seq_len, + max_q_heads=config.num_attention_heads, + head_dim=config.hidden_size // config.num_attention_heads, + device=pool.device, + dtype=pool.dtype, ) def _run_prefill( @@ -156,7 +161,7 @@ class GenerationBenchmark: import time pool = self._make_pool(batch_size, prompt_length) - workspace = self._make_workspace(pool) + workspace = self._make_workspace(pool, self.config) task_ids = [f"bench_prefill_{i}" for i in range(batch_size)] for tid in task_ids: pool.task_alloc(tid, list(range(prompt_length))) @@ -219,7 +224,7 @@ class GenerationBenchmark: # (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) + workspace = self._make_workspace(pool, self.config) task_ids = self._run_prefill(pool, batch_size, prompt_length, workspace) for i in range(5): diff --git a/tests/extension/test_backend_equivalence.py b/tests/extension/test_backend_equivalence.py index acac283..981889f 100644 --- a/tests/extension/test_backend_equivalence.py +++ b/tests/extension/test_backend_equivalence.py @@ -14,7 +14,12 @@ from tests.extension.conftest import D, skip_no_kernel def _ws(pool: PagePool) -> InferenceWorkspace: return InferenceWorkspace( - pool.max_batch_size, pool.max_seq_len, pool.device, pool.dtype + pool.max_batch_size, + pool.max_seq_len, + max_q_heads=2, + head_dim=64, + device=pool.device, + dtype=pool.dtype, ) diff --git a/tests/inference/test_cache.py b/tests/inference/test_cache.py index bfdfa29..d83c91c 100644 --- a/tests/inference/test_cache.py +++ b/tests/inference/test_cache.py @@ -16,7 +16,12 @@ from astrai.inference.core.workspace import InferenceWorkspace def _ws(pool: PagePool) -> InferenceWorkspace: """Workspace sized to the pool (bind_tasks requires it).""" return InferenceWorkspace( - pool.max_batch_size, pool.max_seq_len, pool.device, pool.dtype + pool.max_batch_size, + pool.max_seq_len, + max_q_heads=2, + head_dim=4, + device=pool.device, + dtype=pool.dtype, )