From 50fc2c9e2308114b40d507ca8446e2c3bdd3842e Mon Sep 17 00:00:00 2001 From: Ke Bao Date: Tue, 21 Apr 2026 10:46:45 +0800 Subject: [PATCH] Fix hybrid swa chunked prefill oom (#23174) --- python/sglang/srt/managers/schedule_policy.py | 8 ++- .../unit/managers/test_prefill_adder.py | 57 +++++++++++++++++++ 2 files changed, 64 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index 407de8ee6..10f0fe47b 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -631,10 +631,16 @@ class PrefillAdder: 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)) + # alloc_extend needs extend_num_tokens + page_size per request, + # so reserve one page here to avoid OOM + _rem_tokens = min( + _rem_tokens, int(self.rem_swa_tokens) - self.page_size + ) # 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: + if self.is_hybrid_swa: + return req _rem_tokens = self.rem_chunk_tokens truncated = req.extend_input_len > _rem_tokens diff --git a/test/registered/unit/managers/test_prefill_adder.py b/test/registered/unit/managers/test_prefill_adder.py index 2a1111712..b0b611e8e 100644 --- a/test/registered/unit/managers/test_prefill_adder.py +++ b/test/registered/unit/managers/test_prefill_adder.py @@ -442,6 +442,63 @@ class TestPrefillAdder(CustomTestCase): self.assertEqual(adder2.rem_chunk_tokens, 0) # 3 - 3 = 0 self.assertEqual(result3, AddReqResult.OTHER) + def _build_hybrid_swa_chunked_req( + self, *, page_size, rem_swa, rem_chunk=2048, extend_input_len=500 + ): + self.mock_token_allocator.swa_available_size.return_value = rem_swa + self.mock_token_allocator.full_available_size.return_value = 100_000 + self.mock_token_allocator.available_size.return_value = 100_000 + self.mock_tree_cache.sliding_window_size = 128 + adder = self.create_adder( + self.create_running_batch(), + page_size=page_size, + rem_chunk_tokens=rem_chunk, + ) + adder.is_hybrid_swa = True + + req = self.create_mock_req("chunked", priority=0, max_new_tokens=128) + req.extend_input_len = extend_input_len + req.prefix_indices = [] + req.fill_ids = list(range(extend_input_len)) + req.set_extend_input_len = MagicMock() + return adder, req + + def test_add_chunked_req_hybrid_swa_reserves_page_for_alloc_extend(self): + # alloc_extend needs extend_num_tokens + page_size per request. If the + # scheduler hands out all of rem_swa_tokens, alloc_extend cannot get its + # extra page and OOMs. With the fix, extend_input_len must cap at + # rem_swa_tokens - page_size so the page is reserved. + PAGE_SIZE = 64 + REM_SWA = 100 + adder, req = self._build_hybrid_swa_chunked_req( + page_size=PAGE_SIZE, rem_swa=REM_SWA + ) + + result = adder.add_chunked_req(req) + + self.assertIs(result, req) # truncated → chunked prefill continues + req.set_extend_input_len.assert_called_once() + new_len = req.set_extend_input_len.call_args.args[0] + self.assertLessEqual(new_len + PAGE_SIZE, REM_SWA) + self.assertEqual(new_len, REM_SWA - PAGE_SIZE) + + def test_add_chunked_req_hybrid_swa_defers_when_swa_below_page(self): + # When rem_swa_tokens <= page_size there is no room to serve even the + # reservation, so the chunked req must be deferred (returned unchanged) + # instead of falling back to rem_chunk_tokens and bypassing SWA budget. + PAGE_SIZE = 64 + adder, req = self._build_hybrid_swa_chunked_req( + page_size=PAGE_SIZE, rem_swa=PAGE_SIZE + ) + original_len = req.extend_input_len + + result = adder.add_chunked_req(req) + + self.assertIs(result, req) + req.set_extend_input_len.assert_not_called() + self.assertEqual(req.extend_input_len, original_len) + self.assertEqual(len(adder.can_run_list), 0) + if __name__ == "__main__": unittest.main()