forked from Karylab-cklius/vllm
@@ -40,6 +40,26 @@ from .utils import EOS_TOKEN_ID, create_requests, create_scheduler, mock_kv
|
||||
pytestmark = pytest.mark.cpu_test
|
||||
|
||||
|
||||
def _get_remote_waiting_queue(scheduler: Scheduler):
|
||||
return getattr(scheduler, "waiting_for_remote_kvs", None)
|
||||
|
||||
|
||||
def _num_waiting_requests(scheduler: Scheduler) -> int:
|
||||
remote_waiting = _get_remote_waiting_queue(scheduler)
|
||||
return len(scheduler.waiting) + (len(remote_waiting) if remote_waiting else 0)
|
||||
|
||||
|
||||
def _get_remote_waiting_requests(scheduler: Scheduler) -> list[Request]:
|
||||
remote_waiting = _get_remote_waiting_queue(scheduler)
|
||||
if remote_waiting is not None:
|
||||
return list(remote_waiting)
|
||||
return [
|
||||
req
|
||||
for req in scheduler.waiting
|
||||
if req.status == RequestStatus.WAITING_FOR_REMOTE_KVS
|
||||
]
|
||||
|
||||
|
||||
def test_add_requests():
|
||||
scheduler = create_scheduler()
|
||||
requests = create_requests(num_requests=10)
|
||||
@@ -1120,7 +1140,8 @@ def _step_until_kv_transfer_finished(scheduler: Scheduler, req_ids: list[str]):
|
||||
|
||||
# Requests should first transition to WAITING_FOR_REMOTE_KVS
|
||||
output = scheduler.schedule()
|
||||
assert len(scheduler.waiting) == len(req_ids)
|
||||
assert _num_waiting_requests(scheduler) == len(req_ids)
|
||||
assert len(_get_remote_waiting_requests(scheduler)) == len(req_ids)
|
||||
assert len(scheduler.running) == 0
|
||||
assert len(output.scheduled_new_reqs) == 0
|
||||
for req in scheduler.requests.values():
|
||||
@@ -1139,7 +1160,8 @@ def _step_until_kv_transfer_finished(scheduler: Scheduler, req_ids: list[str]):
|
||||
|
||||
# Simulate KV transfer completion using KVConnectorOutput.finished_recving
|
||||
output = scheduler.schedule()
|
||||
assert len(scheduler.waiting) == len(req_ids)
|
||||
assert _num_waiting_requests(scheduler) == len(req_ids)
|
||||
assert len(_get_remote_waiting_requests(scheduler)) == len(req_ids)
|
||||
assert len(scheduler.running) == 0
|
||||
|
||||
MODEL_RUNNER_OUTPUT = ModelRunnerOutput(
|
||||
@@ -1546,7 +1568,7 @@ def test_kv_connector_handles_preemption(is_async, use_ec_connector, ec_role):
|
||||
# All can be scheduled - 1st token.
|
||||
output = scheduler.schedule()
|
||||
if is_async:
|
||||
assert len(scheduler.waiting) == 2
|
||||
assert _num_waiting_requests(scheduler) == 2
|
||||
assert scheduler.running == []
|
||||
_step_until_kv_transfer_finished(scheduler, req_ids)
|
||||
output = scheduler.schedule()
|
||||
@@ -1604,7 +1626,9 @@ def test_kv_connector_handles_preemption(is_async, use_ec_connector, ec_role):
|
||||
# This will have a local and remote cache hit.
|
||||
output = scheduler.schedule()
|
||||
if is_async:
|
||||
waiting_req_ids = [req.request_id for req in scheduler.waiting]
|
||||
waiting_req_ids = [
|
||||
req.request_id for req in _get_remote_waiting_requests(scheduler)
|
||||
]
|
||||
assert len(waiting_req_ids) == 1
|
||||
_step_until_kv_transfer_finished(scheduler, waiting_req_ids)
|
||||
output = scheduler.schedule()
|
||||
@@ -3614,6 +3638,10 @@ def test_prepend_skipped_requests_order():
|
||||
# simulate first 2 waiting requests are waiting for remote KVs
|
||||
for req in expected_waiting_reqs[:2]:
|
||||
req.status = RequestStatus.WAITING_FOR_REMOTE_KVS
|
||||
remote_waiting = _get_remote_waiting_queue(scheduler)
|
||||
if remote_waiting is not None:
|
||||
scheduler.waiting.remove_request(req)
|
||||
remote_waiting.add_request(req)
|
||||
|
||||
# schedule step
|
||||
# expect the first 2 waiting to be skipped, the third running,
|
||||
@@ -3623,8 +3651,13 @@ def test_prepend_skipped_requests_order():
|
||||
# pop the third request which is expected to be running
|
||||
expected_waiting_reqs.pop(2)
|
||||
|
||||
# verify waiting order is preserved
|
||||
assert list(scheduler.waiting) == expected_waiting_reqs
|
||||
# verify waiting order is preserved for schedulable requests.
|
||||
remote_waiting = _get_remote_waiting_queue(scheduler)
|
||||
if remote_waiting is not None:
|
||||
assert list(scheduler.waiting) == expected_waiting_reqs[2:]
|
||||
assert list(remote_waiting) == expected_waiting_reqs[:2]
|
||||
else:
|
||||
assert list(scheduler.waiting) == expected_waiting_reqs
|
||||
|
||||
|
||||
def test_abort_request_waiting_for_remote_kvs():
|
||||
|
||||
@@ -17,6 +17,26 @@ from .utils import (
|
||||
)
|
||||
|
||||
|
||||
def _get_remote_waiting_queue(scheduler: Scheduler):
|
||||
return getattr(scheduler, "waiting_for_remote_kvs", None)
|
||||
|
||||
|
||||
def _num_waiting_requests(scheduler: Scheduler) -> int:
|
||||
remote_waiting = _get_remote_waiting_queue(scheduler)
|
||||
return len(scheduler.waiting) + (len(remote_waiting) if remote_waiting else 0)
|
||||
|
||||
|
||||
def _get_remote_waiting_requests(scheduler: Scheduler) -> list[Request]:
|
||||
remote_waiting = _get_remote_waiting_queue(scheduler)
|
||||
if remote_waiting is not None:
|
||||
return list(remote_waiting)
|
||||
return [
|
||||
req
|
||||
for req in scheduler.waiting
|
||||
if req.status == RequestStatus.WAITING_FOR_REMOTE_KVS
|
||||
]
|
||||
|
||||
|
||||
def _make_get_num_new_matched_tokens(
|
||||
req_num_new_matched_tokens: dict[str, int],
|
||||
async_load,
|
||||
@@ -76,8 +96,8 @@ def test_async_load_failure(
|
||||
|
||||
scheduler_output = scheduler.schedule()
|
||||
|
||||
assert len(scheduler.waiting) == 3
|
||||
for request in scheduler.waiting:
|
||||
assert _num_waiting_requests(scheduler) == 3
|
||||
for request in _get_remote_waiting_requests(scheduler):
|
||||
assert request.num_computed_tokens == 0
|
||||
assert request.status == RequestStatus.WAITING_FOR_REMOTE_KVS
|
||||
assert scheduler.connector.get_num_new_matched_tokens.call_count == 3
|
||||
@@ -96,8 +116,8 @@ def test_async_load_failure(
|
||||
|
||||
min_invalid_block_idx = min(invalid_block_idxs)
|
||||
|
||||
assert len(scheduler.waiting) == 3
|
||||
for request in scheduler.waiting:
|
||||
assert _num_waiting_requests(scheduler) == 3
|
||||
for request in _get_remote_waiting_requests(scheduler):
|
||||
if request.request_id == request2.request_id:
|
||||
assert request.num_computed_tokens == (
|
||||
min_invalid_block_idx * scheduler.block_size
|
||||
@@ -303,8 +323,10 @@ def test_async_progressive_load_failure(
|
||||
|
||||
scheduler_output = scheduler.schedule()
|
||||
|
||||
assert len(scheduler.waiting) == 1
|
||||
assert scheduler.waiting.peek_request().request_id == request.request_id
|
||||
assert _num_waiting_requests(scheduler) == 1
|
||||
remote_waiting_reqs = _get_remote_waiting_requests(scheduler)
|
||||
assert len(remote_waiting_reqs) == 1
|
||||
assert remote_waiting_reqs[0].request_id == request.request_id
|
||||
assert request.num_computed_tokens == 0
|
||||
assert request.status == RequestStatus.WAITING_FOR_REMOTE_KVS
|
||||
assert scheduler.connector.get_num_new_matched_tokens.call_count == 1
|
||||
@@ -325,8 +347,10 @@ def test_async_progressive_load_failure(
|
||||
|
||||
min_invalid_block_idx = min(min_invalid_block_idx, invalid_block_idx)
|
||||
|
||||
assert len(scheduler.waiting) == 1
|
||||
assert scheduler.waiting.peek_request().request_id == request.request_id
|
||||
assert _num_waiting_requests(scheduler) == 1
|
||||
remote_waiting_reqs = _get_remote_waiting_requests(scheduler)
|
||||
assert len(remote_waiting_reqs) == 1
|
||||
assert remote_waiting_reqs[0].request_id == request.request_id
|
||||
assert request.num_computed_tokens == (
|
||||
min_invalid_block_idx * scheduler.block_size
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user