Skip to content

Commit e928441

Browse files
authored
[Test][P/D] Add request_finished extra UT (vllm-project#11026)
### What this PR does / why we need it? Add E2E UT for `request_finished` function, which includes E2E slidingwindow & MTP blocks, ensuring that `MooncakeConnector` can provide correct `block_ids` to transfer. ### Does this PR introduce _any_ user-facing change? No. ### How was this patch tested? Had been tested on local. - vLLM version: v0.23.0 - vLLM main: vllm-project/vllm@dc68bd8 Signed-off-by: nwpu-zxr <zhouxuerong2@huawei.com>
1 parent 10b244a commit e928441

1 file changed

Lines changed: 109 additions & 0 deletions

File tree

tests/ut/kv_offload/test_mooncake_connector.py

Lines changed: 109 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515
import zmq
1616
from vllm.utils.network_utils import make_zmq_path
1717
from vllm.v1.kv_cache_interface import FullAttentionSpec, UniformTypeKVCacheSpecs
18+
from vllm.v1.request import RequestStatus
1819

1920
fake_engine = types.ModuleType("mooncake.engine")
2021
fake_engine.TransferEngine = MagicMock() # type: ignore[attr-defined]
@@ -1282,6 +1283,14 @@ def setUp(self):
12821283
):
12831284
self.scheduler = MooncakeConnectorScheduler(self.config, "test_engine", MockKVCacheConfig())
12841285

1286+
def _make_remote_decode_request(self, prompt_len: int, request_id: str = "req1"):
1287+
return MockRequest(
1288+
request_id,
1289+
prompt_token_ids=list(range(prompt_len)),
1290+
kv_transfer_params={"do_remote_decode": True},
1291+
status=RequestStatus.FINISHED_LENGTH_CAPPED,
1292+
)
1293+
12851294
def test_get_num_new_matched_tokens_no_remote_prefill(self):
12861295
request = MockRequest("req1")
12871296
tokens, async_flag = self.scheduler.get_num_new_matched_tokens(request, 0)
@@ -1416,6 +1425,106 @@ def test_transfer_block_ids_trims_mtp_before_swa_zero_filter(self):
14161425

14171426
self.assertEqual(block_ids, ([10],))
14181427

1428+
def test_request_finished_trims_mtp_blocks_in_params(self):
1429+
self.scheduler.group_transfer_info = [
1430+
types.SimpleNamespace(
1431+
tokens_per_block=16,
1432+
blocks_per_window=0,
1433+
is_state_group=False,
1434+
)
1435+
]
1436+
request = self._make_remote_decode_request(prompt_len=33, request_id="req_mtp")
1437+
1438+
delay_free, params = self.scheduler.request_finished(request, ([10, 11, 12, 13, 14],))
1439+
1440+
self.assertTrue(delay_free)
1441+
self.assertIsNotNone(params)
1442+
assert params is not None
1443+
self.assertEqual(params["remote_block_ids"], ([10, 11, 12],))
1444+
self.assertEqual(params["num_prompt_blocks"], 3)
1445+
self.assertIn("req_mtp", self.scheduler._reqs_need_send)
1446+
1447+
def test_request_finished_clips_sliding_window_blocks_in_params(self):
1448+
self.scheduler.group_transfer_info = [
1449+
types.SimpleNamespace(
1450+
tokens_per_block=16,
1451+
blocks_per_window=3,
1452+
is_state_group=False,
1453+
)
1454+
]
1455+
request = self._make_remote_decode_request(prompt_len=80, request_id="req_swa")
1456+
1457+
delay_free, params = self.scheduler.request_finished(request, ([10, 11, 12, 13, 14],))
1458+
1459+
self.assertTrue(delay_free)
1460+
self.assertIsNotNone(params)
1461+
assert params is not None
1462+
self.assertEqual(params["remote_block_ids"], ([12, 13, 14],))
1463+
self.assertEqual(params["num_prompt_blocks"], 5)
1464+
self.assertIn("req_swa", self.scheduler._reqs_need_send)
1465+
1466+
def test_request_finished_trims_mtp_before_swa_tail_clip(self):
1467+
self.scheduler.group_transfer_info = [
1468+
types.SimpleNamespace(
1469+
tokens_per_block=16,
1470+
blocks_per_window=3,
1471+
is_state_group=False,
1472+
)
1473+
]
1474+
request = self._make_remote_decode_request(prompt_len=64, request_id="req_mtp_swa")
1475+
1476+
delay_free, params = self.scheduler.request_finished(request, ([0, 10, 11, 12, 13, 14],))
1477+
1478+
self.assertTrue(delay_free)
1479+
self.assertIsNotNone(params)
1480+
assert params is not None
1481+
self.assertEqual(params["remote_block_ids"], ([10, 11, 12],))
1482+
self.assertEqual(params["num_prompt_blocks"], 4)
1483+
self.assertIn("req_mtp_swa", self.scheduler._reqs_need_send)
1484+
1485+
def test_request_finished_handles_mtp_swa_and_state_groups_together(self):
1486+
self.scheduler.group_transfer_info = [
1487+
types.SimpleNamespace(
1488+
tokens_per_block=16,
1489+
blocks_per_window=0,
1490+
is_state_group=False,
1491+
),
1492+
types.SimpleNamespace(
1493+
tokens_per_block=16,
1494+
blocks_per_window=3,
1495+
is_state_group=False,
1496+
),
1497+
types.SimpleNamespace(
1498+
tokens_per_block=16,
1499+
blocks_per_window=0,
1500+
is_state_group=True,
1501+
),
1502+
]
1503+
request = self._make_remote_decode_request(prompt_len=64, request_id="req_mixed_groups")
1504+
1505+
delay_free, params = self.scheduler.request_finished(
1506+
request,
1507+
(
1508+
[100, 101, 102, 103, 104],
1509+
[0, 200, 201, 202, 203, 204],
1510+
[300, 301, 302, 303, 304],
1511+
),
1512+
)
1513+
1514+
self.assertTrue(delay_free)
1515+
self.assertIsNotNone(params)
1516+
assert params is not None
1517+
self.assertEqual(
1518+
params["remote_block_ids"],
1519+
(
1520+
[100, 101, 102, 103],
1521+
[200, 201, 202],
1522+
[300, 301, 302, 303, 304],
1523+
),
1524+
)
1525+
self.assertEqual(params["num_prompt_blocks"], 4)
1526+
self.assertIn("req_mixed_groups", self.scheduler._reqs_need_send)
1527+
14191528

14201529
class TestUtils(unittest.TestCase):
14211530
def test_string_to_int64_hash(self):

0 commit comments

Comments
 (0)