73
tests/ut/distributed/mooncake/test_mooncake_kv_transfer.py
Normal file
73
tests/ut/distributed/mooncake/test_mooncake_kv_transfer.py
Normal file
@@ -0,0 +1,73 @@
|
||||
import threading
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
|
||||
if not hasattr(torch, "npu"):
|
||||
torch.npu = SimpleNamespace(Event=object) # type: ignore[attr-defined]
|
||||
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.config_data import (
|
||||
ChunkedTokenDatabase,
|
||||
KeyMetadata,
|
||||
ReqMeta,
|
||||
)
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer import (
|
||||
KVCacheStoreSendingThread,
|
||||
)
|
||||
|
||||
|
||||
class _FakeStore:
|
||||
def __init__(self, exists_result: list[int]):
|
||||
self.exists_result = exists_result
|
||||
self.put_calls: list[tuple[list[str], list[list[int]], list[list[int]]]] = []
|
||||
|
||||
def set_device(self):
|
||||
return None
|
||||
|
||||
def exists(self, keys: list[str]) -> list[int]:
|
||||
# Return exact number of states for requested keys.
|
||||
return self.exists_result[: len(keys)]
|
||||
|
||||
def put(self, keys, addrs, sizes):
|
||||
self.put_calls.append((list(keys), list(addrs), list(sizes)))
|
||||
|
||||
|
||||
class TestKVTransferMissingKeyPut(unittest.TestCase):
|
||||
def test_sending_thread_only_puts_missing_keys(self):
|
||||
store = _FakeStore(exists_result=[1, 0, 1, 0])
|
||||
token_db = ChunkedTokenDatabase([KeyMetadata("m", 0, 0, 0, 0)], [16], None)
|
||||
token_db.set_group_buffers({0: [1000]}, {0: [16]}, {0: [1]})
|
||||
thread = KVCacheStoreSendingThread(
|
||||
m_store=store,
|
||||
token_database=token_db,
|
||||
block_size=16,
|
||||
tp_rank=0,
|
||||
dcp_size=1,
|
||||
put_step=1,
|
||||
kv_role="kv_producer",
|
||||
ready_event=threading.Event(),
|
||||
group_uses_align_state=[False],
|
||||
enable_kv_event=False,
|
||||
)
|
||||
|
||||
req_meta = ReqMeta(
|
||||
req_id="req-1",
|
||||
token_len_chunk=64,
|
||||
block_ids=[0, 1, 2, 3],
|
||||
block_hashes=[b"h0", b"h1", b"h2", b"h3"], # type: ignore[arg-type]
|
||||
current_event=None,
|
||||
)
|
||||
thread.add_stored_request("req-1")
|
||||
thread.request_queue.put(req_meta)
|
||||
thread._handle_request(req_meta)
|
||||
|
||||
self.assertEqual(len(store.put_calls), 1)
|
||||
put_keys, put_addrs, put_sizes = store.put_calls[0]
|
||||
self.assertEqual(len(put_keys), 2)
|
||||
self.assertEqual(put_addrs, [[1001], [1003]])
|
||||
self.assertEqual(put_sizes, [[16], [16]])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user