[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.
|
# Internal/testing only - users should not need to change this.
|
||||||
SGLANG_PREFILL_TILE_BUDGET_MODE = EnvStr("compact")
|
SGLANG_PREFILL_TILE_BUDGET_MODE = EnvStr("compact")
|
||||||
SGLANG_PREFILL_DELAYER_MAX_PREFILL_BS_WINDOW_SIZE = EnvInt(16)
|
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
|
# Scheduler polling, timeouts, and output
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from sglang.srt.runtime_context import (
|
|||||||
get_disagg,
|
get_disagg,
|
||||||
get_schedule,
|
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")
|
_ROUTING_KEY_POLICY_DEBUG_LOG = get_bool_env_var("SGLANG_ROUTING_KEY_POLICY_DEBUG_LOG")
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -34,6 +34,7 @@ import random
|
|||||||
from collections import Counter
|
from collections import Counter
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from enum import Enum, auto
|
from enum import Enum, auto
|
||||||
|
from functools import lru_cache
|
||||||
from typing import TYPE_CHECKING, Dict, List, Optional, Set, Union
|
from typing import TYPE_CHECKING, Dict, List, Optional, Set, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -78,6 +79,19 @@ CLIP_MAX_NEW_TOKENS = int(
|
|||||||
os.environ.get("SGLANG_CLIP_MAX_NEW_TOKENS_ESTIMATION", "4096")
|
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.
|
# Threshold for in-batch prefix cache.
|
||||||
# If a request has a matched prefix length (against existing cache) less than this value,
|
# 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.
|
# 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_input_tokens = rem_input_tokens - num_mixed_decode_tokens
|
||||||
self.rem_chunk_tokens = rem_chunk_tokens
|
self.rem_chunk_tokens = rem_chunk_tokens
|
||||||
self.dllm_config = dllm_config
|
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:
|
if self.dllm_config is not None:
|
||||||
self._init_dllm_meta(dllm_config)
|
self._init_dllm_meta(dllm_config)
|
||||||
@@ -903,10 +918,24 @@ class PrefillAdder:
|
|||||||
retracted_stain: bool,
|
retracted_stain: bool,
|
||||||
mamba_gap_reserve: int = 0,
|
mamba_gap_reserve: int = 0,
|
||||||
is_chunked_continuation: bool = False,
|
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
|
# 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
|
raw_extend_input_len = extend_input_len
|
||||||
extend_input_len = self.ceil_paged_tokens(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
|
# alloc_extend reserves an extra page_size per request to make sure the budget doesn't over-commit
|
||||||
page_overhead = self.page_size
|
page_overhead = self.page_size
|
||||||
@@ -924,7 +953,7 @@ class PrefillAdder:
|
|||||||
# separately so full_evictable can't cover it — see __init__).
|
# separately so full_evictable can't cover it — see __init__).
|
||||||
if mamba_gap_reserve and self.rem_mamba_slots is not None:
|
if mamba_gap_reserve and self.rem_mamba_slots is not None:
|
||||||
self.rem_mamba_slots -= 1
|
self.rem_mamba_slots -= 1
|
||||||
self.rem_input_tokens -= extend_input_len
|
self.rem_input_tokens -= compute_charge
|
||||||
|
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
# The ring slot is reserved once at first admission; charging it
|
# The ring slot is reserved once at first admission; charging it
|
||||||
@@ -935,9 +964,9 @@ class PrefillAdder:
|
|||||||
)
|
)
|
||||||
|
|
||||||
if self.dllm_config is not None:
|
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:
|
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
|
# reprocessed_log_* is a subset of log_*; metrics_reporter subtracts it
|
||||||
# when computing the first-attempt prefix cache hit rate.
|
# when computing the first-attempt prefix cache hit rate.
|
||||||
@@ -1106,6 +1135,7 @@ class PrefillAdder:
|
|||||||
req.retracted_stain,
|
req.retracted_stain,
|
||||||
mamba_gap_reserve=self._mamba_gap_budget_for_req(req),
|
mamba_gap_reserve=self._mamba_gap_budget_for_req(req),
|
||||||
is_chunked_continuation=True,
|
is_chunked_continuation=True,
|
||||||
|
compute_charge=req.extend_range.length if self.exact_chunk_fill else None,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Return if chunked prefill not finished
|
# Return if chunked prefill not finished
|
||||||
@@ -1236,6 +1266,9 @@ class PrefillAdder:
|
|||||||
min(req.sampling_params.max_new_tokens, CLIP_MAX_NEW_TOKENS),
|
min(req.sampling_params.max_new_tokens, CLIP_MAX_NEW_TOKENS),
|
||||||
req.retracted_stain,
|
req.retracted_stain,
|
||||||
mamba_gap_reserve=self._mamba_gap_budget_for_req(req),
|
mamba_gap_reserve=self._mamba_gap_budget_for_req(req),
|
||||||
|
compute_charge=(
|
||||||
|
req.extend_range.length if self.exact_chunk_fill else None
|
||||||
|
),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
if self.rem_chunk_tokens <= 0:
|
if self.rem_chunk_tokens <= 0:
|
||||||
@@ -1259,6 +1292,7 @@ class PrefillAdder:
|
|||||||
0,
|
0,
|
||||||
req.retracted_stain,
|
req.retracted_stain,
|
||||||
mamba_gap_reserve=self._mamba_gap_budget_for_req(req),
|
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()
|
return self.budget_state()
|
||||||
@@ -1395,8 +1429,15 @@ class PrefillAdder:
|
|||||||
prefix_len = len(req.prefix_indices)
|
prefix_len = len(req.prefix_indices)
|
||||||
req.kv.cache_protected_len = prefix_len
|
req.kv.cache_protected_len = prefix_len
|
||||||
|
|
||||||
input_tokens = self.ceil_paged_tokens(
|
raw_input_tokens = len(req.full_untruncated_fill_ids) - len(
|
||||||
len(req.full_untruncated_fill_ids) - len(req.prefix_indices)
|
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 (
|
if (
|
||||||
@@ -1424,7 +1465,7 @@ class PrefillAdder:
|
|||||||
|
|
||||||
self._add_dllm_req(req, prefix_len)
|
self._add_dllm_req(req, prefix_len)
|
||||||
self._req_inc_lock_ref(req)
|
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 (
|
if (
|
||||||
tile_stop := self._check_prefill_tile_budget(input_tokens)
|
tile_stop := self._check_prefill_tile_budget(input_tokens)
|
||||||
) is not None:
|
) is not None:
|
||||||
@@ -1446,6 +1487,45 @@ class PrefillAdder:
|
|||||||
),
|
),
|
||||||
req.retracted_stain,
|
req.retracted_stain,
|
||||||
mamba_gap_reserve=mamba_gap_reserve,
|
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)
|
self._account_prefill_cache_admission(req, prefix_len)
|
||||||
else:
|
else:
|
||||||
|
|||||||
Reference in New Issue
Block a user