Fix swa input length limitation (#22597)

This commit is contained in:
Ke Bao
2026-04-12 16:03:35 +08:00
committed by GitHub
parent f2377a00cb
commit bc1bfbf607
2 changed files with 60 additions and 21 deletions
+59 -20
View File
@@ -428,6 +428,8 @@ class PrefillAdder:
)
self.is_hybrid_ssm_cache = self.tree_cache.supports_mamba()
self.rem_swa_token_offset = 0
self.priority_scheduling_preemption_threshold = (
priority_scheduling_preemption_threshold
)
@@ -456,11 +458,9 @@ class PrefillAdder:
@property
def rem_total_tokens(self):
if self.is_hybrid_swa:
available_and_evictable = min(
available_and_evictable = (
self.token_to_kv_pool_allocator.full_available_size()
+ self.tree_cache.full_evictable_size(),
self.token_to_kv_pool_allocator.swa_available_size()
+ self.tree_cache.swa_evictable_size(),
+ self.tree_cache.full_evictable_size()
)
elif self.is_hybrid_ssm_cache:
available_and_evictable = (
@@ -474,14 +474,20 @@ class PrefillAdder:
)
return available_and_evictable - self.rem_total_token_offset
@property
def rem_swa_tokens(self):
return (
self.token_to_kv_pool_allocator.swa_available_size()
+ self.tree_cache.swa_evictable_size()
- self.rem_swa_token_offset
)
@property
def cur_rem_tokens(self):
if self.is_hybrid_swa:
available_and_evictable = min(
available_and_evictable = (
self.token_to_kv_pool_allocator.full_available_size()
+ self.tree_cache.full_evictable_size(),
self.token_to_kv_pool_allocator.swa_available_size()
+ self.tree_cache.swa_evictable_size(),
+ self.tree_cache.full_evictable_size()
)
elif self.is_hybrid_ssm_cache:
available_and_evictable = (
@@ -496,11 +502,31 @@ class PrefillAdder:
return available_and_evictable - self.cur_rem_token_offset
def _swa_budget_for_req(self, extend_input_len: int) -> int:
"""SWA pool budget per request. Only valid when is_hybrid_swa is True.
With chunked prefill + overlap scheduler, the peak SWA occupancy is:
chunk N (running, not yet in tree) + sliding window (locked in tree)
+ chunk N+1 (new allocation)
Since chunk N and locked tokens are already excluded from
swa_available + swa_evictable, the budget only needs to cover the
chunk N+1 allocation. We floor at sliding_window_size to reserve
room for the decode phase.
"""
if self.rem_chunk_tokens is not None:
alloc = min(extend_input_len, self.rem_chunk_tokens)
else:
alloc = extend_input_len
return max(alloc, self.tree_cache.sliding_window_size) + self.page_size
def ceil_paged_tokens(self, tokens: int) -> int:
return -(-tokens // self.page_size) * self.page_size
def budget_state(self):
if self.rem_total_tokens <= 0 or self.cur_rem_tokens <= 0:
no_token = self.rem_total_tokens <= 0 or self.cur_rem_tokens <= 0
if not no_token and self.is_hybrid_swa:
no_token = self.rem_swa_tokens <= 0
if no_token:
return AddReqResult.NO_TOKEN
if self.rem_input_tokens <= 0:
@@ -527,6 +553,9 @@ class PrefillAdder:
self.cur_rem_token_offset += extend_input_len + page_overhead
self.rem_input_tokens -= extend_input_len
if self.is_hybrid_swa:
self.rem_swa_token_offset += self._swa_budget_for_req(extend_input_len)
if self.dllm_config is not None:
self.rem_dllm_tokens -= extend_input_len
elif self.rem_chunk_tokens is not None:
@@ -601,6 +630,8 @@ class PrefillAdder:
_rem_tokens = self._get_dllm_remain_tokens()
else:
_rem_tokens = min(self.rem_chunk_tokens, int(self.rem_total_tokens))
if self.is_hybrid_swa:
_rem_tokens = min(_rem_tokens, int(self.rem_swa_tokens))
# The chunked_req must be added to the list; otherwise, it will cause a memory leak.
# Therefore, in certain cases where _rem_tokens <= 0, it should be replaced with rem_chunk_tokens.
if _rem_tokens <= 0:
@@ -639,11 +670,12 @@ class PrefillAdder:
self.tree_cache.dec_lock_ref(last_node)
def add_one_req_ignore_eos(self, req: Req):
# Early exit if no enough tokens for the input tokens
if self.ceil_paged_tokens(req.extend_input_len) > min(
self.cur_rem_tokens, self.rem_total_tokens
):
paged_input = self.ceil_paged_tokens(req.extend_input_len)
if paged_input > min(self.cur_rem_tokens, self.rem_total_tokens):
return AddReqResult.NO_TOKEN
if self.is_hybrid_swa:
if self._swa_budget_for_req(req.extend_input_len) > self.rem_swa_tokens:
return AddReqResult.NO_TOKEN
def add_req_state(r, insert_sort=False):
new_token_ratio = (
@@ -756,14 +788,11 @@ class PrefillAdder:
# _update_prefill_budget already accounts for this in the deduction.
# Without this, admission is more optimistic than the actual budget
# deduction, allowing over-admission when the pool is nearly full.
total_tokens = (
req.extend_input_len
+ min(
max(req.sampling_params.max_new_tokens - len(req.output_ids), 0),
CLIP_MAX_NEW_TOKENS,
)
+ self.page_size
max_new = min(
max(req.sampling_params.max_new_tokens - len(req.output_ids), 0),
CLIP_MAX_NEW_TOKENS,
)
total_tokens = req.extend_input_len + max_new + self.page_size
# adjusting the input_tokens based on host_hit_length and page_size
real_input_tokens = req.extend_input_len - req.host_hit_length
@@ -773,6 +802,11 @@ class PrefillAdder:
if total_tokens >= self.rem_total_tokens:
return AddReqResult.NO_TOKEN
if self.is_hybrid_swa:
swa_needed = self._swa_budget_for_req(req.extend_input_len)
if swa_needed >= self.rem_swa_tokens:
return AddReqResult.NO_TOKEN
if real_input_tokens >= self.rem_input_tokens and len(self.can_run_list) != 0:
return AddReqResult.OTHER
@@ -781,6 +815,11 @@ class PrefillAdder:
if total_tokens >= self.rem_total_tokens:
return AddReqResult.NO_TOKEN
if self.is_hybrid_swa:
swa_needed = self._swa_budget_for_req(req.extend_input_len)
if swa_needed >= self.rem_swa_tokens:
return AddReqResult.NO_TOKEN
if req.host_hit_length > 0:
new_indices, req.last_node = self.tree_cache.init_load_back(
InitLoadBackParams(
@@ -1961,7 +1961,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
def max_token_pool_size(self):
"""Return the max token pool size considering hybrid swa settings."""
if self.is_hybrid_swa:
return min(self.swa_max_total_num_tokens, self.max_total_num_tokens)
return self.full_max_total_num_tokens
else:
return self.max_total_num_tokens