Fix hybrid swa chunked prefill oom (#23174)

This commit is contained in:
Ke Bao
2026-04-21 10:46:45 +08:00
committed by GitHub
parent ab3ce02de9
commit 50fc2c9e23
2 changed files with 64 additions and 1 deletions
@@ -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
@@ -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()