[AMD][gfx95] Fill the chunked-prefill compute budget exactly (#32888)
Co-authored-by: Thomas Wang <thomawan@amd.com>
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user