perf(disaggregation): reuse req pool freelists and alloc_extend tensors

This commit is contained in:
wxiwnd
2026-04-08 20:12:47 +08:00
parent a371eacd86
commit 53a04a9a97
3 changed files with 264 additions and 41 deletions

View 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()