diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index 6909d9483..407de8ee6 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -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( diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index efebc02e8..7dd70af63 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -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