[AMD][gfx95] Fill the chunked-prefill compute budget exactly (#32888)

Co-authored-by: Thomas Wang <thomawan@amd.com>
This commit is contained in:
Jacob0226
2026-09-14 00:35:15 -07:00
committed by GitHub
co-authored by Thomas Wang
parent 3eeb7d37f9
commit 5200508b0f
2 changed files with 91 additions and 7 deletions
+4
View File
@@ -602,6 +602,10 @@ class Envs:
# Internal/testing only - users should not need to change this.
SGLANG_PREFILL_TILE_BUDGET_MODE = EnvStr("compact")
SGLANG_PREFILL_DELAYER_MAX_PREFILL_BS_WINDOW_SIZE = EnvInt(16)
# Charge the chunked-prefill compute budget in tokens, not page-ceiled
# tokens, so a prefill batch runs exactly chunked_prefill_size and the dense
# GEMMs get an aligned M. gfx95 only; see PrefillAdder.exact_chunk_fill.
SGLANG_EXACT_CHUNK_FILL = EnvBool(True)
# ===================================================================
# Scheduler polling, timeouts, and output
+87 -7
View File
@@ -9,7 +9,7 @@ from sglang.srt.runtime_context import (
get_disagg,
get_schedule,
)
from sglang.srt.utils import get_bool_env_var, is_hip
from sglang.srt.utils import get_bool_env_var, is_gfx95_supported, is_hip
_ROUTING_KEY_POLICY_DEBUG_LOG = get_bool_env_var("SGLANG_ROUTING_KEY_POLICY_DEBUG_LOG")
logger = logging.getLogger(__name__)
@@ -34,6 +34,7 @@ import random
from collections import Counter
from contextlib import contextmanager
from enum import Enum, auto
from functools import lru_cache
from typing import TYPE_CHECKING, Dict, List, Optional, Set, Union
import torch
@@ -78,6 +79,19 @@ CLIP_MAX_NEW_TOKENS = int(
os.environ.get("SGLANG_CLIP_MAX_NEW_TOKENS_ESTIMATION", "4096")
)
@lru_cache(maxsize=1)
def _use_exact_chunk_fill() -> bool:
"""Whether to charge the chunked-prefill compute budget in tokens (gfx95 only).
Gated on gfx95 because that is where the win is: the aiter absorb bmm picks
its EVEN_MN specialization from M % BLOCK_SIZE_M, so a short batch costs 2x
on that kernel, against 1% for the hipBLASLt GEMMs that just run one extra
partial tile.
"""
return envs.SGLANG_EXACT_CHUNK_FILL.get() and is_gfx95_supported()
# Threshold for in-batch prefix cache.
# If a request has a matched prefix length (against existing cache) less than this value,
# the scheduler runs the in-batch prefix caching check for this request.
@@ -565,6 +579,7 @@ class PrefillAdder:
self.rem_input_tokens = rem_input_tokens - num_mixed_decode_tokens
self.rem_chunk_tokens = rem_chunk_tokens
self.dllm_config = dllm_config
self.exact_chunk_fill = _use_exact_chunk_fill() and dllm_config is None
if self.dllm_config is not None:
self._init_dllm_meta(dllm_config)
@@ -903,10 +918,24 @@ class PrefillAdder:
retracted_stain: bool,
mamba_gap_reserve: int = 0,
is_chunked_continuation: bool = False,
compute_charge: Optional[int] = None,
):
"""Charge one admitted request against the prefill budgets.
`compute_charge` is what the compute budgets (`rem_chunk_tokens`,
`rem_input_tokens`, `rem_dllm_tokens`) are billed, counted in
forward-pass tokens; the KV budgets are always billed page-ceiled
tokens. It defaults to the ceiled count, so only the exact-chunk-fill
path parts from upstream behaviour. Both compute budgets have to take
it: they are usually configured to the same value, so leaving either one
ceiled makes it hit zero first and stop admission with the rounding
slack unspent.
"""
# TODO(lsyin): check this workaround logic, which only ensures the prefill will not out of memory, and may be too conservative
raw_extend_input_len = extend_input_len
extend_input_len = self.ceil_paged_tokens(extend_input_len)
if compute_charge is None:
compute_charge = extend_input_len
# alloc_extend reserves an extra page_size per request to make sure the budget doesn't over-commit
page_overhead = self.page_size
@@ -924,7 +953,7 @@ class PrefillAdder:
# separately so full_evictable can't cover it — see __init__).
if mamba_gap_reserve and self.rem_mamba_slots is not None:
self.rem_mamba_slots -= 1
self.rem_input_tokens -= extend_input_len
self.rem_input_tokens -= compute_charge
if self.is_hybrid_swa:
# The ring slot is reserved once at first admission; charging it
@@ -935,9 +964,9 @@ class PrefillAdder:
)
if self.dllm_config is not None:
self.rem_dllm_tokens -= extend_input_len
self.rem_dllm_tokens -= compute_charge
elif self.rem_chunk_tokens is not None:
self.rem_chunk_tokens -= extend_input_len
self.rem_chunk_tokens -= compute_charge
# reprocessed_log_* is a subset of log_*; metrics_reporter subtracts it
# when computing the first-attempt prefix cache hit rate.
@@ -1106,6 +1135,7 @@ class PrefillAdder:
req.retracted_stain,
mamba_gap_reserve=self._mamba_gap_budget_for_req(req),
is_chunked_continuation=True,
compute_charge=req.extend_range.length if self.exact_chunk_fill else None,
)
# Return if chunked prefill not finished
@@ -1236,6 +1266,9 @@ class PrefillAdder:
min(req.sampling_params.max_new_tokens, CLIP_MAX_NEW_TOKENS),
req.retracted_stain,
mamba_gap_reserve=self._mamba_gap_budget_for_req(req),
compute_charge=(
req.extend_range.length if self.exact_chunk_fill else None
),
)
else:
if self.rem_chunk_tokens <= 0:
@@ -1259,6 +1292,7 @@ class PrefillAdder:
0,
req.retracted_stain,
mamba_gap_reserve=self._mamba_gap_budget_for_req(req),
compute_charge=trunc_len if self.exact_chunk_fill else None,
)
return self.budget_state()
@@ -1395,8 +1429,15 @@ class PrefillAdder:
prefix_len = len(req.prefix_indices)
req.kv.cache_protected_len = prefix_len
input_tokens = self.ceil_paged_tokens(
len(req.full_untruncated_fill_ids) - len(req.prefix_indices)
raw_input_tokens = len(req.full_untruncated_fill_ids) - len(
req.prefix_indices
)
input_tokens = self.ceil_paged_tokens(raw_input_tokens)
# Whether the request fits whole. Against the raw length under
# exact-chunk-fill, so a request whose ceiled length would spill is
# not needlessly split into a second chunk.
chunk_fit_tokens = (
raw_input_tokens if self.exact_chunk_fill else input_tokens
)
if (
@@ -1424,7 +1465,7 @@ class PrefillAdder:
self._add_dllm_req(req, prefix_len)
self._req_inc_lock_ref(req)
elif chunk_tokens_limit is None or input_tokens <= chunk_tokens_limit:
elif chunk_tokens_limit is None or chunk_fit_tokens <= chunk_tokens_limit:
if (
tile_stop := self._check_prefill_tile_budget(input_tokens)
) is not None:
@@ -1446,6 +1487,45 @@ class PrefillAdder:
),
req.retracted_stain,
mamba_gap_reserve=mamba_gap_reserve,
compute_charge=raw_input_tokens if self.exact_chunk_fill else None,
)
self._account_prefill_cache_admission(req, prefix_len)
elif self.exact_chunk_fill:
# Take the remainder verbatim so the batch hits exactly
# chunked_prefill_size. `chunk_fit_tokens > chunk_tokens_limit`
# here, so this never runs past the end of the prompt. Uses the
# limit rather than rem_chunk_tokens so an SWA-capped chunk stays
# capped.
trunc_len = chunk_tokens_limit
if trunc_len <= 0:
return AddReqResult.OTHER
if truncation_align_size is not None:
if trunc_len < truncation_align_size:
return AddReqResult.OTHER
trunc_len = truncation_align_size * (
trunc_len // truncation_align_size
)
if (
tile_stop := self._check_prefill_tile_budget(trunc_len)
) is not None:
return tile_stop
req.set_extend_range(
len(req.prefix_indices), len(req.prefix_indices) + trunc_len
)
self.can_run_list.append(req)
self.new_chunked_req = req
self._req_inc_lock_ref(req)
self._update_prefill_budget(
prefix_len,
trunc_len,
0,
req.retracted_stain,
mamba_gap_reserve=mamba_gap_reserve,
compute_charge=trunc_len,
)
self._account_prefill_cache_admission(req, prefix_len)
else: