[Performance] Tune FlashInfer EXTEND for DP prefill (#36219)

This commit is contained in:
YAMY
2026-08-25 08:29:57 -07:00
committed by GitHub
parent b760f7fb19
commit e9c9df6a52
3 changed files with 116 additions and 14 deletions
@@ -426,10 +426,12 @@ class BaseRunner(ABC):
capture_forward_mode = ForwardMode.TARGET_VERIFY
num_tokens_per_req = mr.decode_num_tokens_per_req()
if extend_num_tokens_per_req is not None:
assert (
capture_forward_mode == ForwardMode.EXTEND
and not mr.spec_algorithm.is_speculative()
), "extend_num_tokens_per_req requires a non-speculative EXTEND dummy"
assert capture_forward_mode == ForwardMode.EXTEND and (
not mr.spec_algorithm.is_speculative() or _is_pd_prefill_target
), (
"extend_num_tokens_per_req requires an ordinary or PD-prefill "
"target EXTEND dummy"
)
num_tokens_per_req = extend_num_tokens_per_req
num_tokens = batch_size * num_tokens_per_req
@@ -26,6 +26,7 @@ import torch
from sglang.srt.environ import envs
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.runtime_context import (
get_disagg,
get_exec,
get_model,
get_spec,
@@ -261,19 +262,24 @@ def maybe_flashinfer_autotune_extend(
if not envs.SGLANG_FLASHINFER_AUTOTUNE_EXTEND.get():
return
mr = runner.model_runner
# max_prefill_tokens is a per-scheduler (per dp-rank) budget, and warmup
# runs on all dp ranks at once, so the gathered dummy already reaches the
# worst-case serving gather. Do not divide by dp_size.
num_tokens = mr.server_args.max_prefill_tokens
# Prefer the per-rank scheduler buffer while preserving the legacy ceiling
# when chunked prefill is disabled.
num_tokens = (
mr.server_args.max_prefill_buffer_tokens() or mr.server_args.max_prefill_tokens
)
if num_tokens <= (decode_num_tokens or 0):
return # decode-shaped autotune already covered these buckets
if not mr.is_generation or mr.spec_algorithm.is_speculative():
# _dummy_run forces TARGET_VERIFY shapes for speculative runners;
# extend-bucket autotune for spec configs is a follow-up.
return
if mr.model_config.is_multimodal:
# The dummy runs mm_inputs=None, which multimodal prefill paths iterate.
is_pd_prefill_target = (
get_disagg().disaggregation_mode == "prefill" and not mr.is_draft_worker
)
if not mr.is_generation or (
mr.spec_algorithm.is_speculative() and not is_pd_prefill_target
):
# Ordinary speculative runners force TARGET_VERIFY; PD prefill targets
# have no draft-side state and preserve the requested EXTEND mode.
return
# Multimodal generation wrappers can still run this text-only EXTEND dummy;
# an incompatible model should fail the explicit opt-in visibly.
if mr.attn_backend.extend_dummy_seqs_capped_by_req_pool:
pool_size = mr.req_to_token_pool.size
@@ -0,0 +1,94 @@
import sys
from types import SimpleNamespace
from unittest.mock import Mock, patch
import pytest
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode, ForwardMode
from sglang.srt.model_executor.runner import base_runner, flashinfer_autotune
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
@pytest.mark.parametrize(
"mode,error",
[
("prefill", "_dummy_run needs a static buffer"),
("decode", "ordinary or PD-prefill target EXTEND dummy"),
],
)
def test_packed_speculative_extend_is_limited_to_pd_prefill_target(mode, error):
runner = SimpleNamespace(
model_runner=SimpleNamespace(
is_draft_worker=False,
spec_algorithm=SimpleNamespace(is_speculative=lambda: True),
decode_num_tokens_per_req=lambda: 6,
)
)
with (
patch.object(
base_runner,
"get_disagg",
return_value=SimpleNamespace(disaggregation_mode=mode),
),
patch.object(
base_runner,
"get_server_return_hidden_states_mode",
return_value=CaptureHiddenMode.NULL,
),
pytest.raises(AssertionError, match=error),
):
base_runner.BaseRunner._dummy_run(
runner,
batch_size=1,
buffers=None,
forward_mode_override=ForwardMode.EXTEND,
extend_num_tokens_per_req=1,
)
def test_chunked_prefill_disabled_uses_legacy_token_ceiling():
model_runner = SimpleNamespace(
server_args=SimpleNamespace(
max_prefill_buffer_tokens=Mock(return_value=0),
max_prefill_tokens=32768,
),
is_generation=True,
is_draft_worker=False,
spec_algorithm=SimpleNamespace(is_speculative=lambda: False),
attn_backend=SimpleNamespace(extend_dummy_seqs_capped_by_req_pool=False),
canary_manager=None,
)
runner = SimpleNamespace(
model_runner=model_runner,
_alloc_dummy_decode_buffers=Mock(return_value=object()),
_dummy_run=Mock(),
)
with (
patch.object(
flashinfer_autotune.envs.SGLANG_FLASHINFER_AUTOTUNE_EXTEND,
"get",
return_value=True,
),
patch.object(
flashinfer_autotune,
"get_disagg",
return_value=SimpleNamespace(disaggregation_mode="prefill"),
),
patch.object(flashinfer_autotune, "run_flashinfer_autotune_forward"),
patch.object(flashinfer_autotune.torch.cuda, "empty_cache"),
):
flashinfer_autotune.maybe_flashinfer_autotune_extend(
runner, decode_num_tokens=128
)
runner._alloc_dummy_decode_buffers.assert_called_once_with(
32768,
num_tokens_per_req=1,
allocate_logits_buffer=False,
)
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v"]))