Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
67 changes: 66 additions & 1 deletion tests/ut/kv_offload/test_mooncake_connector.py
Original file line number Diff line number Diff line change
Expand Up @@ -2059,7 +2059,8 @@ def get_block_ids(self):


class MockSchedulerOutput:
pass
def __init__(self):
self.num_scheduled_tokens: dict[str, int] = {}


class MockForwardContext:
Expand Down Expand Up @@ -2198,12 +2199,76 @@ def test_update_state_after_alloc_with_remote_prefill(self):
self.assertEqual(self.scheduler._reqs_need_recv["req1"][1], ([4, 5, 6],))
self.assertEqual(self.scheduler._reqs_need_recv["req1"][2], ([1, 2, 4, 5, 6],))

def test_pd_decode_recompute_marker_waits_for_first_forward_step(self):
request = MockRequest(
"req1",
kv_transfer_params={
"do_remote_prefill": True,
"remote_block_ids": [1, 2, 3],
"remote_engine_id": "remote",
"remote_request_id": "remote_req1",
"remote_host": "localhost",
"remote_port": 5000,
},
)
self.scheduler.update_state_after_alloc(request, MockKVCacheBlocks(), 3)

receive_step = MockSchedulerOutput()
receive_step.num_scheduled_tokens = {"req1": 1}
receive_meta = self.scheduler.build_connector_meta(receive_step)
self.assertEqual(receive_meta.pd_decode_recompute_req_ids, set())
self.assertEqual(self.scheduler._pd_decode_recompute_req_ids, {"req1"})

self.scheduler.update_state_after_alloc(request, MockKVCacheBlocks(), 0)
forward_step = MockSchedulerOutput()
forward_step.num_scheduled_tokens = {"req1": 1, "decode": 1}
forward_meta = self.scheduler.build_connector_meta(forward_step)
self.assertEqual(forward_meta.pd_decode_recompute_req_ids, {"req1"})
self.assertEqual(self.scheduler._pd_decode_recompute_req_ids, set())

later_step = MockSchedulerOutput()
later_step.num_scheduled_tokens = {"req1": 1}
later_meta = self.scheduler.build_connector_meta(later_step)
self.assertEqual(later_meta.pd_decode_recompute_req_ids, set())

def test_request_finished_no_remote_decode(self):
request = MockRequest("req1")
delay_free, params = self.scheduler.request_finished(request, [1, 2, 3])
self.assertFalse(delay_free)
self.assertIsNone(params)

def test_request_finished_clears_pd_decode_recompute_marker(self):
for phase in ("pending", "ready"):
with self.subTest(phase=phase):
request = MockRequest("req1")
self.scheduler._pd_decode_recompute_req_ids = {"req1", "other"}
self.scheduler._pd_decode_recompute_ready_req_ids = {"req1", "other"} if phase == "ready" else {"other"}

delay_free, params = self.scheduler.request_finished(request, [1, 2, 3])

self.assertFalse(delay_free)
self.assertIsNone(params)
self.assertEqual(self.scheduler._pd_decode_recompute_req_ids, {"other"})
self.assertEqual(
self.scheduler._pd_decode_recompute_ready_req_ids,
{"other"},
)

metadata = self.scheduler.build_connector_meta(MockSchedulerOutput())
self.assertEqual(metadata.pd_decode_recompute_req_ids, {"other"})

def test_finished_marker_does_not_contaminate_reused_request_id(self):
finished = MockRequest("req1")
self.scheduler._pd_decode_recompute_req_ids.add("req1")
self.scheduler._pd_decode_recompute_ready_req_ids.add("req1")
self.scheduler.request_finished(finished, [1, 2, 3])

reused = MockRequest("req1")
self.scheduler.update_state_after_alloc(reused, MockKVCacheBlocks(), 0)
metadata = self.scheduler.build_connector_meta(MockSchedulerOutput())

