|
15 | 15 | import zmq |
16 | 16 | from vllm.utils.network_utils import make_zmq_path |
17 | 17 | from vllm.v1.kv_cache_interface import FullAttentionSpec, UniformTypeKVCacheSpecs |
| 18 | +from vllm.v1.request import RequestStatus |
18 | 19 |
|
19 | 20 | fake_engine = types.ModuleType("mooncake.engine") |
20 | 21 | fake_engine.TransferEngine = MagicMock() # type: ignore[attr-defined] |
@@ -1282,6 +1283,14 @@ def setUp(self): |
1282 | 1283 | ): |
1283 | 1284 | self.scheduler = MooncakeConnectorScheduler(self.config, "test_engine", MockKVCacheConfig()) |
1284 | 1285 |
|
| 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 | + |
1285 | 1294 | def test_get_num_new_matched_tokens_no_remote_prefill(self): |
1286 | 1295 | request = MockRequest("req1") |
1287 | 1296 | 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): |
1416 | 1425 |
|
1417 | 1426 | self.assertEqual(block_ids, ([10],)) |
1418 | 1427 |
|
| 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 | + |
1419 | 1528 |
|
1420 | 1529 | class TestUtils(unittest.TestCase): |
1421 | 1530 | def test_string_to_int64_hash(self): |
|
0 commit comments