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.is_hybrid_ssm_cache = self.tree_cache.supports_mamba()
|
||||||
|
|
||||||
|
self.rem_swa_token_offset = 0
|
||||||
|
|
||||||
self.priority_scheduling_preemption_threshold = (
|
self.priority_scheduling_preemption_threshold = (
|
||||||
priority_scheduling_preemption_threshold
|
priority_scheduling_preemption_threshold
|
||||||
)
|
)
|
||||||
@@ -456,11 +458,9 @@ class PrefillAdder:
|
|||||||
@property
|
@property
|
||||||
def rem_total_tokens(self):
|
def rem_total_tokens(self):
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
available_and_evictable = min(
|
available_and_evictable = (
|
||||||
self.token_to_kv_pool_allocator.full_available_size()
|
self.token_to_kv_pool_allocator.full_available_size()
|
||||||
+ self.tree_cache.full_evictable_size(),
|
+ self.tree_cache.full_evictable_size()
|
||||||
self.token_to_kv_pool_allocator.swa_available_size()
|
|
||||||
+ self.tree_cache.swa_evictable_size(),
|
|
||||||
)
|
)
|
||||||
elif self.is_hybrid_ssm_cache:
|
elif self.is_hybrid_ssm_cache:
|
||||||
available_and_evictable = (
|
available_and_evictable = (
|
||||||
@@ -474,14 +474,20 @@ class PrefillAdder:
|
|||||||
)
|
)
|
||||||
return available_and_evictable - self.rem_total_token_offset
|
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
|
@property
|
||||||
def cur_rem_tokens(self):
|
def cur_rem_tokens(self):
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
available_and_evictable = min(
|
available_and_evictable = (
|
||||||
self.token_to_kv_pool_allocator.full_available_size()
|
self.token_to_kv_pool_allocator.full_available_size()
|
||||||
+ self.tree_cache.full_evictable_size(),
|
+ self.tree_cache.full_evictable_size()
|
||||||
self.token_to_kv_pool_allocator.swa_available_size()
|
|
||||||
+ self.tree_cache.swa_evictable_size(),
|
|
||||||
)
|
)
|
||||||
elif self.is_hybrid_ssm_cache:
|
elif self.is_hybrid_ssm_cache:
|
||||||
available_and_evictable = (
|
available_and_evictable = (
|
||||||
@@ -496,11 +502,31 @@ class PrefillAdder:
|
|||||||
|
|
||||||
return available_and_evictable - self.cur_rem_token_offset
|
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:
|
def ceil_paged_tokens(self, tokens: int) -> int:
|
||||||
return -(-tokens // self.page_size) * self.page_size
|
return -(-tokens // self.page_size) * self.page_size
|
||||||
|
|
||||||
def budget_state(self):
|
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
|
return AddReqResult.NO_TOKEN
|
||||||
|
|
||||||
if self.rem_input_tokens <= 0:
|
if self.rem_input_tokens <= 0:
|
||||||
@@ -527,6 +553,9 @@ class PrefillAdder:
|
|||||||
self.cur_rem_token_offset += extend_input_len + page_overhead
|
self.cur_rem_token_offset += extend_input_len + page_overhead
|
||||||
self.rem_input_tokens -= extend_input_len
|
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:
|
if self.dllm_config is not None:
|
||||||
self.rem_dllm_tokens -= extend_input_len
|
self.rem_dllm_tokens -= extend_input_len
|
||||||
elif self.rem_chunk_tokens is not None:
|
elif self.rem_chunk_tokens is not None:
|
||||||
@@ -601,6 +630,8 @@ class PrefillAdder:
|
|||||||
_rem_tokens = self._get_dllm_remain_tokens()
|
_rem_tokens = self._get_dllm_remain_tokens()
|
||||||
else:
|
else:
|
||||||
_rem_tokens = min(self.rem_chunk_tokens, int(self.rem_total_tokens))
|
_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.
|
# 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.
|
# Therefore, in certain cases where _rem_tokens <= 0, it should be replaced with rem_chunk_tokens.
|
||||||
if _rem_tokens <= 0:
|
if _rem_tokens <= 0:
|
||||||
@@ -639,11 +670,12 @@ class PrefillAdder:
|
|||||||
self.tree_cache.dec_lock_ref(last_node)
|
self.tree_cache.dec_lock_ref(last_node)
|
||||||
|
|
||||||
def add_one_req_ignore_eos(self, req: Req):
|
def add_one_req_ignore_eos(self, req: Req):
|
||||||
# Early exit if no enough tokens for the input tokens
|
paged_input = self.ceil_paged_tokens(req.extend_input_len)
|
||||||
if self.ceil_paged_tokens(req.extend_input_len) > min(
|
if paged_input > min(self.cur_rem_tokens, self.rem_total_tokens):
|
||||||
self.cur_rem_tokens, self.rem_total_tokens
|
|
||||||
):
|
|
||||||
return AddReqResult.NO_TOKEN
|
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):
|
def add_req_state(r, insert_sort=False):
|
||||||
new_token_ratio = (
|
new_token_ratio = (
|
||||||
@@ -756,14 +788,11 @@ class PrefillAdder:
|
|||||||
# _update_prefill_budget already accounts for this in the deduction.
|
# _update_prefill_budget already accounts for this in the deduction.
|
||||||
# Without this, admission is more optimistic than the actual budget
|
# Without this, admission is more optimistic than the actual budget
|
||||||
# deduction, allowing over-admission when the pool is nearly full.
|
# deduction, allowing over-admission when the pool is nearly full.
|
||||||
total_tokens = (
|
max_new = min(
|
||||||
req.extend_input_len
|
max(req.sampling_params.max_new_tokens - len(req.output_ids), 0),
|
||||||
+ min(
|
CLIP_MAX_NEW_TOKENS,
|
||||||
max(req.sampling_params.max_new_tokens - len(req.output_ids), 0),
|
|
||||||
CLIP_MAX_NEW_TOKENS,
|
|
||||||
)
|
|
||||||
+ self.page_size
|
|
||||||
)
|
)
|
||||||
|
total_tokens = req.extend_input_len + max_new + self.page_size
|
||||||
|
|
||||||
# adjusting the input_tokens based on host_hit_length and 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
|
real_input_tokens = req.extend_input_len - req.host_hit_length
|
||||||
@@ -773,6 +802,11 @@ class PrefillAdder:
|
|||||||
if total_tokens >= self.rem_total_tokens:
|
if total_tokens >= self.rem_total_tokens:
|
||||||
return AddReqResult.NO_TOKEN
|
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:
|
if real_input_tokens >= self.rem_input_tokens and len(self.can_run_list) != 0:
|
||||||
return AddReqResult.OTHER
|
return AddReqResult.OTHER
|
||||||
|
|
||||||
@@ -781,6 +815,11 @@ class PrefillAdder:
|
|||||||
if total_tokens >= self.rem_total_tokens:
|
if total_tokens >= self.rem_total_tokens:
|
||||||
return AddReqResult.NO_TOKEN
|
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:
|
if req.host_hit_length > 0:
|
||||||
new_indices, req.last_node = self.tree_cache.init_load_back(
|
new_indices, req.last_node = self.tree_cache.init_load_back(
|
||||||
InitLoadBackParams(
|
InitLoadBackParams(
|
||||||
|
|||||||
@@ -1961,7 +1961,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
def max_token_pool_size(self):
|
def max_token_pool_size(self):
|
||||||
"""Return the max token pool size considering hybrid swa settings."""
|
"""Return the max token pool size considering hybrid swa settings."""
|
||||||
if self.is_hybrid_swa:
|
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:
|
else:
|
||||||
return self.max_total_num_tokens
|
return self.max_total_num_tokens
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user