diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index df10504e4..222ade3b6 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -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 diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index 9fa67710b..59a7fe2e6 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -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: