diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index fcbaa6f0b..ca2030a16 100644 --- a/python/sglang/srt/model_executor/runner/base_runner.py +++ b/python/sglang/srt/model_executor/runner/base_runner.py @@ -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 diff --git a/python/sglang/srt/model_executor/runner/flashinfer_autotune.py b/python/sglang/srt/model_executor/runner/flashinfer_autotune.py index fea98b385..decfcdd56 100644 --- a/python/sglang/srt/model_executor/runner/flashinfer_autotune.py +++ b/python/sglang/srt/model_executor/runner/flashinfer_autotune.py @@ -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 diff --git a/test/registered/unit/model_executor/runner/test_flashinfer_autotune.py b/test/registered/unit/model_executor/runner/test_flashinfer_autotune.py new file mode 100644 index 000000000..b31b88ce8 --- /dev/null +++ b/test/registered/unit/model_executor/runner/test_flashinfer_autotune.py @@ -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"]))