[Performance] Tune FlashInfer EXTEND for DP prefill (#36219)
This commit is contained in:
@@ -426,10 +426,12 @@ class BaseRunner(ABC):
|
|||||||
capture_forward_mode = ForwardMode.TARGET_VERIFY
|
capture_forward_mode = ForwardMode.TARGET_VERIFY
|
||||||
num_tokens_per_req = mr.decode_num_tokens_per_req()
|
num_tokens_per_req = mr.decode_num_tokens_per_req()
|
||||||
if extend_num_tokens_per_req is not None:
|
if extend_num_tokens_per_req is not None:
|
||||||
assert (
|
assert capture_forward_mode == ForwardMode.EXTEND and (
|
||||||
capture_forward_mode == ForwardMode.EXTEND
|
not mr.spec_algorithm.is_speculative() or _is_pd_prefill_target
|
||||||
and not mr.spec_algorithm.is_speculative()
|
), (
|
||||||
), "extend_num_tokens_per_req requires a non-speculative EXTEND dummy"
|
"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_per_req = extend_num_tokens_per_req
|
||||||
|
|
||||||
num_tokens = batch_size * 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.environ import envs
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
get_disagg,
|
||||||
get_exec,
|
get_exec,
|
||||||
get_model,
|
get_model,
|
||||||
get_spec,
|
get_spec,
|
||||||
@@ -261,19 +262,24 @@ def maybe_flashinfer_autotune_extend(
|
|||||||
if not envs.SGLANG_FLASHINFER_AUTOTUNE_EXTEND.get():
|
if not envs.SGLANG_FLASHINFER_AUTOTUNE_EXTEND.get():
|
||||||
return
|
return
|
||||||
mr = runner.model_runner
|
mr = runner.model_runner
|
||||||
# max_prefill_tokens is a per-scheduler (per dp-rank) budget, and warmup
|
# Prefer the per-rank scheduler buffer while preserving the legacy ceiling
|
||||||
# runs on all dp ranks at once, so the gathered dummy already reaches the
|
# when chunked prefill is disabled.
|
||||||
# worst-case serving gather. Do not divide by dp_size.
|
num_tokens = (
|
||||||
num_tokens = mr.server_args.max_prefill_tokens
|
mr.server_args.max_prefill_buffer_tokens() or mr.server_args.max_prefill_tokens
|
||||||
|
)
|
||||||
if num_tokens <= (decode_num_tokens or 0):
|
if num_tokens <= (decode_num_tokens or 0):
|
||||||
return # decode-shaped autotune already covered these buckets
|
return # decode-shaped autotune already covered these buckets
|
||||||
if not mr.is_generation or mr.spec_algorithm.is_speculative():
|
is_pd_prefill_target = (
|
||||||
# _dummy_run forces TARGET_VERIFY shapes for speculative runners;
|
get_disagg().disaggregation_mode == "prefill" and not mr.is_draft_worker
|
||||||
# extend-bucket autotune for spec configs is a follow-up.
|
)
|
||||||
return
|
if not mr.is_generation or (
|
||||||
if mr.model_config.is_multimodal:
|
mr.spec_algorithm.is_speculative() and not is_pd_prefill_target
|
||||||
# The dummy runs mm_inputs=None, which multimodal prefill paths iterate.
|
):
|
||||||
|
# Ordinary speculative runners force TARGET_VERIFY; PD prefill targets
|
||||||
|
# have no draft-side state and preserve the requested EXTEND mode.
|
||||||
return
|
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:
|
if mr.attn_backend.extend_dummy_seqs_capped_by_req_pool:
|
||||||
pool_size = mr.req_to_token_pool.size
|
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"]))
|
||||||
Reference in New Issue
Block a user