self.assertEqual(metadata.pd_decode_recompute_req_ids, set())

def test_request_finished_rejected_remote_prefill_enqueues_empty_recv(self):
request = MockRequest(
"req1",
Expand Down
178 changes: 177 additions & 1 deletion tests/ut/worker/test_model_runner_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from vllm.config import CUDAGraphMode
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm.v1.worker.gpu import model_runner as vllm_model_runner
from vllm.v1.worker.gpu.model_runner import GPUModelRunner
from vllm.v1.worker.gpu.model_runner import BatchReqState, GPUModelRunner

from vllm_ascend.ascend_forward_context import MoECommType
from vllm_ascend.utils import vllm_version_is
Expand All @@ -36,6 +36,182 @@ def _make_runner(need_timing: bool = True):
return runner


def test_pd_decode_tail_recompute_keeps_uniform_decode_batch():
runner = _make_runner()
runner.decode_query_len = 1
runner.vllm_config = SimpleNamespace(kv_transfer_config=SimpleNamespace(is_kv_consumer=True))
batch_state = BatchReqState(
req_ids=["migrating", "decoding"],
num_scheduled_tokens=np.array([1, 1], dtype=np.int32),
num_tokens=2,
idx_mapping_np=np.array([0, 1], dtype=np.intp),
prefill_len_np=np.array([128, 64], dtype=np.int32),
num_computed_prefill_tokens_np=np.array([127, 64], dtype=np.int32),
is_prefilling_np=np.array([True, False]),
has_prefill=True,
)

transfer_step = SimpleNamespace(
kv_connector_metadata=SimpleNamespace(pd_decode_recompute_req_ids={"migrating"}),
finished_req_ids=set(),
preempted_req_ids=None,
)
runner._update_pd_decode_recompute_requests(transfer_step)

with patch.object(
GPUModelRunner,
"gather_batch_req_state",
return_value=(batch_state, None),
):
scheduler_output = SimpleNamespace(kv_connector_metadata=SimpleNamespace(pd_decode_recompute_req_ids=set()))
gathered, uniform = runner.gather_batch_req_state(scheduler_output, False)

assert gathered is not batch_state
np.testing.assert_array_equal(gathered.is_prefilling_np, [False, False])
assert not gathered.has_prefill
assert uniform == 1
assert not runner._pd_decode_recompute_req_ids


@pytest.mark.parametrize("computed,scheduled", [(64, 8), (127, 1)])
def test_pd_decode_real_prefill_still_blocks_decode_graph(computed, scheduled):
runner = _make_runner()
runner.decode_query_len = 1
runner.vllm_config = SimpleNamespace(kv_transfer_config=SimpleNamespace(is_kv_consumer=True))
batch_state = BatchReqState(
req_ids=["fallback-prefill", "decoding"],
num_scheduled_tokens=np.array([scheduled, 1], dtype=np.int32),
num_tokens=scheduled + 1,
idx_mapping_np=np.array([0, 1], dtype=np.intp),
prefill_len_np=np.array([128, 64], dtype=np.int32),
num_computed_prefill_tokens_np=np.array([computed, 64], dtype=np.int32),
is_prefilling_np=np.array([True, False]),
has_prefill=True,
)

with patch.object(
GPUModelRunner,
"gather_batch_req_state",
return_value=(batch_state, None),
):
scheduler_output = SimpleNamespace(kv_connector_metadata=SimpleNamespace(pd_decode_recompute_req_ids=set()))
gathered, uniform = runner.gather_batch_req_state(scheduler_output, False)

np.testing.assert_array_equal(gathered.is_prefilling_np, [True, False])
assert gathered.has_prefill
assert uniform is None


def test_pd_decode_only_reclassifies_transferred_requests_in_mixed_prefill_batch():
runner = _make_runner()
runner.decode_query_len = 1
runner.vllm_config = SimpleNamespace(kv_transfer_config=SimpleNamespace(is_kv_consumer=True))
batch_state = BatchReqState(
req_ids=["migrating", "local-final-prefill", "decoding"],
num_scheduled_tokens=np.array([1, 1, 1], dtype=np.int32),
num_tokens=3,
idx_mapping_np=np.array([0, 1, 2], dtype=np.intp),
prefill_len_np=np.array([128, 128, 64], dtype=np.int32),
num_computed_prefill_tokens_np=np.array([127, 127, 64], dtype=np.int32),
is_prefilling_np=np.array([True, True, False]),
has_prefill=True,
)

with patch.object(
GPUModelRunner,
"gather_batch_req_state",
return_value=(batch_state, None),
):
scheduler_output = SimpleNamespace(
kv_connector_metadata=SimpleNamespace(pd_decode_recompute_req_ids={"migrating"})
)
gathered, uniform = runner.gather_batch_req_state(scheduler_output, False)

np.testing.assert_array_equal(gathered.is_prefilling_np, [False, True, False])
assert gathered.has_prefill
assert uniform is None


def test_pd_decode_transfer_tracking_clears_finished_and_preempted_requests():
runner = _make_runner()
runner._pd_decode_recompute_req_ids = {"finished", "preempted", "active"}
scheduler_output = SimpleNamespace(
kv_connector_metadata=SimpleNamespace(
metadata=(
SimpleNamespace(pd_decode_recompute_req_ids={"new"}),
SimpleNamespace(pd_decode_recompute_req_ids=set()),
)
),
finished_req_ids={"finished"},
preempted_req_ids={"preempted"},
)

runner._update_pd_decode_recompute_requests(scheduler_output)

assert runner._pd_decode_recompute_req_ids == {"active", "new"}


def test_pd_decode_tail_recompute_supports_multi_token_decode_query():
runner = _make_runner()
runner.decode_query_len = 2
runner.vllm_config = SimpleNamespace(kv_transfer_config=SimpleNamespace(is_kv_consumer=True))
batch_state = BatchReqState(
req_ids=["migrating", "decoding"],
num_scheduled_tokens=np.array([2, 2], dtype=np.int32),
num_tokens=4,
idx_mapping_np=np.array([0, 1], dtype=np.intp),
prefill_len_np=np.array([128, 64], dtype=np.int32),
num_computed_prefill_tokens_np=np.array([127, 64], dtype=np.int32),
is_prefilling_np=np.array([True, False]),
has_prefill=True,
)

with patch.object(
GPUModelRunner,
"gather_batch_req_state",
return_value=(batch_state, None),
):
scheduler_output = SimpleNamespace(
kv_connector_metadata=SimpleNamespace(pd_decode_recompute_req_ids={"migrating"})
)
gathered, uniform = runner.gather_batch_req_state(scheduler_output, False)

np.testing.assert_array_equal(gathered.is_prefilling_np, [False, False])
assert not gathered.has_prefill
assert uniform == 2


def test_pd_decode_non_consumer_does_not_reclassify_transferred_request():
runner = _make_runner()
runner.decode_query_len = 1
runner.vllm_config = SimpleNamespace(kv_transfer_config=SimpleNamespace(is_kv_consumer=False))
batch_state = BatchReqState(
req_ids=["transferred"],
num_scheduled_tokens=np.array([1], dtype=np.int32),
num_tokens=1,
idx_mapping_np=np.array([0], dtype=np.intp),
prefill_len_np=np.array([128], dtype=np.int32),
num_computed_prefill_tokens_np=np.array([127], dtype=np.int32),
is_prefilling_np=np.array([True]),
has_prefill=True,
)

with patch.object(
GPUModelRunner,
"gather_batch_req_state",
return_value=(batch_state, None),
):
scheduler_output = SimpleNamespace(
kv_connector_metadata=SimpleNamespace(pd_decode_recompute_req_ids={"transferred"})
)
gathered, uniform = runner.gather_batch_req_state(scheduler_output, False)

assert gathered is batch_state
np.testing.assert_array_equal(gathered.is_prefilling_np, [True])
assert gathered.has_prefill
assert uniform is None


def test_execute_model_records_profiling_time():
runner = _make_runner()
scheduler_output = SimpleNamespace(disable_profiling_timing=False)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1576,6 +1576,9 @@ def __init__(self):
self.requests: dict[str, ReqMeta] = {}
self.requests_to_send: dict[str, float] = {}
self.reqs_in_batch: set[str] = set()
# Requests whose first local forward after a remote prefill is the
# decode-side tail-token recomputation.
self.pd_decode_recompute_req_ids: set[str] = set()

def add_new_req(
self,
Expand Down Expand Up @@ -1768,6 +1771,8 @@ def __init__(self, vllm_config: VllmConfig, engine_id: str, kv_cache_config: KVC
self._reqs_need_recv: dict[str, tuple[Request, BlockIds, BlockIds, int]] = {}
self._reqs_need_send: dict[str, float] = {}
self._reqs_in_batch: set[str] = set()
self._pd_decode_recompute_req_ids: set[str] = set()
self._pd_decode_recompute_ready_req_ids: set[str] = set()

# master-slave meta information for cross-nodes
self.multi_nodes_meta_mapping: dict[str, dict[str, Any]] = {}
Expand Down Expand Up @@ -1948,6 +1953,12 @@ def update_state_after_alloc(self, request: "Request", blocks: "KVCacheBlocks",

if params is not None and (params.get("do_remote_prefill", False) or params.get("do_remote_decode", False)):
self._reqs_in_batch.add(request.request_id)
if (
request.request_id in self._pd_decode_recompute_req_ids
and params is not None
and not params.get("do_remote_prefill", False)
):
self._pd_decode_recompute_ready_req_ids.add(request.request_id)
if params is not None and params.get("do_remote_prefill"):
if params.get("remote_block_ids"):
if all(p in params for p in ("remote_engine_id", "remote_host", "remote_port", "remote_request_id")):
Expand All @@ -1960,6 +1971,7 @@ def update_state_after_alloc(self, request: "Request", blocks: "KVCacheBlocks",
local_full_block_ids,
num_external_tokens,
)
self._pd_decode_recompute_req_ids.add(request.request_id)
else:
logger.warning("Got invalid KVTransferParams. params=%s. ", params)
else:
Expand Down Expand Up @@ -1994,6 +2006,18 @@ def build_connector_meta(
meta.reqs_in_batch = self._reqs_in_batch
self._reqs_in_batch = set()

# KV receive is scheduled in a no-forward step. update_state_after_alloc
# marks the same request ready when the scheduler later allocates its
# first local forward, which is the decode-side tail recomputation.
meta.pd_decode_recompute_req_ids = self._pd_decode_recompute_ready_req_ids
if meta.pd_decode_recompute_req_ids:
logger.debug(
"Emitting PD decode recompute requests: %s",
sorted(meta.pd_decode_recompute_req_ids),
)
self._pd_decode_recompute_req_ids.difference_update(meta.pd_decode_recompute_req_ids)
self._pd_decode_recompute_ready_req_ids = set()

return meta

def request_finished(
Expand All @@ -2006,6 +2030,13 @@ def request_finished(
should be freed now or will be sent asynchronously and freed later.
"""

# A request can finish or be aborted while its remote prefill is still
# pending, or after it becomes ready but before metadata is built.
# Clear both phases before any early return so stale request IDs cannot
# leak or be applied to a later request that reuses the same ID.
self._pd_decode_recompute_req_ids.discard(request.request_id)
self._pd_decode_recompute_ready_req_ids.discard(request.request_id)

params = request.kv_transfer_params
logger.debug(
"MooncakeConnector request_finished, request_status=%s, kv_transfer_params=%s", request.status, params
Expand Down
Loading
Loading