refactor(runner): reuse a prepared static buffer for every dummy run (#28740)
This commit is contained in:
@@ -1,4 +1,5 @@
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import time
|
||||
from contextlib import contextmanager, nullcontext
|
||||
@@ -434,14 +435,28 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
|
||||
# in-seq-split). _dummy_run does not pad q/hidden like the real flow, so
|
||||
# an unaligned bs makes DSA's padded num_splits longer than the q tokens
|
||||
# and trips FlashMLA's "num_splits must have shape (b+1)" check.
|
||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||
from sglang.srt.layers.utils.cp_utils import get_cp_padding_align_size
|
||||
from sglang.srt.utils.common import require_mlp_sync
|
||||
|
||||
n_sms = torch.cuda.get_device_properties(model_runner.device).multi_processor_count
|
||||
block_m = 64
|
||||
cp = max(get_cp_padding_align_size(), 1)
|
||||
|
||||
attn_tp_size = get_attention_tp_size()
|
||||
mlp_sync = require_mlp_sync(model_runner.server_args)
|
||||
|
||||
def _align(bs: int) -> int:
|
||||
# Align to lcm(cp, attn_tp_size) so the CP multiple isn't undone by a
|
||||
# later attn_tp align (e.g. cp=2, attn_tp=3: 128 -> 128 -> 129).
|
||||
align = cp
|
||||
if mlp_sync and attn_tp_size > 1:
|
||||
align = math.lcm(cp, attn_tp_size)
|
||||
return ceil_align(bs, align)
|
||||
|
||||
batch_sizes = sorted(
|
||||
{
|
||||
ceil_align(bs, cp)
|
||||
_align(bs)
|
||||
for bs in (
|
||||
1,
|
||||
2 * block_m,
|
||||
@@ -467,16 +482,24 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
|
||||
disagg_mode,
|
||||
)
|
||||
|
||||
# One buffer set sized to the largest shape, reused across the sweep
|
||||
# (the decode runner's max_bs is too small for n_sms*block_m).
|
||||
dummy_buffers = runner._alloc_dummy_decode_buffers(max(batch_sizes))
|
||||
|
||||
t0 = time.perf_counter()
|
||||
with torch.inference_mode():
|
||||
for bs in batch_sizes:
|
||||
if run_decode:
|
||||
runner._dummy_run(
|
||||
batch_size=bs, forward_mode_override=ForwardMode.DECODE
|
||||
batch_size=bs,
|
||||
forward_mode_override=ForwardMode.DECODE,
|
||||
buffers=dummy_buffers,
|
||||
)
|
||||
if run_extend:
|
||||
runner._dummy_run(
|
||||
batch_size=bs, forward_mode_override=ForwardMode.EXTEND
|
||||
batch_size=bs,
|
||||
forward_mode_override=ForwardMode.EXTEND,
|
||||
buffers=dummy_buffers,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
|
||||
@@ -22,7 +22,7 @@ import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
from typing import TYPE_CHECKING, Any, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
@@ -52,7 +52,6 @@ from sglang.srt.utils import (
|
||||
require_gathered_buffer,
|
||||
require_mlp_tp_gather,
|
||||
)
|
||||
from sglang.srt.utils.common import ceil_align, require_mlp_sync
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
@@ -210,7 +209,11 @@ class BaseRunner(ABC):
|
||||
self._pre_initialize_flashinfer_allreduce_workspace()
|
||||
|
||||
if self._should_run_flashinfer_autotune():
|
||||
self._flashinfer_autotune()
|
||||
buffers, batch_size = self._autotune_buffers()
|
||||
assert (
|
||||
buffers is not None
|
||||
), "_autotune_buffers() must return a reusable buffer set for autotune"
|
||||
self._flashinfer_autotune(buffers=buffers, batch_size=batch_size)
|
||||
|
||||
if (
|
||||
envs.SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP.get()
|
||||
@@ -313,8 +316,14 @@ class BaseRunner(ABC):
|
||||
|
||||
return True
|
||||
|
||||
def _flashinfer_autotune(self):
|
||||
"""Run flashinfer autotune."""
|
||||
def _flashinfer_autotune(self, *, buffers, batch_size):
|
||||
"""Run flashinfer autotune.
|
||||
|
||||
buffers / batch_size: a prepared static decode-buffer set and its bs,
|
||||
reused for the dummy forward instead of allocating a throwaway set.
|
||||
Supplied by warmup() (the decode runner's captured buffers when a graph
|
||||
runner exists; a freshly-allocated dummy set in the eager path).
|
||||
"""
|
||||
from flashinfer.autotuner import autotune
|
||||
|
||||
from sglang.srt.layers.logits_processor import autotune_dummy_run_mode
|
||||
@@ -346,7 +355,7 @@ class BaseRunner(ABC):
|
||||
autotune(True, cache=str(autotune_cache)),
|
||||
autotune_dummy_run_mode(),
|
||||
):
|
||||
self._dummy_run(batch_size=mr.req_to_token_pool.size)
|
||||
self._dummy_run(batch_size=batch_size, buffers=buffers)
|
||||
torch.cuda.current_stream().wait_stream(mr.forward_stream)
|
||||
logger.info("FlashInfer autotune completed.")
|
||||
|
||||
@@ -386,16 +395,66 @@ class BaseRunner(ABC):
|
||||
cache_dir / f"rank_tp{mr.tp_rank}_pp{mr.pp_rank}_dp{mr.dp_rank or 0}.json"
|
||||
)
|
||||
|
||||
def _alloc_dummy_decode_buffers(self, max_bs: int, *, num_tokens_per_bs: int = 1):
|
||||
"""Allocate one static decode-buffer set for a dummy forward, sized to
|
||||
(max_bs, max_bs * num_tokens_per_bs).
|
||||
|
||||
The PP-parallel DeepGEMM warmup sweeps batch sizes far larger than any
|
||||
runner's max_bs (up to ~n_sms*block_m), so no pre-allocated runner buffer
|
||||
set fits; it builds one here and hands it to _dummy_run (reused across the
|
||||
sweep; _dummy_run slices it per shape). The flashinfer autotune does NOT
|
||||
use this -- it reuses an existing runner's buffers via _autotune_buffers
|
||||
(the eager input registry, or the decode cuda-graph runner's captured
|
||||
buffers).
|
||||
"""
|
||||
mr = self.model_runner
|
||||
return _allocate_decode_buffers(
|
||||
device=mr.device,
|
||||
max_bs=max_bs,
|
||||
max_num_token=max_bs * num_tokens_per_bs,
|
||||
hidden_size=mr.model_config.hidden_size,
|
||||
vocab_size=mr.model_config.vocab_size,
|
||||
dtype=mr.model_config.dtype,
|
||||
dp_size=mr.server_args.dp_size,
|
||||
pp_size=mr.server_args.pp_size,
|
||||
is_encoder_decoder=mr.model_config.is_encoder_decoder,
|
||||
require_mlp_tp_gather=require_mlp_tp_gather(mr.server_args),
|
||||
seq_len_fill_value=mr.attn_backend.get_cuda_graph_seq_len_fill_value(),
|
||||
encoder_len_fill_value=(
|
||||
getattr(mr.model_config.hf_config, "max_source_positions", 0)
|
||||
if mr.model_config.is_encoder_decoder
|
||||
else 0
|
||||
),
|
||||
num_tokens_per_bs=num_tokens_per_bs,
|
||||
cache_loc_dtype=torch.int64,
|
||||
enable_mamba_track=False,
|
||||
hc_hidden_size=getattr(mr.model_config, "hc_hidden_size", None),
|
||||
)
|
||||
|
||||
def _dummy_run(
|
||||
self,
|
||||
batch_size: int,
|
||||
run_ctx=None,
|
||||
forward_mode_override: Optional[ForwardMode] = None,
|
||||
*,
|
||||
buffers,
|
||||
):
|
||||
"""Run a dummy forward pass for warmup/profiling.
|
||||
|
||||
forward_mode_override forces EXTEND/DECODE regardless of
|
||||
is_generation (used by the PP-parallel DeepGEMM warmup).
|
||||
|
||||
buffers: a prepared static buffer set (or lightweight adapter exposing
|
||||
the same fields), sized >= this dummy shape, which _dummy_run slices to
|
||||
(batch_size, num_tokens). The caller owns the shape and the allocation --
|
||||
the flashinfer autotune reuses an existing runner's buffers via
|
||||
_autotune_buffers (the eager input registry, or the decode cuda-graph
|
||||
runner's captured buffers); the PP-DeepGEMM warmup builds one via
|
||||
_alloc_dummy_decode_buffers. _dummy_run never allocates and never re-pads
|
||||
(autotune must run at the reused shape; the PP warmup pre-pads and sizes
|
||||
its buffer to match). next_token_logits_buffer is optional -- a live
|
||||
autotune forward returns logits fresh, so the eager-reuse path passes
|
||||
None (only the PP warmup set still carries one).
|
||||
"""
|
||||
mr = self.model_runner
|
||||
if forward_mode_override is not None:
|
||||
@@ -422,12 +481,22 @@ class BaseRunner(ABC):
|
||||
|
||||
num_tokens = batch_size * num_tokens_per_bs
|
||||
|
||||
# Keep warmup aligned with scheduler MLP-sync padding.
|
||||
if require_mlp_sync(mr.server_args):
|
||||
attn_tp_size = get_parallel().attn_tp_size
|
||||
if attn_tp_size > 1 and num_tokens % attn_tp_size != 0:
|
||||
num_tokens = ceil_align(num_tokens, attn_tp_size)
|
||||
batch_size = num_tokens // num_tokens_per_bs
|
||||
# Caller owns the shape: passes a static buffer >= the dummy shape; no
|
||||
# allocation, no re-padding (would overflow the reused buffers).
|
||||
assert (
|
||||
buffers is not None
|
||||
and num_tokens <= buffers.input_ids.shape[0]
|
||||
and batch_size <= buffers.seq_lens.shape[0]
|
||||
), (
|
||||
f"_dummy_run needs a static buffer >= (num_tokens={num_tokens}, "
|
||||
f"batch_size={batch_size}); got "
|
||||
+ (
|
||||
"None"
|
||||
if buffers is None
|
||||
else f"(input_ids={buffers.input_ids.shape[0]}, "
|
||||
f"seq_lens={buffers.seq_lens.shape[0]})"
|
||||
)
|
||||
)
|
||||
|
||||
seq_len_fill_value = mr.attn_backend.get_cuda_graph_seq_len_fill_value()
|
||||
|
||||
@@ -451,28 +520,25 @@ class BaseRunner(ABC):
|
||||
if require_gathered_buffer(mr.server_args):
|
||||
assert require_mlp_tp_gather_ or require_attn_tp_gather(mr.server_args)
|
||||
|
||||
buffers = _allocate_decode_buffers(
|
||||
device=mr.device,
|
||||
max_bs=batch_size,
|
||||
max_num_token=num_tokens,
|
||||
hidden_size=mr.model_config.hidden_size,
|
||||
vocab_size=mr.model_config.vocab_size,
|
||||
dtype=mr.model_config.dtype,
|
||||
dp_size=mr.server_args.dp_size,
|
||||
pp_size=mr.server_args.pp_size,
|
||||
is_encoder_decoder=mr.model_config.is_encoder_decoder,
|
||||
require_mlp_tp_gather=require_mlp_tp_gather_,
|
||||
seq_len_fill_value=seq_len_fill_value,
|
||||
encoder_len_fill_value=(
|
||||
getattr(mr.model_config.hf_config, "max_source_positions", 0)
|
||||
if mr.model_config.is_encoder_decoder
|
||||
else 0
|
||||
),
|
||||
num_tokens_per_bs=num_tokens_per_bs,
|
||||
cache_loc_dtype=torch.int64,
|
||||
enable_mamba_track=False,
|
||||
hc_hidden_size=getattr(mr.model_config, "hc_hidden_size", None),
|
||||
input_ids = buffers.input_ids[:num_tokens]
|
||||
positions = buffers.positions[:num_tokens]
|
||||
out_cache_loc = buffers.out_cache_loc[:num_tokens]
|
||||
# Eager-reuse drops the logits buffer; only buffer sets that carry one slice it.
|
||||
next_token_logits_buffer = (
|
||||
buffers.next_token_logits_buffer[:num_tokens]
|
||||
if buffers.next_token_logits_buffer is not None
|
||||
else None
|
||||
)
|
||||
mrope_positions = buffers.mrope_positions[:, :num_tokens]
|
||||
req_pool_indices = buffers.req_pool_indices[:batch_size]
|
||||
seq_lens = buffers.seq_lens[:batch_size]
|
||||
seq_lens_cpu = buffers.seq_lens_cpu[:batch_size]
|
||||
encoder_lens = (
|
||||
buffers.encoder_lens[:batch_size]
|
||||
if buffers.encoder_lens is not None
|
||||
else None
|
||||
)
|
||||
|
||||
buffers.num_token_non_padded[...] = num_tokens
|
||||
|
||||
# For extend mode
|
||||
@@ -558,17 +624,17 @@ class BaseRunner(ABC):
|
||||
forward_batch = ForwardBatch(
|
||||
forward_mode=capture_forward_mode,
|
||||
batch_size=batch_size,
|
||||
input_ids=buffers.input_ids,
|
||||
req_pool_indices=buffers.req_pool_indices,
|
||||
seq_lens=buffers.seq_lens,
|
||||
seq_lens_cpu=buffers.seq_lens_cpu,
|
||||
next_token_logits_buffer=buffers.next_token_logits_buffer,
|
||||
orig_seq_lens=buffers.seq_lens,
|
||||
out_cache_loc=buffers.out_cache_loc,
|
||||
seq_lens_sum=buffers.seq_lens.sum().item(),
|
||||
encoder_lens=buffers.encoder_lens,
|
||||
input_ids=input_ids,
|
||||
req_pool_indices=req_pool_indices,
|
||||
seq_lens=seq_lens,
|
||||
seq_lens_cpu=seq_lens_cpu,
|
||||
next_token_logits_buffer=next_token_logits_buffer,
|
||||
orig_seq_lens=seq_lens,
|
||||
out_cache_loc=out_cache_loc,
|
||||
seq_lens_sum=seq_lens.sum().item(),
|
||||
encoder_lens=encoder_lens,
|
||||
return_logprob=False,
|
||||
positions=buffers.positions,
|
||||
positions=positions,
|
||||
extend_num_tokens=extend_num_tokens,
|
||||
extend_seq_lens=extend_seq_lens,
|
||||
extend_prefix_lens=extend_prefix_lens,
|
||||
@@ -580,7 +646,7 @@ class BaseRunner(ABC):
|
||||
global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu,
|
||||
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
|
||||
global_dp_buffer_len=global_dp_buffer_len,
|
||||
mrope_positions=buffers.mrope_positions,
|
||||
mrope_positions=mrope_positions,
|
||||
spec_algorithm=mr.spec_algorithm,
|
||||
spec_info=spec_info,
|
||||
capture_hidden_mode=capture_hidden_mode,
|
||||
@@ -616,7 +682,7 @@ class BaseRunner(ABC):
|
||||
kwargs["get_embedding"] = True
|
||||
|
||||
logits_output_or_pp_proxy_tensors = mr.model.forward(
|
||||
buffers.input_ids,
|
||||
input_ids,
|
||||
forward_batch.positions,
|
||||
forward_batch,
|
||||
**kwargs,
|
||||
@@ -629,6 +695,11 @@ class BaseRunner(ABC):
|
||||
with torch.inference_mode(), run_ctx or empty_context():
|
||||
run_once()
|
||||
|
||||
def _autotune_buffers(self) -> Tuple[Optional[Any], Optional[int]]:
|
||||
"""Return (buffers, bs) for the autotune dummy forward to reuse; the
|
||||
EagerRunner and DecodeCudaGraphRunner override this."""
|
||||
return None, None
|
||||
|
||||
@abstractmethod
|
||||
def can_run_graph(self, forward_batch: ForwardBatch) -> bool: ...
|
||||
|
||||
|
||||
@@ -358,6 +358,20 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
f"Capture cuda graph failed: {e}\n" f"{CUDA_GRAPH_CAPTURE_FAILED_MSG}"
|
||||
)
|
||||
|
||||
def _autotune_buffers(self):
|
||||
"""Reuse these static decode buffers (sized to max_bs) for the warmup
|
||||
flashinfer-autotune dummy forward instead of allocating a throwaway set
|
||||
— see BaseRunner._autotune_buffers / BaseRunner._dummy_run.
|
||||
|
||||
The dummy forward derives its shape from max_bs and must match these
|
||||
buffers exactly; _dummy_run asserts that. Every autotune-reachable
|
||||
decode shape (plain decode, spec target-verify) matches. DLLM would not
|
||||
(its buffers hold block_size tokens/bs while the dummy run derives 1),
|
||||
but DLLM does not use a flashinfer MoE backend, so autotune never runs
|
||||
for it and this is never reached there.
|
||||
"""
|
||||
return self.buffers, self.max_bs
|
||||
|
||||
def maybe_init_pdmux(self):
|
||||
if self.enable_pdmux:
|
||||
self.stream_groups = get_stream_groups()
|
||||
|
||||
@@ -18,7 +18,8 @@ from __future__ import annotations
|
||||
import contextlib
|
||||
import logging
|
||||
from dataclasses import replace
|
||||
from typing import TYPE_CHECKING, Any, Tuple, Union
|
||||
from types import SimpleNamespace
|
||||
from typing import TYPE_CHECKING, Any, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
|
||||
@@ -41,7 +42,7 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
|
||||
enable_tc_piecewise_cuda_graph,
|
||||
set_tc_piecewise_forward_context,
|
||||
)
|
||||
from sglang.srt.utils import is_hip
|
||||
from sglang.srt.utils import is_hip, require_mlp_tp_gather
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -99,6 +100,8 @@ class EagerRunner(BaseRunner):
|
||||
else mr.max_total_num_tokens
|
||||
)
|
||||
max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_bs)
|
||||
self._eager_max_bs = max_bs
|
||||
self._eager_num_tokens_per_bs = num_tokens_per_bs
|
||||
is_encoder_decoder = mr.model_config.is_encoder_decoder
|
||||
self._eager_registry = build_eager_registry(
|
||||
device=mr.device,
|
||||
@@ -119,6 +122,83 @@ class EagerRunner(BaseRunner):
|
||||
# Eager has no capture step, so warm up here (run-once via mr._kernel_warmed_up).
|
||||
self.warmup()
|
||||
|
||||
def _autotune_buffers(self) -> Tuple[Any, int]:
|
||||
"""Adapter over the eager registry for the autotune dummy forward; fills
|
||||
in fields the registry omits (logits buffer, pp_proxy, custom_mask)."""
|
||||
mr = self.model_runner
|
||||
reg = self._eager_registry
|
||||
max_bs = self._eager_max_bs
|
||||
|
||||
def _slot(name):
|
||||
return reg.get_slot(name).buffer if reg.has_slot(name) else None
|
||||
|
||||
# num_token_non_padded / global_num_tokens_* are not registered on the
|
||||
# eager registry (build_eager_registry passes enable_num_token_non_padded
|
||||
# =False, register_global_num_tokens=False); _dummy_run writes + reads
|
||||
# them unconditionally, so supply tiny fresh tensors here.
|
||||
num_token_non_padded = torch.zeros((1,), dtype=torch.int32, device=mr.device)
|
||||
global_dim = (
|
||||
mr.server_args.dp_size if require_mlp_tp_gather(mr.server_args) else 1
|
||||
)
|
||||
global_num_tokens_gpu = torch.zeros(
|
||||
(global_dim,), dtype=torch.int32, device=mr.device
|
||||
)
|
||||
global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||
(global_dim,), dtype=torch.int32, device=mr.device
|
||||
)
|
||||
|
||||
# custom_mask: only consumed by create_dummy_verify_input (spec). Size it
|
||||
# like the decode path's custom_mask for a spec target worker.
|
||||
custom_mask: Optional[torch.Tensor] = None
|
||||
if mr.spec_algorithm.is_speculative():
|
||||
num_tokens_per_bs = self._eager_num_tokens_per_bs
|
||||
max_num_token = reg.max_num_tokens
|
||||
seq_len_fill_value = mr.attn_backend.get_cuda_graph_seq_len_fill_value()
|
||||
custom_mask = torch.ones(
|
||||
(max_bs * seq_len_fill_value + max_num_token) * num_tokens_per_bs,
|
||||
dtype=torch.bool,
|
||||
device=mr.device,
|
||||
)
|
||||
|
||||
# pp_proxy_tensors: only read when pp_size>1. _dummy_run slices each value
|
||||
# [:pp_hidden_tokens] (pp_hidden_tokens <= num_tokens), so size the first
|
||||
# dim to the registry's token ceiling. Mirror _allocate_decode_buffers'
|
||||
# keys/dtypes (mHC flattens residual into hidden_states of hc_hidden_size).
|
||||
pp_proxy_tensors = None
|
||||
if mr.server_args.pp_size > 1:
|
||||
hidden_size = mr.model_config.hidden_size
|
||||
hc_hidden_size = getattr(mr.model_config, "hc_hidden_size", None)
|
||||
is_mhc = hc_hidden_size is not None
|
||||
hs = hc_hidden_size if is_mhc else hidden_size
|
||||
rows = reg.max_num_tokens
|
||||
pp_proxy_tensors = {
|
||||
"hidden_states": torch.zeros(
|
||||
(rows, hs), dtype=mr.dtype, device=mr.device
|
||||
),
|
||||
}
|
||||
if not is_mhc:
|
||||
pp_proxy_tensors["residual"] = torch.zeros(
|
||||
(rows, hidden_size), dtype=mr.dtype, device=mr.device
|
||||
)
|
||||
|
||||
adapter = SimpleNamespace(
|
||||
input_ids=_slot("input_ids"),
|
||||
positions=_slot("positions"),
|
||||
out_cache_loc=_slot("out_cache_loc"),
|
||||
req_pool_indices=_slot("req_pool_indices"),
|
||||
seq_lens=_slot("seq_lens"),
|
||||
seq_lens_cpu=_slot("seq_lens_cpu"),
|
||||
mrope_positions=_slot("mrope_positions"),
|
||||
encoder_lens=_slot("encoder_lens"),
|
||||
next_token_logits_buffer=None,
|
||||
num_token_non_padded=num_token_non_padded,
|
||||
global_num_tokens_gpu=global_num_tokens_gpu,
|
||||
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
|
||||
custom_mask=custom_mask,
|
||||
pp_proxy_tensors=pp_proxy_tensors,
|
||||
)
|
||||
return adapter, max_bs
|
||||
|
||||
def can_run_graph(self, forward_batch: ForwardBatch) -> bool:
|
||||
# Eager never runs a cuda graph; callers dispatch on isinstance(...,
|
||||
# EagerRunner) and must not route an eager batch into a replay branch.
|
||||
|
||||
Reference in New Issue
Block a user