0
tests/ut/kv_offload/__init__.py
Normal file
0
tests/ut/kv_offload/__init__.py
Normal file
0
tests/ut/kv_offload/a2/__init__.py
Normal file
0
tests/ut/kv_offload/a2/__init__.py
Normal file
146
tests/ut/kv_offload/a2/test_remote_decode_lifecycle.py
Normal file
146
tests/ut/kv_offload/a2/test_remote_decode_lifecycle.py
Normal file
@@ -0,0 +1,146 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
# Copyright 2023 The vLLM team.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
# Adapted from vllm-project/vllm/blob/main/tests/v1/kv_connector/unit/test_remote_decode_lifecycle.py
|
||||
#
|
||||
import copy
|
||||
|
||||
from vllm.v1.outputs import EMPTY_MODEL_RUNNER_OUTPUT, KVConnectorOutput
|
||||
from vllm.v1.request import FinishReason, RequestStatus
|
||||
|
||||
from tests.ut.kv_offload.utils import (
|
||||
assert_scheduler_empty,
|
||||
create_model_runner_output,
|
||||
create_request,
|
||||
create_scheduler,
|
||||
create_vllm_config,
|
||||
)
|
||||
|
||||
|
||||
def test_basic_lifecycle():
|
||||
"""Test lifecycle of a Remote Decode request."""
|
||||
|
||||
vllm_config = create_vllm_config()
|
||||
scheduler = create_scheduler(vllm_config)
|
||||
|
||||
BLOCK_SIZE = vllm_config.cache_config.block_size
|
||||
NUM_EXTERNAL_FULL_BLOCKS = 2
|
||||
NUM_TOKENS = int(BLOCK_SIZE * (NUM_EXTERNAL_FULL_BLOCKS + 0.5))
|
||||
|
||||
request = create_request(
|
||||
request_id=1, max_tokens=1, num_tokens=NUM_TOKENS, do_remote_decode=True, block_size=BLOCK_SIZE
|
||||
)
|
||||
|
||||
scheduler.add_request(request)
|
||||
request_id = request.request_id
|
||||
|
||||
# STEP (1): Prefill.
|
||||
# (1a): schedule()
|
||||
scheduler_output = scheduler.schedule()
|
||||
assert len(scheduler.requests) == 1
|
||||
assert len(scheduler.running) == 1
|
||||
assert len(scheduler_output.scheduled_new_reqs) == 1
|
||||
|
||||
# (1b): execute_model()
|
||||
model_runner_output = create_model_runner_output(reqs=[request])
|
||||
|
||||
# (1c): update_from_output()
|
||||
engine_core_outputs = scheduler.update_from_output(scheduler_output, model_runner_output)
|
||||
|
||||
assert request.is_finished()
|
||||
assert request.status == RequestStatus.FINISHED_LENGTH_CAPPED
|
||||
output = engine_core_outputs[0].outputs[0]
|
||||
assert output.finish_reason == FinishReason.LENGTH
|
||||
# MooncakeConnector.request_finished returns (delay_free_blocks, None),
|
||||
# so kv_transfer_params is None in the output.
|
||||
assert output.kv_transfer_params is None
|
||||
|
||||
# Request freed in Scheduler but blocks should not be freed.
|
||||
assert request_id in scheduler.finished_req_ids
|
||||
assert len(scheduler.running) == 0
|
||||
assert len(scheduler.waiting) == 0
|
||||
assert len(scheduler.requests) == 1
|
||||
blocks = scheduler.kv_cache_manager.coordinator.single_type_managers[0].req_to_blocks[request_id]
|
||||
for block in blocks:
|
||||
assert block.ref_cnt == 1
|
||||
|
||||
# STEP (2): Send Finished to PB.
|
||||
scheduler_output = scheduler.schedule()
|
||||
assert len(scheduler.requests) == 1
|
||||
assert len(scheduler.running) == 0
|
||||
assert len(scheduler_output.finished_req_ids) == 1
|
||||
assert request_id in scheduler_output.finished_req_ids
|
||||
assert len(scheduler_output.scheduled_new_reqs) == 0
|
||||
assert scheduler_output.scheduled_cached_reqs.num_reqs == 0
|
||||
assert len(scheduler.finished_req_ids) == 0
|
||||
|
||||
model_runner_output = EMPTY_MODEL_RUNNER_OUTPUT
|
||||
scheduler.update_from_output(scheduler_output, model_runner_output)
|
||||
|
||||
# STEP (3): Finished sending.
|
||||
scheduler_output = scheduler.schedule()
|
||||
assert len(scheduler.requests) == 1
|
||||
assert len(scheduler.running) == 0
|
||||
assert len(scheduler_output.finished_req_ids) == 0
|
||||
assert len(scheduler_output.scheduled_new_reqs) == 0
|
||||
assert scheduler_output.scheduled_cached_reqs.num_reqs == 0
|
||||
assert len(scheduler.finished_req_ids) == 0
|
||||
|
||||
model_runner_output = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
|
||||
model_runner_output.kv_connector_output = KVConnectorOutput(finished_sending={request_id})
|
||||
|
||||
scheduler.update_from_output(scheduler_output, model_runner_output)
|
||||
|
||||
assert_scheduler_empty(scheduler)
|
||||
|
||||
|
||||
def test_prefix_cache_lifecycle():
|
||||
"""Test that remote decode params still works with a prefix cache hit."""
|
||||
|
||||
vllm_config = create_vllm_config()
|
||||
scheduler = create_scheduler(vllm_config)
|
||||
|
||||
# Prime the KVCache.
|
||||
BLOCK_SIZE = vllm_config.cache_config.block_size
|
||||
NUM_EXTERNAL_FULL_BLOCKS = 3
|
||||
NUM_TOKENS = int(BLOCK_SIZE * (NUM_EXTERNAL_FULL_BLOCKS + 0.5))
|
||||
|
||||
request_normal = create_request(request_id=1, num_tokens=NUM_TOKENS, block_size=BLOCK_SIZE)
|
||||
|
||||
scheduler.add_request(request_normal)
|
||||
scheduler_output = scheduler.schedule()
|
||||
model_runner_output = create_model_runner_output(reqs=[request_normal], use_eos=True)
|
||||
scheduler.update_from_output(scheduler_output, model_runner_output)
|
||||
scheduler_output = scheduler.schedule()
|
||||
scheduler.update_from_output(scheduler_output, EMPTY_MODEL_RUNNER_OUTPUT)
|
||||
|
||||
# Step (1): Send the KV Transfer.
|
||||
NUM_EXTERNAL_FULL_BLOCKS -= 1
|
||||
NUM_TOKENS = int(BLOCK_SIZE * (NUM_EXTERNAL_FULL_BLOCKS + 0.5))
|
||||
|
||||
request_remote = create_request(request_id=1, num_tokens=NUM_TOKENS, do_remote_decode=True, block_size=BLOCK_SIZE)
|
||||
|
||||
scheduler.add_request(request_remote)
|
||||
scheduler_output = scheduler.schedule()
|
||||
model_runner_output = create_model_runner_output(reqs=[request_remote])
|
||||
scheduler.update_from_output(scheduler_output, model_runner_output)
|
||||
|
||||
# STEP (2): Ensure it is freed.
|
||||
scheduler_output = scheduler.schedule()
|
||||
model_runner_output = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
|
||||
model_runner_output.kv_connector_output = KVConnectorOutput(finished_sending={request_remote.request_id})
|
||||
scheduler.update_from_output(scheduler_output, model_runner_output)
|
||||
assert_scheduler_empty(scheduler)
|
||||
209
tests/ut/kv_offload/a2/test_remote_prefill_lifecycle.py
Normal file
209
tests/ut/kv_offload/a2/test_remote_prefill_lifecycle.py
Normal file
@@ -0,0 +1,209 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
# Copyright 2023 The vLLM team.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
# Adapted from vllm-project/vllm/blob/main/tests/v1/kv_connector/unit/test_remote_prefill_lifecycle.py
|
||||
#
|
||||
import copy
|
||||
|
||||
from vllm.v1.outputs import EMPTY_MODEL_RUNNER_OUTPUT, KVConnectorOutput
|
||||
from vllm.v1.request import RequestStatus
|
||||
|
||||
from tests.ut.kv_offload.utils import (
|
||||
assert_scheduler_empty,
|
||||
create_model_runner_output,
|
||||
create_request,
|
||||
create_scheduler,
|
||||
create_vllm_config,
|
||||
)
|
||||
|
||||
|
||||
def _num_waiting_requests(scheduler) -> int:
|
||||
return len(scheduler.waiting) + len(scheduler.skipped_waiting)
|
||||
|
||||
|
||||
def test_basic_lifecycle():
|
||||
"""Test lifecycle of a remote prefill."""
|
||||
|
||||
vllm_config = create_vllm_config()
|
||||
scheduler = create_scheduler(vllm_config)
|
||||
|
||||
BLOCK_SIZE = vllm_config.cache_config.block_size
|
||||
NUM_EXTERNAL_FULL_BLOCKS = 2
|
||||
NUM_TOKENS = int(BLOCK_SIZE * (NUM_EXTERNAL_FULL_BLOCKS + 0.5))
|
||||
START_FREE_BLOCK_QUEUE_SIZE = scheduler.kv_cache_manager.block_pool.free_block_queue.num_free_blocks
|
||||
|
||||
request = create_request(request_id=1, num_tokens=NUM_TOKENS, do_remote_prefill=True, block_size=BLOCK_SIZE)
|
||||
|
||||
scheduler.add_request(request)
|
||||
request_id = request.request_id
|
||||
|
||||
# STEP (1):
|
||||
# (1a): schedule()
|
||||
scheduler_output = scheduler.schedule()
|
||||
|
||||
assert len(scheduler.running) == 0
|
||||
assert len(scheduler_output.scheduled_new_reqs) == 0
|
||||
assert scheduler_output.scheduled_cached_reqs.num_reqs == 0
|
||||
assert len(scheduler_output.num_scheduled_tokens) == 0
|
||||
assert scheduler_output.total_num_scheduled_tokens == 0
|
||||
|
||||
assert _num_waiting_requests(scheduler) == 1
|
||||
assert request in scheduler.skipped_waiting
|
||||
assert request.status == RequestStatus.WAITING_FOR_REMOTE_KVS
|
||||
assert request.num_computed_tokens == NUM_TOKENS
|
||||
|
||||
block_pool = scheduler.kv_cache_manager.block_pool
|
||||
assert block_pool.free_block_queue.num_free_blocks < START_FREE_BLOCK_QUEUE_SIZE
|
||||
assert len(block_pool.cached_block_hash_to_block) == 0
|
||||
blocks = scheduler.kv_cache_manager.coordinator.single_type_managers[0].req_to_blocks[request_id]
|
||||
for block in blocks:
|
||||
assert block._block_hash is None
|
||||
|
||||
# (1b): forward()
|
||||
model_runner_output = EMPTY_MODEL_RUNNER_OUTPUT
|
||||
|
||||
# (1c): update_from_output()
|
||||
engine_core_outputs = scheduler.update_from_output(scheduler_output, model_runner_output)
|
||||
assert not engine_core_outputs or not engine_core_outputs[0].outputs
|
||||
|
||||
# STEP (2):
|
||||
# (2a): schedule(): nothing happens!
|
||||
scheduler_output = scheduler.schedule()
|
||||
assert _num_waiting_requests(scheduler) == 1
|
||||
assert len(scheduler.running) == 0
|
||||
|
||||
# (2b): forward(): request finishes recv.
|
||||
model_runner_output = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
|
||||
model_runner_output.kv_connector_output = KVConnectorOutput(finished_recving={request_id})
|
||||
|
||||
# (2c): update_from_output():
|
||||
engine_core_outputs = scheduler.update_from_output(scheduler_output, model_runner_output)
|
||||
assert _num_waiting_requests(scheduler) == 1
|
||||
assert request_id in scheduler.finished_recving_kv_req_ids
|
||||
|
||||
# STEP (3):
|
||||
# (3a): schedule(): this should actually schedule.
|
||||
scheduler_output = scheduler.schedule()
|
||||
assert len(scheduler.running) == 1
|
||||
|
||||
num_hashed_blocks = 0
|
||||
blocks = scheduler.kv_cache_manager.coordinator.single_type_managers[0].req_to_blocks[request_id]
|
||||
for block in blocks:
|
||||
assert block.ref_cnt == 1
|
||||
num_hashed_blocks += 1 if block._block_hash is not None else 0
|
||||
assert num_hashed_blocks == NUM_EXTERNAL_FULL_BLOCKS
|
||||
|
||||
scheduled_req = scheduler_output.scheduled_new_reqs[0]
|
||||
num_scheduled_tokens = scheduler_output.num_scheduled_tokens[request_id]
|
||||
num_computed_tokens = scheduled_req.num_computed_tokens
|
||||
total_prompt_tokens = len(scheduled_req.prompt_token_ids)
|
||||
assert num_scheduled_tokens == total_prompt_tokens - num_computed_tokens
|
||||
|
||||
# (3b): execute_model()
|
||||
model_runner_output = create_model_runner_output([request])
|
||||
# (3c): update_from_output()
|
||||
scheduler.update_from_output(scheduler_output, model_runner_output)
|
||||
|
||||
# Step (4): Hit EOS.
|
||||
scheduler_output = scheduler.schedule()
|
||||
model_runner_output = create_model_runner_output([request], use_eos=True)
|
||||
scheduler.update_from_output(scheduler_output, model_runner_output)
|
||||
scheduler.schedule()
|
||||
|
||||
assert_scheduler_empty(scheduler)
|
||||
|
||||
|
||||
def test_no_spurious_prefix_caching():
|
||||
"""With P/D, blocks can be allocated but uncomputed for multiple engine steps.
|
||||
This test confirms that we do not accidentally have cache hits against
|
||||
uncomputed blocks."""
|
||||
|
||||
vllm_config = create_vllm_config()
|
||||
scheduler = create_scheduler(vllm_config)
|
||||
|
||||
BLOCK_SIZE = vllm_config.cache_config.block_size
|
||||
NUM_EXTERNAL_FULL_BLOCKS = 2
|
||||
NUM_TOKENS = int(BLOCK_SIZE * (NUM_EXTERNAL_FULL_BLOCKS + 0.5))
|
||||
|
||||
request_remote = create_request(
|
||||
request_id=1,
|
||||
num_tokens=NUM_TOKENS,
|
||||
do_remote_prefill=True,
|
||||
block_size=BLOCK_SIZE,
|
||||
)
|
||||
|
||||
scheduler.add_request(request_remote)
|
||||
scheduler_output = scheduler.schedule()
|
||||
scheduler.update_from_output(scheduler_output, EMPTY_MODEL_RUNNER_OUTPUT)
|
||||
assert _num_waiting_requests(scheduler) == 1
|
||||
|
||||
remote_blocks = scheduler.kv_cache_manager.coordinator.single_type_managers[0].req_to_blocks[
|
||||
request_remote.request_id
|
||||
]
|
||||
|
||||
for block in remote_blocks:
|
||||
assert block.ref_cnt == 1
|
||||
assert block._block_hash is None
|
||||
|
||||
|
||||
def test_full_block_prompt():
|
||||
"""Test that we handle a prompt that is the full block size."""
|
||||
|
||||
vllm_config = create_vllm_config()
|
||||
scheduler = create_scheduler(vllm_config)
|
||||
|
||||
BLOCK_SIZE = vllm_config.cache_config.block_size
|
||||
NUM_EXTERNAL_FULL_BLOCKS = 2
|
||||
NUM_TOKENS = int(BLOCK_SIZE * NUM_EXTERNAL_FULL_BLOCKS)
|
||||
|
||||
request = create_request(request_id=1, num_tokens=NUM_TOKENS, do_remote_prefill=True, block_size=BLOCK_SIZE)
|
||||
|
||||
scheduler.add_request(request)
|
||||
request_id = request.request_id
|
||||
|
||||
# STEP (1): Initialize a recv.
|
||||
scheduler_output = scheduler.schedule()
|
||||
num_blocks = len(scheduler.kv_cache_manager.coordinator.single_type_managers[0].req_to_blocks[request_id])
|
||||
assert num_blocks == NUM_EXTERNAL_FULL_BLOCKS
|
||||
model_runner_output = EMPTY_MODEL_RUNNER_OUTPUT
|
||||
scheduler.update_from_output(scheduler_output, model_runner_output)
|
||||
|
||||
# STEP (2): Recv.
|
||||
scheduler_output = scheduler.schedule()
|
||||
model_runner_output = copy.deepcopy(EMPTY_MODEL_RUNNER_OUTPUT)
|
||||
model_runner_output.kv_connector_output = KVConnectorOutput(finished_recving={request_id})
|
||||
scheduler.update_from_output(scheduler_output, model_runner_output)
|
||||
assert _num_waiting_requests(scheduler) == 1
|
||||
assert request_id in scheduler.finished_recving_kv_req_ids
|
||||
|
||||
# STEP (3): Run as usual.
|
||||
scheduler_output = scheduler.schedule()
|
||||
|
||||
num_blocks = len(scheduler.kv_cache_manager.coordinator.single_type_managers[0].req_to_blocks[request_id])
|
||||
assert num_blocks == NUM_EXTERNAL_FULL_BLOCKS
|
||||
assert scheduler_output.scheduled_new_reqs[0].num_computed_tokens == NUM_TOKENS - 1
|
||||
assert scheduler_output.num_scheduled_tokens[request_id] == 1
|
||||
|
||||
model_runner_output = create_model_runner_output([request])
|
||||
scheduler.update_from_output(scheduler_output, model_runner_output)
|
||||
|
||||
# Step (4): Hit EOS.
|
||||
scheduler_output = scheduler.schedule()
|
||||
model_runner_output = create_model_runner_output([request], use_eos=True)
|
||||
scheduler.update_from_output(scheduler_output, model_runner_output)
|
||||
scheduler.schedule()
|
||||
|
||||
assert_scheduler_empty(scheduler)
|
||||
3172
tests/ut/kv_offload/test_mooncake_connector.py
Normal file
3172
tests/ut/kv_offload/test_mooncake_connector.py
Normal file
File diff suppressed because it is too large
Load Diff
296
tests/ut/kv_offload/test_mooncake_hybrid_connector.py
Normal file
296
tests/ut/kv_offload/test_mooncake_hybrid_connector.py
Normal file
@@ -0,0 +1,296 @@
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import types
|
||||
import unittest
|
||||
from collections import defaultdict, deque
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
fake_engine = types.ModuleType("mooncake.engine")
|
||||
fake_engine.TransferEngine = MagicMock() # type: ignore[attr-defined]
|
||||
sys.modules["mooncake.engine"] = fake_engine
|
||||
|
||||
from vllm.v1.request import RequestStatus # noqa: E402
|
||||
|
||||
from vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_hybrid_connector import ( # noqa: E402
|
||||
MAX_REQUESTS_PER_PEER_HANDLER,
|
||||
KVCacheRecvingThread,
|
||||
MooncakeConnectorScheduler,
|
||||
)
|
||||
|
||||
|
||||
class MockRequest:
|
||||
def __init__(
|
||||
self,
|
||||
request_id,
|
||||
prompt_token_ids,
|
||||
kv_transfer_params,
|
||||
status,
|
||||
num_prompt_tokens=None,
|
||||
):
|
||||
self.request_id = request_id
|
||||
self.prompt_token_ids = prompt_token_ids
|
||||
if num_prompt_tokens is None:
|
||||
num_prompt_tokens = len(prompt_token_ids) if prompt_token_ids is not None else 0
|
||||
self.num_prompt_tokens = num_prompt_tokens
|
||||
self.kv_transfer_params = kv_transfer_params
|
||||
self.status = status
|
||||
self.output_token_ids = [101]
|
||||
|
||||
|
||||
class TestHybridKVCacheRecvingThreadDispatch(unittest.TestCase):
|
||||
def _make_thread(self):
|
||||
thread = object.__new__(KVCacheRecvingThread)
|
||||
thread.executor = ThreadPoolExecutor(max_workers=2)
|
||||
thread.peer_request_queues = defaultdict(deque)
|
||||
thread.active_peer_request_handlers = set()
|
||||
thread.peer_request_queues_lock = threading.Lock()
|
||||
thread.request_task_counts = defaultdict(int)
|
||||
thread.finished_request_markers = set()
|
||||
thread.request_task_counts_lock = threading.Lock()
|
||||
return thread
|
||||
|
||||
def test_executor_workers_bind_kv_cache_device_before_handling_requests(self):
|
||||
expected_device = torch.device("npu:5")
|
||||
kv_cache = MagicMock(device=expected_device)
|
||||
model_config = types.SimpleNamespace(
|
||||
is_deepseek_mla=False,
|
||||
hf_config=types.SimpleNamespace(compress_ratios=[1]),
|
||||
hf_text_config=types.SimpleNamespace(num_hidden_layers=1),
|
||||
)
|
||||
vllm_config = types.SimpleNamespace(
|
||||
model_config=model_config,
|
||||
cache_config=types.SimpleNamespace(block_size=16),
|
||||
)
|
||||
kv_cache_config = types.SimpleNamespace(kv_cache_groups=[])
|
||||
worker_events: defaultdict[int, list[tuple[str, int | str]]] = defaultdict(list)
|
||||
events_lock = threading.Lock()
|
||||
both_workers_started = threading.Event()
|
||||
release_workers = threading.Event()
|
||||
|
||||
def record_set_device(device):
|
||||
device_index = device if isinstance(device, int) else torch.device(device).index
|
||||
with events_lock:
|
||||
worker_events[threading.get_ident()].append(("set_device", device_index))
|
||||
|
||||
with (
|
||||
patch("torch.npu.set_device", side_effect=record_set_device),
|
||||
patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake_hybrid_connector.is_vl_model",
|
||||
return_value=False,
|
||||
),
|
||||
):
|
||||
thread = KVCacheRecvingThread(
|
||||
tp_rank=1,
|
||||
tp_size=2,
|
||||
_prefill_pp_size=1,
|
||||
engine=MagicMock(),
|
||||
local_engine_id="local_engine",
|
||||
local_handshake_port=5555,
|
||||
side_channel_port=30000,
|
||||
local_kv_caches_base_addr=[0x1000],
|
||||
block_len_per_addr=[1024],
|
||||
block_stride_per_addr=[1024],
|
||||
addr_group_idx=[0],
|
||||
mamba_ssm_size=(0, 0),
|
||||
use_hybrid=False,
|
||||
has_mamba=False,
|
||||
hma_group_size=1,
|
||||
ready_event=threading.Event(),
|
||||
vllm_config=vllm_config,
|
||||
kv_cache_config=kv_cache_config,
|
||||
kv_caches={"layer.0": (kv_cache, kv_cache)},
|
||||
)
|
||||
|
||||
def handle_request(req_meta: dict[str, Any]):
|
||||
with events_lock:
|
||||
worker_events[threading.get_ident()].append(("handle", req_meta["request_id"]))
|
||||
handled_worker_count = sum(
|
||||
any(event == "handle" for event, _ in events) for events in worker_events.values()
|
||||
)
|
||||
if handled_worker_count == 2:
|
||||
both_workers_started.set()
|
||||
release_workers.wait()
|
||||
|
||||
thread._handle_request = handle_request # type: ignore[method-assign]
|
||||
try:
|
||||
for index in range(2):
|
||||
thread._submit_request(
|
||||
{
|
||||
"request_id": f"req-{index}",
|
||||
"remote_host": f"host-{index}",
|
||||
"remote_handshake_port": 6000 + index,
|
||||
"all_task_done": True,
|
||||
}
|
||||
)
|
||||
self.assertTrue(both_workers_started.wait(timeout=5.0), "executor did not start two workers")
|
||||
finally:
|
||||
release_workers.set()
|
||||
thread.executor.shutdown(wait=True, cancel_futures=True)
|
||||
|
||||
handled_worker_events = [events for events in worker_events.values() if any(e == "handle" for e, _ in events)]
|
||||
self.assertEqual(len(handled_worker_events), 2)
|
||||
for events in handled_worker_events:
|
||||
self.assertEqual(events[0], ("set_device", expected_device.index))
|
||||
self.assertEqual(events[1][0], "handle")
|
||||
|
||||
def test_submit_request_serializes_same_peer_fifo(self):
|
||||
thread = self._make_thread()
|
||||
release_first_request = threading.Event()
|
||||
first_request_started = threading.Event()
|
||||
other_peer_started = threading.Event()
|
||||
handled_requests: list[str] = []
|
||||
active_by_peer: defaultdict[tuple[str, int], int] = defaultdict(int)
|
||||
max_active_by_peer: defaultdict[tuple[str, int], int] = defaultdict(int)
|
||||
state_lock = threading.Lock()
|
||||
|
||||
def handle_request(req_meta: dict[str, Any]):
|
||||
peer_key = (req_meta["remote_host"], req_meta["remote_handshake_port"])
|
||||
with state_lock:
|
||||
active_by_peer[peer_key] += 1
|
||||
max_active_by_peer[peer_key] = max(max_active_by_peer[peer_key], active_by_peer[peer_key])
|
||||
handled_requests.append(req_meta["request_id"])
|
||||
|
||||
if req_meta["request_id"] == "same-peer-1":
|
||||
first_request_started.set()
|
||||
self.assertTrue(release_first_request.wait(timeout=2.0))
|
||||
elif req_meta["request_id"] == "other-peer-1":
|
||||
other_peer_started.set()
|
||||
|
||||
time.sleep(0.01)
|
||||
with state_lock:
|
||||
active_by_peer[peer_key] -= 1
|
||||
|
||||
thread._handle_request = handle_request # type: ignore[method-assign]
|
||||
same_peer_1 = {
|
||||
"request_id": "same-peer-1",
|
||||
"remote_host": "host-a",
|
||||
"remote_handshake_port": 6000,
|
||||
"all_task_done": False,
|
||||
}
|
||||
same_peer_2 = {
|
||||
"request_id": "same-peer-2",
|
||||
"remote_host": "host-a",
|
||||
"remote_handshake_port": 6000,
|
||||
"all_task_done": True,
|
||||
}
|
||||
other_peer = {
|
||||
"request_id": "other-peer-1",
|
||||
"remote_host": "host-b",
|
||||
"remote_handshake_port": 6001,
|
||||
"all_task_done": True,
|
||||
}
|
||||
|
||||
try:
|
||||
thread._submit_request(same_peer_1)
|
||||
self.assertTrue(first_request_started.wait(timeout=1.0))
|
||||
thread._submit_request(same_peer_2)
|
||||
thread._submit_request(other_peer)
|
||||
|
||||
self.assertTrue(other_peer_started.wait(timeout=1.0))
|
||||
time.sleep(0.05)
|
||||
self.assertNotIn("same-peer-2", handled_requests)
|
||||
finally:
|
||||
release_first_request.set()
|
||||
thread.executor.shutdown(wait=True, cancel_futures=True)
|
||||
|
||||
self.assertLess(handled_requests.index("same-peer-1"), handled_requests.index("same-peer-2"))
|
||||
self.assertEqual(max_active_by_peer[("host-a", 6000)], 1)
|
||||
self.assertEqual(max_active_by_peer[("host-b", 6001)], 1)
|
||||
|
||||
def test_peer_handler_yields_after_batch_limit(self):
|
||||
thread = self._make_thread()
|
||||
peer_key = ("host-a", 6000)
|
||||
requests = [
|
||||
{
|
||||
"request_id": f"req-{idx}",
|
||||
"remote_host": peer_key[0],
|
||||
"remote_handshake_port": peer_key[1],
|
||||
}
|
||||
for idx in range(MAX_REQUESTS_PER_PEER_HANDLER + 1)
|
||||
]
|
||||
handled_requests: list[str] = []
|
||||
thread.peer_request_queues[peer_key].extend(requests)
|
||||
thread.active_peer_request_handlers.add(peer_key)
|
||||
thread.executor = MagicMock()
|
||||
|
||||
def handle_request(req_meta: dict[str, Any]):
|
||||
handled_requests.append(req_meta["request_id"])
|
||||
|
||||
thread._handle_request = handle_request # type: ignore[method-assign]
|
||||
|
||||
thread._handle_peer_requests(peer_key)
|
||||
|
||||
self.assertEqual(handled_requests, [f"req-{idx}" for idx in range(MAX_REQUESTS_PER_PEER_HANDLER)])
|
||||
self.assertEqual(
|
||||
[req["request_id"] for req in thread.peer_request_queues[peer_key]],
|
||||
[f"req-{MAX_REQUESTS_PER_PEER_HANDLER}"],
|
||||
)
|
||||
self.assertIn(peer_key, thread.active_peer_request_handlers)
|
||||
thread.executor.submit.assert_called_once_with(thread._handle_peer_requests, peer_key)
|
||||
|
||||
|
||||
class TestMooncakeHybridConnectorScheduler(unittest.TestCase):
|
||||
def _make_scheduler(self):
|
||||
scheduler = object.__new__(MooncakeConnectorScheduler)
|
||||
scheduler.use_hybrid = True
|
||||
scheduler.use_compress = True
|
||||
scheduler.num_swa_blocks = [0, 2]
|
||||
scheduler.group_block_size = [128, 128]
|
||||
scheduler.group_compress_ratio = [4, 1]
|
||||
scheduler._reqs_need_send = {}
|
||||
scheduler.block_size = 128
|
||||
scheduler.engine_id = "engine"
|
||||
scheduler.side_channel_host = "127.0.0.1"
|
||||
scheduler.side_channel_port = 12345
|
||||
scheduler.tp_size = 1
|
||||
scheduler.multi_nodes_meta_mapping = {}
|
||||
return scheduler
|
||||
|
||||
def test_compute_transfer_block_ids_trims_swa_groups(self):
|
||||
scheduler = self._make_scheduler()
|
||||
block_ids = (list(range(10)), [100, 101, 102, 103])
|
||||
|
||||
transfer_block_ids = scheduler._compute_transfer_block_ids(block_ids, prompt_len=129)
|
||||
|
||||
self.assertEqual(transfer_block_ids, ([0], [100, 101]))
|
||||
|
||||
def test_request_finished_trims_before_swa_clip(self):
|
||||
scheduler = self._make_scheduler()
|
||||
request = MockRequest(
|
||||
"req1",
|
||||
prompt_token_ids=list(range(129)),
|
||||
kv_transfer_params={"do_remote_decode": True},
|
||||
status=RequestStatus.FINISHED_LENGTH_CAPPED,
|
||||
)
|
||||
block_ids = (list(range(10)), [100, 101, 102, 103])
|
||||
|
||||
delay_free, params = scheduler.request_finished_all_groups(request, block_ids)
|
||||
|
||||
self.assertTrue(delay_free)
|
||||
self.assertIsNotNone(params)
|
||||
self.assertEqual(params["remote_block_ids"], ([0], [100, 101]))
|
||||
self.assertEqual(params["num_prompt_blocks"], 2)
|
||||
self.assertIn("req1", scheduler._reqs_need_send)
|
||||
|
||||
def test_request_finished_uses_num_prompt_tokens(self):
|
||||
scheduler = self._make_scheduler()
|
||||
request = MockRequest(
|
||||
"req1",
|
||||
prompt_token_ids=None,
|
||||
kv_transfer_params={"do_remote_decode": True},
|
||||
status=RequestStatus.FINISHED_LENGTH_CAPPED,
|
||||
num_prompt_tokens=129,
|
||||
)
|
||||
block_ids = (list(range(10)), [100, 101, 102, 103])
|
||||
|
||||
delay_free, params = scheduler.request_finished_all_groups(request, block_ids)
|
||||
|
||||
self.assertTrue(delay_free)
|
||||
self.assertIsNotNone(params)
|
||||
self.assertEqual(params["remote_block_ids"], ([0], [100, 101]))
|
||||
self.assertEqual(params["num_prompt_blocks"], 2)
|
||||
1172
tests/ut/kv_offload/test_mooncake_layerwise_connector.py
Normal file
1172
tests/ut/kv_offload/test_mooncake_layerwise_connector.py
Normal file
File diff suppressed because it is too large
Load Diff
640
tests/ut/kv_offload/test_recompute_cpu_offload.py
Normal file
640
tests/ut/kv_offload/test_recompute_cpu_offload.py
Normal file
@@ -0,0 +1,640 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
import sys
|
||||
import types
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from vllm.v1.outputs import KVConnectorOutput
|
||||
from vllm.v1.sample.rejection_sampler import PLACEHOLDER_TOKEN_ID
|
||||
|
||||
# Clean up stale mock modules installed by other kv offload tests that replace
|
||||
# real kv_transfer packages with fake modules, breaking imports of this package.
|
||||
_kv_xfer = "vllm_ascend.distributed.kv_transfer"
|
||||
_vllm_kv_xfer = "vllm.distributed.kv_transfer"
|
||||
_saved_modules: dict[str, types.ModuleType] = {}
|
||||
_to_remove = []
|
||||
for _module_name in list(sys.modules):
|
||||
if _module_name.startswith(_kv_xfer) or _module_name.startswith(_vllm_kv_xfer):
|
||||
_to_remove.append(_module_name)
|
||||
for _module_name in _to_remove:
|
||||
_saved_modules[_module_name] = sys.modules.pop(_module_name)
|
||||
|
||||
from vllm_ascend.core.recompute_scheduler import RecomputeScheduler # noqa: E402
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.recompute_cpu_offload.manager import ( # noqa: E402
|
||||
PreemptedRequestState,
|
||||
RecomputeCPUOffloadScheduler,
|
||||
TransferMeta,
|
||||
)
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.recompute_cpu_offload.metadata import ( # noqa: E402
|
||||
INVALID_JOB_ID,
|
||||
RecomputeCPUOffloadMetadata,
|
||||
RecomputeCPUOffloadWorkerMetadata,
|
||||
)
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.recompute_cpu_offload.recompute_cpu_offload_connector import ( # noqa: E402
|
||||
RecomputeCPUOffloadConnectorV1,
|
||||
)
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.recompute_cpu_offload.worker import ( # noqa: E402
|
||||
RecomputeCPUOffloadWorker,
|
||||
)
|
||||
|
||||
for _module_name, _module in _saved_modules.items():
|
||||
sys.modules[_module_name] = _module
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_worker_metadata_aggregate():
|
||||
metadata = RecomputeCPUOffloadWorkerMetadata(completed_store_events={1: 1, 2: 2})
|
||||
other = RecomputeCPUOffloadWorkerMetadata(completed_store_events={2: 3, 4: 1})
|
||||
|
||||
merged = metadata.aggregate(other)
|
||||
|
||||
assert isinstance(merged, RecomputeCPUOffloadWorkerMetadata)
|
||||
assert merged.completed_store_events == {1: 1, 2: 5, 4: 1}
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_metadata_defaults_are_empty():
|
||||
metadata = RecomputeCPUOffloadMetadata()
|
||||
|
||||
assert metadata.need_flush is False
|
||||
assert metadata.preempt_store_event == INVALID_JOB_ID
|
||||
assert metadata.preempt_store_gpu_blocks == []
|
||||
assert metadata.preempt_store_cpu_blocks == []
|
||||
assert metadata.preempt_load_event == INVALID_JOB_ID
|
||||
assert metadata.preempt_load_gpu_blocks == []
|
||||
assert metadata.preempt_load_cpu_blocks == []
|
||||
assert metadata.preempt_load_event_to_reqs == {}
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_connector_scheduler_methods_forward():
|
||||
connector = RecomputeCPUOffloadConnectorV1.__new__(RecomputeCPUOffloadConnectorV1)
|
||||
scheduler_manager = MagicMock()
|
||||
scheduler_manager.get_num_new_matched_tokens.return_value = (8, True)
|
||||
scheduler_manager.update_state_before_preempt.return_value = True
|
||||
scheduler_manager.has_pending_transfers.return_value = True
|
||||
scheduler_manager.has_preempted_request.return_value = True
|
||||
connector.scheduler_manager = scheduler_manager
|
||||
|
||||
request = SimpleNamespace(request_id="req-1")
|
||||
blocks = MagicMock()
|
||||
block_ids = ([1, 2],)
|
||||
|
||||
assert connector.get_num_new_matched_tokens(request, 4) == (8, True)
|
||||
connector.update_state_after_alloc(request, blocks, 8)
|
||||
assert connector.update_state_before_preempt(request, block_ids, 16) is True
|
||||
assert connector.has_pending_transfers() is True
|
||||
assert connector.has_preempted_request("req-1") is True
|
||||
|
||||
scheduler_manager.get_num_new_matched_tokens.assert_called_once_with(request, 4)
|
||||
scheduler_manager.update_state_after_alloc.assert_called_once_with(request, blocks, 8)
|
||||
scheduler_manager.update_state_before_preempt.assert_called_once_with(request, block_ids, 16)
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_connector_worker_methods_forward():
|
||||
connector = RecomputeCPUOffloadConnectorV1.__new__(RecomputeCPUOffloadConnectorV1)
|
||||
worker_handler = MagicMock()
|
||||
worker_handler.get_finished.return_value = (None, {"req-1"})
|
||||
worker_handler.build_connector_worker_meta.return_value = RecomputeCPUOffloadWorkerMetadata(
|
||||
completed_store_events={3: 1}
|
||||
)
|
||||
connector.worker_handler = worker_handler
|
||||
|
||||
metadata = RecomputeCPUOffloadMetadata(preempt_load_event=3)
|
||||
connector.bind_connector_metadata(metadata)
|
||||
connector.handle_preemptions(metadata)
|
||||
connector.start_load_kv(MagicMock())
|
||||
connector.wait_for_layer_load("layer.0")
|
||||
|
||||
assert connector.get_finished(set()) == (None, {"req-1"})
|
||||
assert connector.build_connector_worker_meta().completed_store_events == {3: 1}
|
||||
|
||||
worker_handler.bind_connector_metadata.assert_called_once_with(metadata)
|
||||
worker_handler.handle_preemptions.assert_called_once_with(metadata)
|
||||
worker_handler.start_load_kv.assert_called_once_with()
|
||||
worker_handler.wait_for_layer_load.assert_called_once_with()
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_connector_defaults_without_scheduler_manager():
|
||||
connector = RecomputeCPUOffloadConnectorV1.__new__(RecomputeCPUOffloadConnectorV1)
|
||||
connector.scheduler_manager = None
|
||||
|
||||
assert connector.get_num_new_matched_tokens(MagicMock(), 0) == (0, False)
|
||||
assert connector.update_state_before_preempt(MagicMock(), ([],), 1) is False
|
||||
assert isinstance(
|
||||
connector.build_connector_meta(MagicMock()),
|
||||
RecomputeCPUOffloadMetadata,
|
||||
)
|
||||
assert connector.request_finished(MagicMock(), []) == (False, None)
|
||||
assert connector.request_finished_all_groups(MagicMock(), ([],)) == (
|
||||
False,
|
||||
None,
|
||||
)
|
||||
assert connector.has_pending_transfers() is False
|
||||
assert connector.has_preempted_request("req-1") is False
|
||||
assert connector.take_events() == []
|
||||
assert connector.reset_cache() is None
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_scheduler_get_num_new_matched_tokens_states():
|
||||
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
|
||||
scheduler._preempted_req_states = {}
|
||||
scheduler._cleanup_preempt_cache_request = MagicMock()
|
||||
request = SimpleNamespace(request_id="req-1", num_tokens=10)
|
||||
|
||||
assert scheduler.get_num_new_matched_tokens(request, 0) == (0, False)
|
||||
|
||||
scheduler._preempted_req_states["req-1"] = PreemptedRequestState(
|
||||
req_id="req-1",
|
||||
cpu_block_ids=([1],),
|
||||
num_computed_tokens=8,
|
||||
store_transfer_meta=TransferMeta([11], [1]),
|
||||
ready=False,
|
||||
)
|
||||
assert scheduler.get_num_new_matched_tokens(request, 0) == (None, False)
|
||||
|
||||
scheduler._preempted_req_states["req-1"].ready = True
|
||||
assert scheduler.get_num_new_matched_tokens(request, 3) == (5, True)
|
||||
assert scheduler._preempted_req_states["req-1"].load_start_tokens == 3
|
||||
|
||||
assert scheduler.get_num_new_matched_tokens(request, 8) == (0, False)
|
||||
scheduler._cleanup_preempt_cache_request.assert_called_once_with("req-1")
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_scheduler_update_state_after_alloc_errors():
|
||||
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
|
||||
scheduler._prepare_preempt_load_after_alloc = MagicMock(return_value=False)
|
||||
request = SimpleNamespace(request_id="req-1")
|
||||
blocks = MagicMock()
|
||||
blocks.get_block_ids.return_value = ([1, 2],)
|
||||
|
||||
scheduler.update_state_after_alloc(request, blocks, 0)
|
||||
scheduler._prepare_preempt_load_after_alloc.assert_not_called()
|
||||
|
||||
try:
|
||||
scheduler.update_state_after_alloc(request, blocks, 2)
|
||||
except RuntimeError as exc:
|
||||
assert "Failed to prepare recompute H2D load" in str(exc)
|
||||
else:
|
||||
raise AssertionError("Expected RuntimeError when load mapping fails")
|
||||
|
||||
scheduler._prepare_preempt_load_after_alloc.assert_called_once_with(request, ([1, 2],), 2)
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_scheduler_aligns_sliding_window_blocks():
|
||||
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
|
||||
scheduler._group_is_sliding_window = [True, False]
|
||||
|
||||
assert scheduler._align_group_block_ids(0, [7, 8], 4) == [0, 0, 7, 8]
|
||||
assert scheduler._align_group_block_ids(0, [5, 6, 7, 8, 9], 4) == [
|
||||
5,
|
||||
6,
|
||||
7,
|
||||
8,
|
||||
]
|
||||
assert scheduler._align_group_block_ids(1, [7, 8], 4) == [7, 8]
|
||||
assert scheduler._align_group_block_ids(0, [7, 8], 0) == []
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_scheduler_d2h_keeps_sliding_window_offsets():
|
||||
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
|
||||
scheduler._group_is_sliding_window = [True]
|
||||
scheduler.cpu_kv_cache_config = SimpleNamespace(
|
||||
kv_cache_groups=[SimpleNamespace(kv_cache_spec=SimpleNamespace(block_size=16))]
|
||||
)
|
||||
scheduler.enable_offload_prefix_caching = False
|
||||
scheduler._pending_hash_blocks = {}
|
||||
scheduler._gpu_block_pool = SimpleNamespace(
|
||||
blocks={
|
||||
20: SimpleNamespace(block_id=20, block_hash=None),
|
||||
21: SimpleNamespace(block_id=21, block_hash=None),
|
||||
},
|
||||
_maybe_evict_cached_block=MagicMock(),
|
||||
)
|
||||
cpu_blocks = [
|
||||
SimpleNamespace(block_id=101, _block_hash=None),
|
||||
SimpleNamespace(block_id=102, _block_hash=None),
|
||||
]
|
||||
scheduler.cpu_block_pool = SimpleNamespace(
|
||||
get_num_free_blocks=MagicMock(return_value=8),
|
||||
get_new_blocks=MagicMock(return_value=cpu_blocks),
|
||||
cached_block_hash_to_block=SimpleNamespace(get_one_block=MagicMock(return_value=None)),
|
||||
)
|
||||
scheduler._preempted_req_states = {}
|
||||
|
||||
assert scheduler._create_preempt_state("req-1", ([20, 21],), 64) is True
|
||||
|
||||
state = scheduler._preempted_req_states["req-1"]
|
||||
assert state.cpu_block_ids == ([0, 0, 101, 102],)
|
||||
assert state.store_transfer_meta == TransferMeta([20, 21], [101, 102])
|
||||
assert state.ready is False
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_scheduler_h2d_skips_sliding_window_null_blocks():
|
||||
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
|
||||
scheduler._group_is_sliding_window = [True]
|
||||
scheduler.cpu_kv_cache_config = SimpleNamespace(
|
||||
kv_cache_groups=[SimpleNamespace(kv_cache_spec=SimpleNamespace(block_size=16))]
|
||||
)
|
||||
scheduler._gpu_block_pool = SimpleNamespace(blocks={30: "gpu30", 31: "gpu31"}, touch=MagicMock())
|
||||
scheduler._preempted_req_states = {
|
||||
"req-1": PreemptedRequestState(
|
||||
req_id="req-1",
|
||||
cpu_block_ids=([0, 0, 4, 5],),
|
||||
num_computed_tokens=64,
|
||||
store_transfer_meta=TransferMeta([20, 21], [4, 5]),
|
||||
load_start_tokens=0,
|
||||
ready=True,
|
||||
)
|
||||
}
|
||||
|
||||
prepared = scheduler._prepare_preempt_load_after_alloc(
|
||||
SimpleNamespace(request_id="req-1"),
|
||||
([30, 31],),
|
||||
num_external_tokens=64,
|
||||
)
|
||||
|
||||
assert prepared is True
|
||||
state = scheduler._preempted_req_states["req-1"]
|
||||
assert state.load_transfer_meta == TransferMeta([30, 31], [4, 5])
|
||||
touched = list(scheduler._gpu_block_pool.touch.call_args.args[0])
|
||||
assert touched == ["gpu30", "gpu31"]
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_scheduler_h2d_clips_mtp_tail_blocks():
|
||||
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
|
||||
scheduler._group_is_sliding_window = [False]
|
||||
scheduler.cpu_kv_cache_config = SimpleNamespace(
|
||||
kv_cache_groups=[SimpleNamespace(kv_cache_spec=SimpleNamespace(block_size=16))]
|
||||
)
|
||||
scheduler._gpu_block_pool = SimpleNamespace(
|
||||
blocks={10: "gpu10", 11: "gpu11", 12: "gpu12"},
|
||||
touch=MagicMock(),
|
||||
)
|
||||
scheduler._preempted_req_states = {
|
||||
"req-1": PreemptedRequestState(
|
||||
req_id="req-1",
|
||||
cpu_block_ids=([1, 2, 3, 4],),
|
||||
num_computed_tokens=64,
|
||||
store_transfer_meta=TransferMeta([20, 21, 22, 23], [1, 2, 3, 4]),
|
||||
load_start_tokens=0,
|
||||
ready=True,
|
||||
)
|
||||
}
|
||||
|
||||
prepared = scheduler._prepare_preempt_load_after_alloc(
|
||||
SimpleNamespace(request_id="req-1"),
|
||||
([10, 11, 12],),
|
||||
num_external_tokens=64,
|
||||
)
|
||||
|
||||
assert prepared is True
|
||||
state = scheduler._preempted_req_states["req-1"]
|
||||
assert state.load_transfer_meta == TransferMeta([10, 11, 12], [1, 2, 3])
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_scheduler_build_connector_meta_assigns_events():
|
||||
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
|
||||
scheduler._store_event_counter = 4
|
||||
scheduler._load_event_counter = 7
|
||||
scheduler._preempt_store_event_to_blocks = {}
|
||||
scheduler._preempt_store_event_to_reqs = {}
|
||||
scheduler._preempt_load_event_to_reqs = {}
|
||||
scheduler._pending_hash_blocks = {"hash": MagicMock()}
|
||||
scheduler._preempted_req_states = {
|
||||
"store-req": PreemptedRequestState(
|
||||
req_id="store-req",
|
||||
cpu_block_ids=([2],),
|
||||
num_computed_tokens=8,
|
||||
store_transfer_meta=TransferMeta([10], [2]),
|
||||
ready=False,
|
||||
),
|
||||
"load-req": PreemptedRequestState(
|
||||
req_id="load-req",
|
||||
cpu_block_ids=([3],),
|
||||
num_computed_tokens=8,
|
||||
store_transfer_meta=TransferMeta([], []),
|
||||
load_transfer_meta=TransferMeta([11], [3]),
|
||||
ready=True,
|
||||
),
|
||||
}
|
||||
scheduler_output = SimpleNamespace(preempted_req_ids={"store-req"})
|
||||
|
||||
metadata = scheduler.build_connector_meta(scheduler_output)
|
||||
|
||||
assert metadata.need_flush is True
|
||||
assert metadata.preempt_store_event == 4
|
||||
assert metadata.preempt_store_gpu_blocks == [10]
|
||||
assert metadata.preempt_store_cpu_blocks == [2]
|
||||
assert metadata.preempt_load_event == 7
|
||||
assert metadata.preempt_load_gpu_blocks == [11]
|
||||
assert metadata.preempt_load_cpu_blocks == [3]
|
||||
assert metadata.preempt_load_event_to_reqs == {7: ["load-req"]}
|
||||
assert scheduler._preempted_req_states["store-req"].store_event == 4
|
||||
assert scheduler._preempted_req_states["load-req"].load_event == 7
|
||||
assert scheduler._pending_hash_blocks == {}
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_scheduler_update_connector_output_marks_store_ready():
|
||||
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
|
||||
scheduler._expected_worker_count = 2
|
||||
scheduler._store_event_pending_counts = {}
|
||||
scheduler._preempted_req_states = {}
|
||||
scheduler._process_preempt_store_event = MagicMock()
|
||||
output = KVConnectorOutput(
|
||||
finished_recving=set(),
|
||||
kv_connector_worker_meta=RecomputeCPUOffloadWorkerMetadata(completed_store_events={5: 1}),
|
||||
)
|
||||
|
||||
scheduler.update_connector_output(output)
|
||||
|
||||
assert scheduler._store_event_pending_counts == {5: 1}
|
||||
scheduler._process_preempt_store_event.assert_not_called()
|
||||
|
||||
scheduler.update_connector_output(output)
|
||||
|
||||
assert scheduler._store_event_pending_counts == {}
|
||||
scheduler._process_preempt_store_event.assert_called_once_with(5)
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_scheduler_request_finished_ready_and_pending():
|
||||
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
|
||||
scheduler._preempted_req_states = {
|
||||
"ready": PreemptedRequestState(
|
||||
req_id="ready",
|
||||
cpu_block_ids=([1],),
|
||||
num_computed_tokens=8,
|
||||
store_transfer_meta=TransferMeta([11], [1]),
|
||||
ready=True,
|
||||
),
|
||||
"pending": PreemptedRequestState(
|
||||
req_id="pending",
|
||||
cpu_block_ids=([2],),
|
||||
num_computed_tokens=8,
|
||||
store_transfer_meta=TransferMeta([12], [2]),
|
||||
ready=False,
|
||||
),
|
||||
"loading": PreemptedRequestState(
|
||||
req_id="loading",
|
||||
cpu_block_ids=([3],),
|
||||
num_computed_tokens=8,
|
||||
store_transfer_meta=TransferMeta([13], [3]),
|
||||
load_event=5,
|
||||
ready=True,
|
||||
),
|
||||
}
|
||||
scheduler._cleanup_preempt_cache_request = MagicMock()
|
||||
|
||||
assert scheduler.request_finished(SimpleNamespace(request_id="ready"), []) == (
|
||||
False,
|
||||
None,
|
||||
)
|
||||
assert scheduler.request_finished(SimpleNamespace(request_id="pending"), []) == (False, None)
|
||||
assert scheduler.request_finished(SimpleNamespace(request_id="loading"), []) == (False, None)
|
||||
|
||||
scheduler._cleanup_preempt_cache_request.assert_called_once_with("ready")
|
||||
assert scheduler._preempted_req_states["pending"].finished is True
|
||||
assert scheduler._preempted_req_states["loading"].finished is False
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_scheduler_process_store_event_finishes_pending_req():
|
||||
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
|
||||
cpu_block = MagicMock()
|
||||
cpu_block.block_hash = None
|
||||
scheduler.cpu_block_pool = SimpleNamespace(blocks={4: cpu_block})
|
||||
scheduler._preempt_store_event_to_blocks = {7: TransferMeta([1], [4])}
|
||||
scheduler._preempt_store_event_to_reqs = {7: ["req-1"]}
|
||||
scheduler._preempted_req_states = {
|
||||
"req-1": PreemptedRequestState(
|
||||
req_id="req-1",
|
||||
cpu_block_ids=([4],),
|
||||
num_computed_tokens=8,
|
||||
store_transfer_meta=TransferMeta([1], [4]),
|
||||
ready=False,
|
||||
finished=True,
|
||||
)
|
||||
}
|
||||
scheduler._cleanup_preempt_cache_request = MagicMock()
|
||||
|
||||
scheduler._process_preempt_store_event(7)
|
||||
|
||||
assert scheduler._preempted_req_states["req-1"].ready is True
|
||||
scheduler._cleanup_preempt_cache_request.assert_called_once_with("req-1")
|
||||
assert scheduler._preempt_store_event_to_blocks == {}
|
||||
assert scheduler._preempt_store_event_to_reqs == {}
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_scheduler_pending_and_reset_cache_paths():
|
||||
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
|
||||
scheduler._store_event_pending_counts = {}
|
||||
scheduler._preempt_store_event_to_blocks = {}
|
||||
scheduler._preempted_req_states = {}
|
||||
|
||||
assert scheduler.has_pending_transfers() is False
|
||||
|
||||
scheduler._preempted_req_states["not-ready"] = PreemptedRequestState(
|
||||
req_id="not-ready",
|
||||
cpu_block_ids=([1],),
|
||||
num_computed_tokens=8,
|
||||
store_transfer_meta=TransferMeta([11], [1]),
|
||||
ready=False,
|
||||
)
|
||||
assert scheduler.has_pending_transfers() is True
|
||||
|
||||
scheduler._preempted_req_states.clear()
|
||||
scheduler._preempt_store_event_to_reqs = {"unused": []}
|
||||
scheduler._preempt_load_event_to_reqs = {1: ["req-1"]}
|
||||
scheduler._pending_hash_blocks = {"hash": MagicMock()}
|
||||
scheduler.cpu_block_pool = MagicMock()
|
||||
scheduler.cpu_block_pool.reset_prefix_cache.return_value = True
|
||||
scheduler._cleanup_preempt_cache_request = MagicMock()
|
||||
|
||||
assert scheduler.reset_cache() is True
|
||||
scheduler.cpu_block_pool.reset_prefix_cache.assert_called_once_with()
|
||||
assert scheduler._preempt_store_event_to_reqs == {}
|
||||
assert scheduler._preempt_load_event_to_reqs == {}
|
||||
assert scheduler._pending_hash_blocks == {}
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_scheduler_cleanup_preempt_load_request():
|
||||
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
|
||||
scheduler._preempt_load_event_to_reqs = {2: ["req-1"]}
|
||||
scheduler._preempted_req_states = {
|
||||
"req-1": PreemptedRequestState(
|
||||
req_id="req-1",
|
||||
cpu_block_ids=([4],),
|
||||
num_computed_tokens=8,
|
||||
store_transfer_meta=TransferMeta([1], [4]),
|
||||
load_event=2,
|
||||
load_transfer_meta=TransferMeta([10, 11], [4, 5]),
|
||||
ready=True,
|
||||
)
|
||||
}
|
||||
scheduler._gpu_block_pool = SimpleNamespace(
|
||||
blocks={10: "gpu10", 11: "gpu11"},
|
||||
free_blocks=MagicMock(),
|
||||
)
|
||||
scheduler._cleanup_preempt_cache_request = MagicMock()
|
||||
|
||||
scheduler._cleanup_preempt_load_request("req-1")
|
||||
|
||||
assert scheduler._preempt_load_event_to_reqs == {}
|
||||
freed = list(scheduler._gpu_block_pool.free_blocks.call_args.args[0])
|
||||
assert freed == ["gpu10", "gpu11"]
|
||||
scheduler._cleanup_preempt_cache_request.assert_called_once_with("req-1")
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_scheduler_cleanup_skips_null_cpu_blocks():
|
||||
scheduler = RecomputeCPUOffloadScheduler.__new__(RecomputeCPUOffloadScheduler)
|
||||
scheduler._preempted_req_states = {
|
||||
"req-1": PreemptedRequestState(
|
||||
req_id="req-1",
|
||||
cpu_block_ids=([0, 4], [0]),
|
||||
num_computed_tokens=32,
|
||||
store_transfer_meta=TransferMeta([10], [4]),
|
||||
ready=True,
|
||||
)
|
||||
}
|
||||
scheduler.cpu_block_pool = SimpleNamespace(blocks={4: "cpu4"}, free_blocks=MagicMock())
|
||||
|
||||
scheduler._cleanup_preempt_cache_request("req-1")
|
||||
|
||||
freed = list(scheduler.cpu_block_pool.free_blocks.call_args.args[0])
|
||||
assert freed == ["cpu4"]
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_worker_metadata_and_empty_transfers():
|
||||
worker = RecomputeCPUOffloadWorker.__new__(RecomputeCPUOffloadWorker)
|
||||
worker._connector_metadata = None
|
||||
worker._pending_load_event_indices = set()
|
||||
worker._submitted_load_event_indices = set()
|
||||
worker._completed_store_events = {}
|
||||
worker._load_events = []
|
||||
worker._load_hwm = -1
|
||||
worker.load_stream = None
|
||||
worker._load_stream_waited = False
|
||||
|
||||
metadata = RecomputeCPUOffloadMetadata(
|
||||
preempt_store_event=1,
|
||||
preempt_load_event=2,
|
||||
preempt_load_event_to_reqs={2: ["req-1"]},
|
||||
)
|
||||
worker.bind_connector_metadata(metadata)
|
||||
assert worker._connector_metadata is metadata
|
||||
assert worker._pending_load_event_indices == {2}
|
||||
|
||||
worker._submit_transfer([], [], 1, is_store=True)
|
||||
assert worker.build_connector_worker_meta().completed_store_events == {1: 1}
|
||||
assert worker.build_connector_worker_meta() is None
|
||||
|
||||
worker._submit_transfer([], [], 2, is_store=False)
|
||||
assert worker.get_finished(set()) == (None, {"req-1"})
|
||||
assert worker.get_finished(set()) == (None, None)
|
||||
|
||||
worker.clear_connector_metadata()
|
||||
assert worker._connector_metadata is None
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_worker_preempt_and_load_entrypoints():
|
||||
worker = RecomputeCPUOffloadWorker.__new__(RecomputeCPUOffloadWorker)
|
||||
worker._submit_transfer = MagicMock()
|
||||
worker._flush_and_sync_all = MagicMock()
|
||||
worker._connector_metadata = None
|
||||
metadata = RecomputeCPUOffloadMetadata(
|
||||
need_flush=True,
|
||||
preempt_store_event=3,
|
||||
preempt_store_gpu_blocks=[1],
|
||||
preempt_store_cpu_blocks=[2],
|
||||
preempt_load_event=4,
|
||||
preempt_load_gpu_blocks=[5],
|
||||
preempt_load_cpu_blocks=[6],
|
||||
)
|
||||
|
||||
worker.handle_preemptions(metadata)
|
||||
|
||||
worker._flush_and_sync_all.assert_called_once_with()
|
||||
worker._submit_transfer.assert_called_once_with(
|
||||
[1],
|
||||
[2],
|
||||
3,
|
||||
is_store=True,
|
||||
sync=True,
|
||||
)
|
||||
|
||||
worker._submit_transfer.reset_mock()
|
||||
worker.start_load_kv()
|
||||
worker._submit_transfer.assert_not_called()
|
||||
|
||||
worker._connector_metadata = metadata
|
||||
worker.start_load_kv()
|
||||
worker._submit_transfer.assert_called_once_with(
|
||||
[6],
|
||||
[5],
|
||||
4,
|
||||
is_store=False,
|
||||
sync=True,
|
||||
)
|
||||
|
||||
|
||||
def test_recompute_cpu_offload_worker_wait_for_layer_load_once():
|
||||
worker = RecomputeCPUOffloadWorker.__new__(RecomputeCPUOffloadWorker)
|
||||
stream = MagicMock()
|
||||
current_stream = MagicMock()
|
||||
worker.load_stream = stream
|
||||
worker._connector_metadata = RecomputeCPUOffloadMetadata(preempt_load_event=1)
|
||||
worker._load_stream_waited = False
|
||||
|
||||
with patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.recompute_cpu_offload.worker.torch.npu.current_stream",
|
||||
return_value=current_stream,
|
||||
):
|
||||
worker.wait_for_layer_load()
|
||||
worker.wait_for_layer_load()
|
||||
|
||||
current_stream.wait_stream.assert_called_once_with(stream)
|
||||
assert worker._load_stream_waited is True
|
||||
|
||||
|
||||
def test_recompute_scheduler_remote_kv_restore_keeps_exact_token_position():
|
||||
scheduler = RecomputeScheduler.__new__(RecomputeScheduler)
|
||||
scheduler.connector = MagicMock()
|
||||
scheduler.failed_recving_kv_req_ids = set()
|
||||
scheduler.finished_recving_kv_req_ids = {"req-1"}
|
||||
scheduler.kv_cache_manager = MagicMock()
|
||||
scheduler.is_mtp_kv_consumer = True
|
||||
scheduler.num_spec_tokens = 2
|
||||
scheduler.max_model_len = 32
|
||||
|
||||
request = SimpleNamespace(
|
||||
request_id="req-1",
|
||||
num_computed_tokens=9,
|
||||
num_tokens=9,
|
||||
num_preemptions=1,
|
||||
spec_token_ids=[],
|
||||
)
|
||||
|
||||
scheduler._update_waiting_for_remote_kv(request)
|
||||
|
||||
scheduler.kv_cache_manager.cache_blocks.assert_called_once_with(request, 8)
|
||||
assert request.num_computed_tokens == 8
|
||||
assert request.spec_token_ids == [PLACEHOLDER_TOKEN_ID] * 2
|
||||
assert scheduler.finished_recving_kv_req_ids == set()
|
||||
|
||||
|
||||
def test_recompute_scheduler_remote_kv_restore_frees_failed_empty_load():
|
||||
scheduler = RecomputeScheduler.__new__(RecomputeScheduler)
|
||||
scheduler.connector = MagicMock()
|
||||
scheduler.failed_recving_kv_req_ids = {"req-1"}
|
||||
scheduler.finished_recving_kv_req_ids = {"req-1"}
|
||||
scheduler.kv_cache_manager = MagicMock()
|
||||
|
||||
request = SimpleNamespace(
|
||||
request_id="req-1",
|
||||
num_computed_tokens=0,
|
||||
)
|
||||
|
||||
scheduler._update_waiting_for_remote_kv(request)
|
||||
|
||||
scheduler.kv_cache_manager.free.assert_called_once_with(request)
|
||||
scheduler.kv_cache_manager.cache_blocks.assert_not_called()
|
||||
assert scheduler.failed_recving_kv_req_ids == set()
|
||||
assert scheduler.finished_recving_kv_req_ids == set()
|
||||
200
tests/ut/kv_offload/utils.py
Normal file
200
tests/ut/kv_offload/utils.py
Normal file
@@ -0,0 +1,200 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# This code is from: https://github.com/vllm-project/vllm/tests/v1/kv_connector/unit/utils.py
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from vllm import SamplingParams
|
||||
from vllm.config import CacheConfig, DeviceConfig, KVTransferConfig, ModelConfig, SchedulerConfig, VllmConfig
|
||||
from vllm.utils.hashing import sha256
|
||||
from vllm.v1.core.kv_cache_utils import get_request_block_hasher, init_none_hash
|
||||
from vllm.v1.core.sched.scheduler import Scheduler
|
||||
from vllm.v1.kv_cache_interface import FullAttentionSpec, KVCacheConfig, KVCacheGroupSpec
|
||||
from vllm.v1.outputs import KVConnectorOutput, ModelRunnerOutput
|
||||
from vllm.v1.request import Request
|
||||
from vllm.v1.structured_output import StructuredOutputManager
|
||||
|
||||
EOS_TOKEN_ID = 50256
|
||||
|
||||
|
||||
def assert_scheduler_empty(scheduler: Scheduler):
|
||||
"""Confirm the scheduler is "empty" - i.e. no leaks."""
|
||||
# Scheduler Metadata.
|
||||
assert len(scheduler.requests) == 0
|
||||
assert len(scheduler.waiting) == 0
|
||||
assert len(scheduler.running) == 0
|
||||
assert len(scheduler.finished_req_ids) == 0
|
||||
assert len(scheduler.finished_recving_kv_req_ids) == 0
|
||||
|
||||
# EncoderCacheManager.
|
||||
assert len(scheduler.encoder_cache_manager.freed) == 0
|
||||
assert len(scheduler.encoder_cache_manager.cached) == 0
|
||||
|
||||
# KVCache Manager.
|
||||
assert len(scheduler.kv_cache_manager.coordinator.single_type_managers[0].req_to_blocks) == 0
|
||||
assert len(scheduler.kv_cache_manager.coordinator.single_type_managers[0].num_cached_block) == 0
|
||||
num_free_blocks = scheduler.kv_cache_manager.block_pool.free_block_queue.num_free_blocks
|
||||
assert num_free_blocks == (scheduler.kv_cache_manager.block_pool.num_gpu_blocks - 1)
|
||||
|
||||
for block in scheduler.kv_cache_manager.block_pool.blocks:
|
||||
assert block.ref_cnt == 0
|
||||
|
||||
|
||||
def create_vllm_config(
|
||||
max_num_seqs: int = 16,
|
||||
max_num_batched_tokens: int = 1024,
|
||||
block_size: int = 128,
|
||||
) -> VllmConfig:
|
||||
"""Initialize VllmConfig For Testing."""
|
||||
fake_weight_path = os.path.join(os.path.dirname(__file__), "..", "_fake_weight")
|
||||
model_config = ModelConfig(
|
||||
model=fake_weight_path,
|
||||
skip_tokenizer_init=True,
|
||||
)
|
||||
scheduler_config = SchedulerConfig(
|
||||
max_num_seqs=max_num_seqs,
|
||||
max_num_batched_tokens=max_num_batched_tokens,
|
||||
max_model_len=max_num_batched_tokens,
|
||||
enable_chunked_prefill=True,
|
||||
is_encoder_decoder=model_config.is_encoder_decoder,
|
||||
)
|
||||
cache_config = CacheConfig(
|
||||
block_size=block_size,
|
||||
gpu_memory_utilization=0.9,
|
||||
cache_dtype="auto",
|
||||
enable_prefix_caching=True,
|
||||
)
|
||||
kv_transfer_config = KVTransferConfig(kv_connector="MooncakeConnector", kv_role="kv_both")
|
||||
return VllmConfig(
|
||||
scheduler_config=scheduler_config,
|
||||
model_config=model_config,
|
||||
cache_config=cache_config,
|
||||
kv_transfer_config=kv_transfer_config,
|
||||
device_config=DeviceConfig("cpu"),
|
||||
)
|
||||
|
||||
|
||||
def create_scheduler(
|
||||
vllm_config: VllmConfig,
|
||||
num_blocks: int = 10000,
|
||||
) -> Scheduler:
|
||||
"""Initialize Scheduler For Testing."""
|
||||
block_size = vllm_config.cache_config.block_size
|
||||
kv_cache_config = KVCacheConfig(
|
||||
num_blocks=num_blocks,
|
||||
kv_cache_tensors=[],
|
||||
kv_cache_groups=[
|
||||
KVCacheGroupSpec(
|
||||
["layer"], FullAttentionSpec(block_size=block_size, num_kv_heads=1, head_size=1, dtype=torch.float16)
|
||||
)
|
||||
],
|
||||
)
|
||||
vllm_config.cache_config.num_gpu_blocks = num_blocks
|
||||
|
||||
return Scheduler(
|
||||
vllm_config=vllm_config,
|
||||
kv_cache_config=kv_cache_config,
|
||||
log_stats=True,
|
||||
block_size=block_size,
|
||||
structured_output_manager=StructuredOutputManager(vllm_config),
|
||||
)
|
||||
|
||||
|
||||
_none_hash_initialized = False
|
||||
|
||||
|
||||
def create_request(
|
||||
request_id: int,
|
||||
num_tokens: int = 10,
|
||||
max_tokens: int = 128,
|
||||
do_remote_decode: bool = False,
|
||||
do_remote_prefill: bool = False,
|
||||
num_remote_blocks: int = 3,
|
||||
block_size: int = 16,
|
||||
) -> Request:
|
||||
"""Make dummy request for testing."""
|
||||
global _none_hash_initialized
|
||||
if not _none_hash_initialized:
|
||||
init_none_hash(sha256)
|
||||
_none_hash_initialized = True
|
||||
|
||||
kv_transfer_params: dict[str, Any] | None = None
|
||||
|
||||
if do_remote_decode:
|
||||
assert not do_remote_prefill
|
||||
kv_transfer_params = dict(
|
||||
do_remote_prefill=False,
|
||||
do_remote_decode=True,
|
||||
transfer_id=f"transfer-{request_id}",
|
||||
)
|
||||
elif do_remote_prefill:
|
||||
kv_transfer_params = dict(
|
||||
do_remote_prefill=True,
|
||||
do_remote_decode=False,
|
||||
remote_engine_id="my-engine-id",
|
||||
remote_block_ids=list(range(num_remote_blocks)),
|
||||
remote_host="my-host",
|
||||
remote_port=1234,
|
||||
remote_bootstrap_addr="my-bootstrap",
|
||||
transfer_id=f"transfer-{request_id}",
|
||||
remote_tp_size=1,
|
||||
remote_pcp_size=1,
|
||||
remote_dcp_size=1,
|
||||
)
|
||||
|
||||
max_tokens = 1 if do_remote_decode else max_tokens
|
||||
sampling_params = SamplingParams(max_tokens=max_tokens)
|
||||
sampling_params.update_from_generation_config({}, EOS_TOKEN_ID)
|
||||
|
||||
prompt_token_ids = [i * request_id for i in range(num_tokens)]
|
||||
|
||||
block_hasher = get_request_block_hasher(block_size, sha256)
|
||||
|
||||
req = Request(
|
||||
request_id=f"id-{request_id}",
|
||||
prompt_token_ids=prompt_token_ids,
|
||||
sampling_params=sampling_params,
|
||||
pooling_params=None,
|
||||
block_hasher=block_hasher,
|
||||
)
|
||||
req.kv_transfer_params = kv_transfer_params
|
||||
return req
|
||||
|
||||
|
||||
def create_model_runner_output(
|
||||
reqs: list[Request],
|
||||
finished_sending: set[str] | None = None,
|
||||
finished_recving: set[str] | None = None,
|
||||
use_eos: bool = False,
|
||||
) -> ModelRunnerOutput:
|
||||
"""Make dummy model runner output for testing."""
|
||||
|
||||
req_ids = [req.request_id for req in reqs]
|
||||
req_id_to_index = {req_id: idx for idx, req_id in enumerate(req_ids)}
|
||||
|
||||
sampled_token = EOS_TOKEN_ID if use_eos else 0
|
||||
sampled_token_ids = [[sampled_token] for _ in req_ids]
|
||||
|
||||
kv_connector_output = (
|
||||
None
|
||||
if (finished_sending is None and finished_recving is None)
|
||||
else KVConnectorOutput(
|
||||
finished_sending=finished_sending,
|
||||
finished_recving=finished_recving,
|
||||
)
|
||||
)
|
||||
|
||||
model_runner_output = ModelRunnerOutput(
|
||||
req_ids=req_ids,
|
||||
req_id_to_index=req_id_to_index,
|
||||
sampled_token_ids=sampled_token_ids,
|
||||
logprobs=None,
|
||||
prompt_logprobs_dict={},
|
||||
pooler_output=None,
|
||||
kv_connector_output=kv_connector_output,
|
||||
)
|
||||
|
||||
return model_runner_output
|
||||
Reference in New Issue
Block a user