[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