perf(disaggregation): reuse req pool freelists and alloc_extend tensors
This commit is contained in:
167
test/registered/unit/mem_cache/test_req_to_token_pool.py
Normal file
167
test/registered/unit/mem_cache/test_req_to_token_pool.py
Normal file
@@ -0,0 +1,167 @@
|
||||
"""Unit tests for req-pool and decode prealloc optimizations."""
|
||||
|
||||
import unittest
|
||||
from collections import deque
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.disaggregation.decode import DecodePreallocQueue, DecodeReqToTokenPool
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=6, suite="stage-a-test-cpu")
|
||||
|
||||
|
||||
class TestReqToTokenPool(CustomTestCase):
|
||||
def _make_req(self, *, req_pool_idx=None, is_chunked=0, kv_committed_len=0):
|
||||
req = MagicMock(spec=Req)
|
||||
req.req_pool_idx = req_pool_idx
|
||||
req.is_chunked = is_chunked
|
||||
req.kv_committed_len = kv_committed_len
|
||||
return req
|
||||
|
||||
def test_req_to_token_pool_uses_deque_backed_free_slots(self):
|
||||
pool = ReqToTokenPool(
|
||||
size=4,
|
||||
max_context_len=8,
|
||||
device="cpu",
|
||||
enable_memory_saver=False,
|
||||
)
|
||||
|
||||
self.assertIsInstance(pool.free_slots, deque)
|
||||
|
||||
req = self._make_req()
|
||||
self.assertEqual(pool.alloc([req]), [0])
|
||||
pool.free(req)
|
||||
pool.clear()
|
||||
|
||||
self.assertIsInstance(pool.free_slots, deque)
|
||||
self.assertEqual(list(pool.free_slots), [0, 1, 2, 3])
|
||||
|
||||
def test_decode_req_to_token_pool_preserves_reuse_and_uses_deque(self):
|
||||
pool = DecodeReqToTokenPool(
|
||||
size=2,
|
||||
max_context_len=8,
|
||||
device="cpu",
|
||||
enable_memory_saver=False,
|
||||
pre_alloc_size=2,
|
||||
)
|
||||
|
||||
self.assertIsInstance(pool.free_slots, deque)
|
||||
self.assertEqual(pool.available_size(), 4)
|
||||
|
||||
reused_req = self._make_req(is_chunked=1)
|
||||
self.assertEqual(pool.alloc([reused_req]), [0])
|
||||
|
||||
fresh_req = self._make_req()
|
||||
self.assertEqual(pool.alloc([fresh_req, reused_req]), [1, 0])
|
||||
|
||||
pool.clear()
|
||||
|
||||
self.assertIsInstance(pool.free_slots, deque)
|
||||
self.assertEqual(list(pool.free_slots), [0, 1, 2, 3])
|
||||
|
||||
|
||||
class TestDecodePreallocQueue(CustomTestCase):
|
||||
def _make_req(self, origin_input_ids, output_ids):
|
||||
req = MagicMock(spec=Req)
|
||||
req.req_pool_idx = None
|
||||
req.origin_input_ids = origin_input_ids
|
||||
req.output_ids = output_ids
|
||||
req.kv_allocated_len = 0
|
||||
req.kv_committed_len = 0
|
||||
req.is_chunked = 0
|
||||
req.set_extend_input_len = MagicMock()
|
||||
return req
|
||||
|
||||
def _build_prealloc_queue(self):
|
||||
req_to_token_pool = DecodeReqToTokenPool(
|
||||
size=4,
|
||||
max_context_len=16,
|
||||
device="cpu",
|
||||
enable_memory_saver=False,
|
||||
pre_alloc_size=2,
|
||||
)
|
||||
|
||||
allocator_calls = []
|
||||
|
||||
def alloc_extend(**kwargs):
|
||||
allocator_calls.append(kwargs)
|
||||
return torch.arange(kwargs["extend_num_tokens"], dtype=torch.int64)
|
||||
|
||||
token_to_kv_pool_allocator = MagicMock()
|
||||
token_to_kv_pool_allocator.page_size = 16
|
||||
token_to_kv_pool_allocator.device = "cpu"
|
||||
token_to_kv_pool_allocator.alloc_extend.side_effect = alloc_extend
|
||||
token_to_kv_pool_allocator.get_kvcache.return_value = MagicMock()
|
||||
|
||||
scheduler = SimpleNamespace(tp_worker=SimpleNamespace(is_hybrid_swa=False))
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
DecodePreallocQueue, "_init_kv_manager", return_value=MagicMock()
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.disaggregation.decode.is_mla_backend", return_value=False
|
||||
),
|
||||
):
|
||||
queue = DecodePreallocQueue(
|
||||
req_to_token_pool=req_to_token_pool,
|
||||
token_to_kv_pool_allocator=token_to_kv_pool_allocator,
|
||||
draft_token_to_kv_pool=None,
|
||||
req_to_metadata_buffer_idx_allocator=MagicMock(),
|
||||
metadata_buffers=MagicMock(),
|
||||
scheduler=scheduler,
|
||||
transfer_queue=SimpleNamespace(queue=[]),
|
||||
tree_cache=MagicMock(),
|
||||
gloo_group=MagicMock(),
|
||||
tp_rank=0,
|
||||
tp_size=1,
|
||||
dp_size=1,
|
||||
gpu_id=0,
|
||||
bootstrap_port=1234,
|
||||
max_total_num_tokens=256,
|
||||
pp_rank=0,
|
||||
num_reserved_decode_tokens=1,
|
||||
transfer_backend=MagicMock(),
|
||||
)
|
||||
|
||||
return queue, allocator_calls
|
||||
|
||||
def test_pre_alloc_reuses_alloc_extend_tensors(self):
|
||||
queue, allocator_calls = self._build_prealloc_queue()
|
||||
|
||||
req1 = self._make_req(origin_input_ids=[1, 2, 3], output_ids=[4, 5])
|
||||
req2 = self._make_req(origin_input_ids=[1, 2, 3, 4], output_ids=[5, 6])
|
||||
|
||||
queue._pre_alloc(req1)
|
||||
queue._pre_alloc(req2)
|
||||
|
||||
self.assertEqual(len(allocator_calls), 2)
|
||||
self.assertIs(
|
||||
allocator_calls[0]["prefix_lens"], allocator_calls[1]["prefix_lens"]
|
||||
)
|
||||
self.assertIs(
|
||||
allocator_calls[0]["prefix_lens_cpu"],
|
||||
allocator_calls[1]["prefix_lens_cpu"],
|
||||
)
|
||||
self.assertIs(allocator_calls[0]["seq_lens"], allocator_calls[1]["seq_lens"])
|
||||
self.assertIs(
|
||||
allocator_calls[0]["seq_lens_cpu"], allocator_calls[1]["seq_lens_cpu"]
|
||||
)
|
||||
self.assertIs(allocator_calls[0]["last_loc"], allocator_calls[1]["last_loc"])
|
||||
|
||||
self.assertEqual(allocator_calls[1]["prefix_lens"].item(), 0)
|
||||
self.assertEqual(allocator_calls[1]["prefix_lens_cpu"].item(), 0)
|
||||
self.assertEqual(allocator_calls[1]["seq_lens"].item(), 5)
|
||||
self.assertEqual(allocator_calls[1]["seq_lens_cpu"].item(), 5)
|
||||
self.assertEqual(allocator_calls[1]["last_loc"].item(), -1)
|
||||
self.assertEqual(allocator_calls[1]["extend_num_tokens"], 5)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user