Fix swa input length limitation (#22597)
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user