Share one logits output buffer across prefill/decode/draft cuda-graph runners (#29779)

This commit is contained in:
cctry
2026-07-01 23:30:56 -07:00
committed by GitHub
parent 926140d789
commit 697b400d70
9 changed files with 130 additions and 23 deletions
+4 -2
View File
@@ -980,8 +980,10 @@ class LogitsProcessor(nn.Module):
def _copy_logits_to_buffer(
self, logits: torch.Tensor, logits_metadata: LogitsMetadata
) -> torch.Tensor:
if logits_metadata.next_token_logits_buffer is not None:
logits_buffer = logits_metadata.next_token_logits_buffer
logits_buffer = logits_metadata.next_token_logits_buffer
# The shared logits buffer is keyed by vocab width; skip it when this
# model's vocab doesn't match (e.g. hot-vocab draft vs full-vocab target).
if logits_buffer is not None and logits_buffer.shape[-1] == self.vocab_size:
assert logits_buffer.dtype == torch.float
logits_buffer.copy_(logits[:, : self.vocab_size])
logits = logits_buffer
@@ -0,0 +1,66 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Dict, Optional
import torch
from sglang.srt.model_executor.cuda_graph_config import Backend
if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner
class GraphSharedOutput:
"""``(max_rows, vocab)`` logits buffer, shared by every cuda-graph runner."""
_process_shared: Optional[GraphSharedOutput] = None
def __init__(
self,
*,
device: torch.device,
max_rows: int,
) -> None:
self.device = torch.device(device)
self.max_rows = max_rows
self._logits_buffers: Dict[int, torch.Tensor] = {}
@classmethod
def create_for_model_runner(
cls, model_runner: ModelRunner
) -> Optional[GraphSharedOutput]:
cuda_graph_config = model_runner.server_args.cuda_graph_config
if cuda_graph_config is None:
return None
max_rows = 0
decode = cuda_graph_config.decode
if decode.backend != Backend.DISABLED and decode.bs:
max_rows = max(max_rows, model_runner.max_decode_logits_rows())
if max_rows <= 0:
return None
device = torch.device(model_runner.device)
shared = cls._process_shared
if (
shared is not None
and shared.device == device
and shared.max_rows >= max_rows
):
return shared
cls._process_shared = cls(device=device, max_rows=max_rows)
return cls._process_shared
def get_logits_buffer(self, vocab_size: int, *, rows: int) -> torch.Tensor:
assert rows <= self.max_rows, (
f"shared logits buffer holds {self.max_rows} rows but caller "
f"needs {rows} (vocab_size={vocab_size})"
)
buffer = self._logits_buffers.get(vocab_size)
if buffer is None:
buffer = torch.zeros(
(self.max_rows, vocab_size), dtype=torch.float, device=self.device
)
self._logits_buffers[vocab_size] = buffer
return buffer[:rows]
@@ -83,6 +83,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory,
)
from sglang.srt.distributed.parallel_state import monkey_patch_vllm_parallel_state
from sglang.srt.dllm.config import DllmConfig
from sglang.srt.elastic_ep.elastic_ep import (
ElasticEPStateManager,
join_process_groups,
@@ -159,6 +160,7 @@ from sglang.srt.model_executor.forward_context import (
forward_context,
has_forward_context,
)
from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput
from sglang.srt.model_executor.hook_manager import register_forward_hooks
from sglang.srt.model_executor.model_runner_kv_cache_mixin import (
ModelRunnerKVCacheMixin,
@@ -818,6 +820,25 @@ class ModelRunner(ModelRunnerKVCacheMixin):
return None
return getattr(hf_config, "index_topk", None)
def decode_num_tokens_per_bs(
self, *, num_draft_tokens: Optional[int] = None
) -> int:
"""Logits rows per decode batch slot."""
if self.spec_algorithm.is_speculative():
if num_draft_tokens is None:
num_draft_tokens = self.server_args.speculative_num_draft_tokens
return self.spec_algorithm.get_num_tokens_per_bs_for_target_verify(
num_draft_tokens, self.is_draft_worker
)
dllm_config = DllmConfig.from_server_args(self.server_args)
return dllm_config.block_size if dllm_config is not None else 1
def max_decode_logits_rows(self) -> int:
"""Rows the shared logits buffer needs."""
num_tokens_per_bs = self.decode_num_tokens_per_bs()
capture_bs, _ = get_batch_sizes_to_capture(self, num_tokens_per_bs)
return max(capture_bs) * num_tokens_per_bs
def alloc_memory_pool(self, memory_pool_config: Optional[MemoryPoolConfig] = None):
"""Allocate KV cache memory pools only (no backends or cuda graphs)."""
if memory_pool_config is not None:
@@ -869,6 +890,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.decode_cuda_graph_runner = None
self.graph_mem_usage = 0
self.prefill_cuda_graph_runner = None
self.graph_shared_output = None
def init_attention_backends(self):
"""Initialize attention backends only (no cuda graph capture)."""
@@ -905,6 +927,8 @@ class ModelRunner(ModelRunnerKVCacheMixin):
because they capture their own decode-style graphs separately.
"""
self.graph_shared_output = GraphSharedOutput.create_for_model_runner(self)
# The eager (no-cuda-graph) phase runner, built AFTER the attention
# backend so its __init__ can warm up kernels (run-once) and allocate the
# fixed-max static buffer — both before the cuda-graph runners, so that
@@ -239,7 +239,9 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
# --- capture mode + tokens-per-bs ------------------------------
self.capture_forward_mode = ForwardMode.DECODE
self.capture_hidden_mode = CaptureHiddenMode.NULL
self.num_tokens_per_bs = 1
self.num_tokens_per_bs = model_runner.decode_num_tokens_per_bs(
num_draft_tokens=self.speculative_num_draft_tokens
)
if model_runner.spec_algorithm.is_speculative():
if self.model_runner.is_draft_worker:
# Draft workers can use TARGET_VERIFY mode.
@@ -248,14 +250,8 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
):
raise RuntimeError("This should not happen")
self.capture_forward_mode = ForwardMode.TARGET_VERIFY
self.num_tokens_per_bs = (
model_runner.spec_algorithm.get_num_tokens_per_bs_for_target_verify(
self.speculative_num_draft_tokens, model_runner.is_draft_worker
)
)
elif self.is_dllm:
self.capture_forward_mode = ForwardMode.DLLM_EXTEND
self.num_tokens_per_bs = self.dllm_config.block_size
# --- bucket sizes ---------------------------------------------
self.capture_bs, self.compile_bs = get_batch_sizes_to_capture(
@@ -314,7 +310,9 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
max_bs=self.max_bs,
max_num_token=self.max_num_token,
hidden_size=self.model_runner.model_config.hidden_size,
vocab_size=self.model_runner.model_config.vocab_size,
next_token_logits_buffer=self.model_runner.graph_shared_output.get_logits_buffer(
self.model_runner.model_config.vocab_size, rows=self.max_num_token
),
dtype=self.model_runner.model_config.dtype,
dp_size=self.dp_size,
pp_size=self.pp_size,
@@ -348,6 +348,17 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
def _cache_loc_dtype(self):
return torch.int64 if not is_npu() else torch.int32
def _next_token_logits_buffer(self, rows: int) -> Optional[torch.Tensor]:
if not self.model_runner.pp_group.is_last_rank:
return None
graph_shared_output = self.model_runner.graph_shared_output
# Fall back to eager logits when the shared buffer can't hold prefill rows.
if graph_shared_output is None or rows > graph_shared_output.max_rows:
return None
return graph_shared_output.get_logits_buffer(
self.model_runner.model_config.vocab_size, rows=rows
)
_aiter_chip_info_cached = False
@classmethod
@@ -627,7 +638,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
),
req_pool_indices=shape_inputs["req_pool_indices"],
seq_lens=shape_inputs["seq_lens"],
next_token_logits_buffer=None,
next_token_logits_buffer=self._next_token_logits_buffer(bs),
orig_seq_lens=shape_inputs["orig_seq_lens"],
seq_lens_cpu=torch.tensor([num_tokens], device="cpu"),
out_cache_loc=_slot("out_cache_loc"),
@@ -823,7 +834,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
input_embeds=input_embeds,
req_pool_indices=forward_batch.req_pool_indices,
seq_lens=forward_batch.seq_lens,
next_token_logits_buffer=None,
next_token_logits_buffer=self._next_token_logits_buffer(bs),
orig_seq_lens=forward_batch.orig_seq_lens,
seq_lens_cpu=forward_batch.seq_lens_cpu,
out_cache_loc=out_cache_loc,
@@ -92,7 +92,7 @@ class DecodeInputBuffers(ForwardInputBuffers):
max_bs: int,
max_num_token: int,
hidden_size: int,
vocab_size: int,
next_token_logits_buffer: torch.Tensor,
dtype: torch.dtype,
dp_size: int,
pp_size: int,
@@ -120,10 +120,6 @@ class DecodeInputBuffers(ForwardInputBuffers):
(max_bs * seq_len_fill_value + max_num_token) * num_tokens_per_bs,
dtype=torch.bool,
)
next_token_logits_buffer = torch.zeros(
(max_num_token, vocab_size),
dtype=torch.float,
)
mamba_track_indices = (
torch.zeros((max_bs,), dtype=torch.int64)
if enable_mamba_track
@@ -223,12 +223,10 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
else:
vocab_size = self.model_runner.model_config.vocab_size
next_token_logits_buffer = torch.zeros(
(
self.max_bs * self.num_tokens_per_bs,
vocab_size,
),
dtype=torch.float,
next_token_logits_buffer = (
self.model_runner.graph_shared_output.get_logits_buffer(
vocab_size, rows=self.max_bs * self.num_tokens_per_bs
)
)
seq_lens_cpu = torch.full(
@@ -17,6 +17,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import set_global_server_args_for_scheduler
@@ -403,6 +404,11 @@ class MockModelRunner(ModelRunner):
self.is_hybrid_swa = case.sliding_window_size is not None
self.sliding_window_size = case.sliding_window_size
self.use_mla_backend = False
# Runner-mode helpers mutate speculative graph sizes after construction.
self.graph_shared_output = GraphSharedOutput(
device=self.device,
max_rows=pool_batch_size * max_context_len,
)
@property
def hybrid_gdn_config(self):
@@ -21,6 +21,7 @@ from sglang.srt.model_executor.forward_context import (
forward_context,
get_token_to_kv_pool,
)
from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import get_parallel
from sglang.srt.server_args import set_global_server_args_for_scheduler
@@ -308,6 +309,11 @@ class MockMLAModelRunner(ModelRunner):
self.use_mla_backend = True
self.is_draft_worker = False
self._kernel_warmed_up = True
# Runner-mode helpers mutate speculative graph sizes after construction.
self.graph_shared_output = GraphSharedOutput(
device=self.device,
max_rows=pool_batch_size * max_context_len,
)
@property
def hybrid_gdn_config(self):