Fix hybrid swa chunked prefill oom (#23174)
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user