0
tests/ut/distributed/__init__.py
Normal file
0
tests/ut/distributed/__init__.py
Normal file
0
tests/ut/distributed/ascend_store/__init__.py
Normal file
0
tests/ut/distributed/ascend_store/__init__.py
Normal file
463
tests/ut/distributed/ascend_store/_mock_deps.py
Normal file
463
tests/ut/distributed/ascend_store/_mock_deps.py
Normal file
@@ -0,0 +1,463 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
"""Mock heavy dependencies (torch, vllm, etc.) for ascend_store unit tests.
|
||||
|
||||
IMPORTANT: This module MUST be imported before any vllm_ascend or vllm
|
||||
imports in each test file.
|
||||
|
||||
Usage at the top of each test file:
|
||||
import tests.ut.distributed.ascend_store._mock_deps # noqa: F401, E402
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
import types
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Mock torch / torch_npu
|
||||
# ---------------------------------------------------------------------------
|
||||
if "torch" not in sys.modules and importlib.util.find_spec("torch") is None:
|
||||
_torch = types.ModuleType("torch")
|
||||
_torch.Tensor = MagicMock # type: ignore[attr-defined]
|
||||
_torch.bool = "bool" # type: ignore[attr-defined]
|
||||
_torch.float16 = "float16" # type: ignore[attr-defined]
|
||||
_torch.float32 = "float32" # type: ignore[attr-defined]
|
||||
_torch.zeros = MagicMock(return_value=MagicMock()) # type: ignore[attr-defined]
|
||||
_torch.sum = MagicMock(return_value=0) # type: ignore[attr-defined]
|
||||
_torch.device = MagicMock() # type: ignore[attr-defined]
|
||||
_torch.distributed = MagicMock() # type: ignore[attr-defined]
|
||||
_npu = MagicMock()
|
||||
_npu.Event = MagicMock
|
||||
_npu.current_device = MagicMock(return_value=0)
|
||||
_npu.set_device = MagicMock()
|
||||
_torch.npu = _npu # type: ignore[attr-defined]
|
||||
sys.modules["torch"] = _torch
|
||||
sys.modules["torch.distributed"] = _torch.distributed # type: ignore[attr-defined]
|
||||
|
||||
if "torch_npu" not in sys.modules:
|
||||
sys.modules["torch_npu"] = MagicMock()
|
||||
sys.modules["torch_npu._inductor"] = MagicMock()
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Mock vllm modules
|
||||
# ---------------------------------------------------------------------------
|
||||
_MOCK_VLLM_DEPS = importlib.util.find_spec("vllm") is None
|
||||
_vllm_mock_modules = [
|
||||
"vllm",
|
||||
"vllm.config",
|
||||
"vllm.distributed",
|
||||
"vllm.distributed.kv_events",
|
||||
"vllm.distributed.kv_transfer",
|
||||
"vllm.distributed.kv_transfer.kv_connector",
|
||||
"vllm.distributed.kv_transfer.kv_connector.factory",
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1",
|
||||
"vllm.distributed.kv_transfer.kv_connector.v1.base",
|
||||
"vllm.distributed.parallel_state",
|
||||
"vllm.envs",
|
||||
"vllm.forward_context",
|
||||
"vllm.logger",
|
||||
"vllm.model_executor",
|
||||
"vllm.model_executor.layers",
|
||||
"vllm.model_executor.layers.linear",
|
||||
"vllm.model_executor.layers.quantization",
|
||||
"vllm.platforms",
|
||||
"vllm.utils",
|
||||
"vllm.utils.hashing",
|
||||
"vllm.utils.math_utils",
|
||||
"vllm.utils.network_utils",
|
||||
"vllm.v1",
|
||||
"vllm.v1.attention",
|
||||
"vllm.v1.attention.backend",
|
||||
"vllm.v1.core",
|
||||
"vllm.v1.core.block_pool",
|
||||
"vllm.v1.core.kv_cache_manager",
|
||||
"vllm.v1.core.kv_cache_utils",
|
||||
"vllm.v1.core.sched",
|
||||
"vllm.v1.core.sched.output",
|
||||
"vllm.v1.core.single_type_kv_cache_manager",
|
||||
"vllm.v1.kv_cache_interface",
|
||||
"vllm.v1.kv_cache_spec_registry",
|
||||
"vllm.v1.outputs",
|
||||
"vllm.v1.request",
|
||||
"vllm.v1.serial_utils",
|
||||
]
|
||||
if _MOCK_VLLM_DEPS:
|
||||
for _mod_name in _vllm_mock_modules:
|
||||
if _mod_name not in sys.modules:
|
||||
sys.modules[_mod_name] = MagicMock()
|
||||
|
||||
if _MOCK_VLLM_DEPS:
|
||||
sys.modules["vllm.utils.math_utils"].cdiv = lambda a, b: -(-a // b) # type: ignore[attr-defined]
|
||||
sys.modules["vllm.logger"].logger = logging.getLogger("vllm") # type: ignore[attr-defined]
|
||||
|
||||
_base_mod: Any = (
|
||||
sys.modules["vllm.distributed.kv_transfer.kv_connector.v1.base"] if _MOCK_VLLM_DEPS else types.SimpleNamespace()
|
||||
)
|
||||
_base_mod.KVConnectorBase_V1 = type("KVConnectorBase_V1", (), {"__init__": lambda self, **kw: None}) # type: ignore[attr-defined]
|
||||
_base_mod.KVConnectorMetadata = type("KVConnectorMetadata", (), {}) # type: ignore[attr-defined]
|
||||
_base_mod.KVConnectorWorkerMetadata = type("KVConnectorWorkerMetadata", (), {}) # type: ignore[attr-defined]
|
||||
_base_mod.KVConnectorRole = MagicMock() # type: ignore[attr-defined]
|
||||
_base_mod.KVConnectorRole.SCHEDULER = "SCHEDULER"
|
||||
_base_mod.KVConnectorRole.WORKER = "WORKER"
|
||||
_base_mod.SupportsHMA = type("SupportsHMA", (), {}) # type: ignore[attr-defined]
|
||||
|
||||
_events_mod: Any = sys.modules["vllm.distributed.kv_events"] if _MOCK_VLLM_DEPS else types.SimpleNamespace()
|
||||
_events_mod.KVCacheEvent = type("KVCacheEvent", (), {}) # type: ignore[attr-defined]
|
||||
_events_mod.KVConnectorKVEvents = type("KVConnectorKVEvents", (), {}) # type: ignore[attr-defined]
|
||||
|
||||
|
||||
class _FakeAggregator:
|
||||
def __init__(self, *args, **kwargs):
|
||||
self._mock = MagicMock()
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._mock, name)
|
||||
|
||||
|
||||
_events_mod.KVEventAggregator = _FakeAggregator # type: ignore[attr-defined]
|
||||
_events_mod.BlockStored = type( # type: ignore[attr-defined]
|
||||
"BlockStored",
|
||||
(),
|
||||
{"__init__": lambda self, **kwargs: self.__dict__.update(kwargs)},
|
||||
)
|
||||
|
||||
_kv_cache_utils_mod: Any = sys.modules["vllm.v1.core.kv_cache_utils"] if _MOCK_VLLM_DEPS else types.SimpleNamespace()
|
||||
_kv_cache_utils_mod.BlockHash = bytes # type: ignore[attr-defined]
|
||||
_kv_cache_utils_mod.maybe_convert_block_hash = lambda x: x # type: ignore[attr-defined]
|
||||
|
||||
|
||||
class _FakeKVCacheBlock:
|
||||
def __init__(self, block_id=0, **kwargs):
|
||||
self.block_id = block_id
|
||||
self.__dict__.update(kwargs)
|
||||
|
||||
|
||||
class _FakeKVCacheSpec:
|
||||
def __init__(self, block_size=16, **kwargs):
|
||||
self.block_size = block_size
|
||||
for key, value in kwargs.items():
|
||||
setattr(self, key, value)
|
||||
|
||||
def __eq__(self, other):
|
||||
return type(self) is type(other) and self.__dict__ == getattr(other, "__dict__", {})
|
||||
|
||||
def copy_with_new_block_size(self, block_size):
|
||||
kwargs = self.__dict__.copy()
|
||||
kwargs["block_size"] = block_size
|
||||
return type(self)(**kwargs)
|
||||
|
||||
@property
|
||||
def page_size_bytes(self):
|
||||
num_kv_heads = getattr(self, "num_kv_heads", 1)
|
||||
head_size = getattr(self, "head_size", 1)
|
||||
dtype = getattr(self, "dtype", None)
|
||||
dtype_size = getattr(dtype, "itemsize", None)
|
||||
if dtype_size is None and dtype is not None and hasattr(dtype, "element_size"):
|
||||
dtype_size = dtype.element_size()
|
||||
return self.block_size * num_kv_heads * head_size * int(dtype_size or 1) * 2
|
||||
|
||||
|
||||
class _FakeFullAttentionSpec(_FakeKVCacheSpec):
|
||||
pass
|
||||
|
||||
|
||||
class _FakeSlidingWindowSpec(_FakeKVCacheSpec):
|
||||
def __init__(self, block_size=16, sliding_window=32, **kwargs):
|
||||
super().__init__(block_size=block_size, sliding_window=sliding_window, **kwargs)
|
||||
|
||||
|
||||
class _FakeMambaSpec(_FakeKVCacheSpec):
|
||||
def __init__(self, block_size=16, **kwargs):
|
||||
super().__init__(block_size=block_size, **kwargs)
|
||||
self.num_speculative_blocks = getattr(self, "num_speculative_blocks", 0)
|
||||
|
||||
|
||||
class _FakeUniformTypeKVCacheSpecs(_FakeKVCacheSpec):
|
||||
def __init__(self, block_size=16, kv_cache_specs=None, **kwargs):
|
||||
super().__init__(block_size=block_size, **kwargs)
|
||||
self.kv_cache_specs = kv_cache_specs or {}
|
||||
|
||||
@classmethod
|
||||
def from_specs(cls, kv_cache_specs):
|
||||
if not kv_cache_specs:
|
||||
return None
|
||||
first_spec = next(iter(kv_cache_specs.values()))
|
||||
return cls(
|
||||
block_size=getattr(first_spec, "block_size", 16),
|
||||
kv_cache_specs=kv_cache_specs,
|
||||
)
|
||||
|
||||
|
||||
class _FakeKVCacheGroupSpec:
|
||||
def __init__(self, layer_names=None, kv_cache_spec=None, is_eagle_group=False):
|
||||
self.layer_names = layer_names or []
|
||||
self.kv_cache_spec = kv_cache_spec or _FakeFullAttentionSpec()
|
||||
self.is_eagle_group = is_eagle_group
|
||||
|
||||
|
||||
class _FakeKVCacheConfig:
|
||||
def __init__(self, num_blocks=1, kv_cache_tensors=None, kv_cache_groups=None):
|
||||
self.num_blocks = num_blocks
|
||||
self.kv_cache_tensors = kv_cache_tensors or []
|
||||
self.kv_cache_groups = kv_cache_groups or []
|
||||
|
||||
|
||||
_kv_cache_utils_mod.KVCacheBlock = _FakeKVCacheBlock # type: ignore[attr-defined]
|
||||
_kv_cache_utils_mod.BlockHashList = list # type: ignore[attr-defined]
|
||||
|
||||
|
||||
class _FakeBlockPool:
|
||||
def __init__(self, *args, **kwargs):
|
||||
self.null_block = _FakeKVCacheBlock(block_id=0)
|
||||
self._next_block_id = 1
|
||||
|
||||
def get_new_blocks(self, num_blocks):
|
||||
blocks = []
|
||||
for _ in range(num_blocks):
|
||||
blocks.append(_FakeKVCacheBlock(block_id=self._next_block_id))
|
||||
self._next_block_id += 1
|
||||
return blocks
|
||||
|
||||
|
||||
if _MOCK_VLLM_DEPS:
|
||||
sys.modules["vllm.v1.core.block_pool"].BlockPool = _FakeBlockPool # type: ignore[attr-defined]
|
||||
|
||||
|
||||
class _FakeSingleTypeKVCacheManager:
|
||||
def __init__(self, *args, **kwargs):
|
||||
self._mock = MagicMock()
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._mock, name)
|
||||
|
||||
@classmethod
|
||||
def reachable_block_mask(
|
||||
cls,
|
||||
start_block,
|
||||
end_block,
|
||||
alignment_tokens,
|
||||
kv_cache_spec,
|
||||
use_eagle,
|
||||
retention_interval=None,
|
||||
num_prompt_tokens=None,
|
||||
):
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def find_longest_cache_hit(
|
||||
cls,
|
||||
block_hashes,
|
||||
max_length,
|
||||
kv_cache_group_ids,
|
||||
block_pool,
|
||||
kv_cache_spec,
|
||||
drop_eagle_block=False,
|
||||
alignment_tokens=16,
|
||||
dcp_world_size=1,
|
||||
pcp_world_size=1,
|
||||
):
|
||||
computed: tuple[list[object], ...] = tuple([] for _ in kv_cache_group_ids)
|
||||
max_blocks = max_length // kv_cache_spec.block_size
|
||||
for block_hash in list(block_hashes)[:max_blocks]:
|
||||
cached = block_pool.get_cached_block(block_hash, kv_cache_group_ids)
|
||||
if not cached:
|
||||
break
|
||||
for blocks, block in zip(computed, cached):
|
||||
blocks.append(block)
|
||||
if drop_eagle_block and computed and computed[0]:
|
||||
for blocks in computed:
|
||||
blocks.pop()
|
||||
return computed
|
||||
|
||||
|
||||
class _FakeSlidingWindowManager(_FakeSingleTypeKVCacheManager):
|
||||
@classmethod
|
||||
def reachable_block_mask(
|
||||
cls,
|
||||
start_block,
|
||||
end_block,
|
||||
alignment_tokens,
|
||||
kv_cache_spec,
|
||||
use_eagle,
|
||||
retention_interval=None,
|
||||
num_prompt_tokens=None,
|
||||
):
|
||||
if alignment_tokens is None:
|
||||
return None
|
||||
per_segment = max(alignment_tokens // kv_cache_spec.block_size, 1)
|
||||
return [(idx + 1) % per_segment == 0 for idx in range(start_block, end_block)]
|
||||
|
||||
|
||||
_single_type_mod: Any = (
|
||||
sys.modules["vllm.v1.core.single_type_kv_cache_manager"] if _MOCK_VLLM_DEPS else types.SimpleNamespace()
|
||||
)
|
||||
_single_type_mod.SingleTypeKVCacheManager = _FakeSingleTypeKVCacheManager # type: ignore[attr-defined]
|
||||
_single_type_mod.FullAttentionManager = _FakeSingleTypeKVCacheManager # type: ignore[attr-defined]
|
||||
_single_type_mod.SlidingWindowManager = _FakeSlidingWindowManager # type: ignore[attr-defined]
|
||||
_single_type_mod.MambaManager = _FakeSingleTypeKVCacheManager # type: ignore[attr-defined]
|
||||
_single_type_mod.spec_manager_map = { # type: ignore[attr-defined]
|
||||
_FakeFullAttentionSpec: _FakeSingleTypeKVCacheManager,
|
||||
_FakeSlidingWindowSpec: _FakeSlidingWindowManager,
|
||||
_FakeMambaSpec: _FakeSingleTypeKVCacheManager,
|
||||
}
|
||||
|
||||
_kv_interface_mod: Any = sys.modules["vllm.v1.kv_cache_interface"] if _MOCK_VLLM_DEPS else types.SimpleNamespace()
|
||||
_kv_interface_mod.KVCacheSpec = _FakeKVCacheSpec # type: ignore[attr-defined]
|
||||
_kv_interface_mod.FullAttentionSpec = _FakeFullAttentionSpec # type: ignore[attr-defined]
|
||||
_kv_interface_mod.SlidingWindowSpec = _FakeSlidingWindowSpec # type: ignore[attr-defined]
|
||||
_kv_interface_mod.MambaSpec = _FakeMambaSpec # type: ignore[attr-defined]
|
||||
_kv_interface_mod.UniformTypeKVCacheSpecs = _FakeUniformTypeKVCacheSpecs # type: ignore[attr-defined]
|
||||
_kv_interface_mod.KVCacheGroupSpec = _FakeKVCacheGroupSpec # type: ignore[attr-defined]
|
||||
_kv_interface_mod.KVCacheConfig = _FakeKVCacheConfig # type: ignore[attr-defined]
|
||||
|
||||
|
||||
class _FakeKVCacheSpecRegistry:
|
||||
@classmethod
|
||||
def get_manager_class(cls, kv_cache_spec):
|
||||
if isinstance(kv_cache_spec, _FakeSlidingWindowSpec):
|
||||
return _FakeSlidingWindowManager
|
||||
return _FakeSingleTypeKVCacheManager
|
||||
|
||||
|
||||
if _MOCK_VLLM_DEPS:
|
||||
sys.modules["vllm.v1.kv_cache_spec_registry"].KVCacheSpecRegistry = _FakeKVCacheSpecRegistry # type: ignore[attr-defined]
|
||||
|
||||
_sched_output_mod: Any = sys.modules["vllm.v1.core.sched.output"] if _MOCK_VLLM_DEPS else types.SimpleNamespace()
|
||||
_sched_output_mod.NewRequestData = MagicMock # type: ignore[attr-defined]
|
||||
|
||||
if _MOCK_VLLM_DEPS:
|
||||
sys.modules["vllm.envs"].VLLM_RPC_BASE_PATH = "/tmp/vllm_rpc" # type: ignore[attr-defined]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Mock external backends
|
||||
# ---------------------------------------------------------------------------
|
||||
for _mod_name in [
|
||||
"mooncake",
|
||||
"mooncake.engine",
|
||||
"mooncake.store",
|
||||
"memcache_hybrid",
|
||||
"yr",
|
||||
"yr.datasystem",
|
||||
"yr.datasystem.hetero_client",
|
||||
"yr.datasystem.kv_client",
|
||||
"yr.datasystem.object_client",
|
||||
"zmq",
|
||||
]:
|
||||
if _mod_name not in sys.modules:
|
||||
sys.modules[_mod_name] = MagicMock()
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Mock vllm_ascend transitive imports
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_pkg(name, path=""):
|
||||
mod = types.ModuleType(name)
|
||||
mod.__path__ = [path] # type: ignore[attr-defined]
|
||||
mod.__package__ = name # type: ignore[attr-defined]
|
||||
return mod
|
||||
|
||||
|
||||
_vllm_ascend_real_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "..", "..", "vllm_ascend"))
|
||||
_vllm_ascend_package_paths = {
|
||||
"vllm_ascend": _vllm_ascend_real_path,
|
||||
"vllm_ascend.distributed": os.path.join(_vllm_ascend_real_path, "distributed"),
|
||||
}
|
||||
for _pkg, _path in _vllm_ascend_package_paths.items():
|
||||
if _pkg not in sys.modules:
|
||||
sys.modules[_pkg] = _make_pkg(_pkg, _path)
|
||||
|
||||
_distributed_utils = types.ModuleType("vllm_ascend.distributed.utils")
|
||||
_distributed_utils.get_decode_context_model_parallel_rank = MagicMock( # type: ignore[attr-defined]
|
||||
return_value=0
|
||||
)
|
||||
_distributed_utils.get_decode_context_model_parallel_world_size = MagicMock( # type: ignore[attr-defined]
|
||||
return_value=1
|
||||
)
|
||||
sys.modules["vllm_ascend.distributed.utils"] = _distributed_utils
|
||||
|
||||
_kv_transfer_init = _make_pkg("vllm_ascend.distributed.kv_transfer")
|
||||
_kv_transfer_init.register_connector = MagicMock() # type: ignore[attr-defined]
|
||||
sys.modules["vllm_ascend.distributed.kv_transfer"] = _kv_transfer_init
|
||||
|
||||
_kv_utils_pkg = _make_pkg("vllm_ascend.distributed.kv_transfer.utils")
|
||||
sys.modules["vllm_ascend.distributed.kv_transfer.utils"] = _kv_utils_pkg
|
||||
sys.modules["vllm_ascend.distributed.kv_transfer.utils.mooncake_transfer_engine"] = MagicMock()
|
||||
|
||||
_kv_pool_pkg = _make_pkg("vllm_ascend.distributed.kv_transfer.kv_pool")
|
||||
sys.modules["vllm_ascend.distributed.kv_transfer.kv_pool"] = _kv_pool_pkg
|
||||
|
||||
_ascend_store_real_path = os.path.join(
|
||||
os.path.dirname(__file__),
|
||||
"..",
|
||||
"..",
|
||||
"..",
|
||||
"..",
|
||||
"vllm_ascend",
|
||||
"distributed",
|
||||
"kv_transfer",
|
||||
"kv_pool",
|
||||
"ascend_store",
|
||||
)
|
||||
_ascend_store_pkg = _make_pkg(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store",
|
||||
os.path.abspath(_ascend_store_real_path),
|
||||
)
|
||||
sys.modules["vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store"] = _ascend_store_pkg
|
||||
|
||||
_backend_pkg = _make_pkg(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend",
|
||||
os.path.join(os.path.abspath(_ascend_store_real_path), "backend"),
|
||||
)
|
||||
sys.modules["vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend"] = _backend_pkg
|
||||
|
||||
# Mirror the real backend/__init__.py entry points. The scheduler/worker resolve
|
||||
# the backend class dynamically via ``importlib.import_module(path)``; tests that
|
||||
# exercise those paths patch ``<module>.importlib`` locally (see
|
||||
# test_pool_scheduler.py / test_pool_worker.py) so the backend resolves to a
|
||||
# MagicMock. Do NOT register the backends in sys.modules or globally wrap
|
||||
# import_module here: test_backend.py imports the real backend classes and also
|
||||
# relies on ``mock.patch`` (which itself calls importlib.import_module) resolving
|
||||
# those real modules.
|
||||
_backend_module_paths = {
|
||||
"mooncake": "vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend",
|
||||
"memcache": "vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.memcache_backend",
|
||||
"yuanrong": "vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.yuanrong_backend",
|
||||
}
|
||||
_backend_pkg.backend_map = { # type: ignore[attr-defined]
|
||||
"mooncake": {"name": "MooncakeBackend", "path": _backend_module_paths["mooncake"]},
|
||||
"memcache": {"name": "MemcacheBackend", "path": _backend_module_paths["memcache"]},
|
||||
"yuanrong": {"name": "YuanrongBackend", "path": _backend_module_paths["yuanrong"]},
|
||||
}
|
||||
|
||||
if "vllm_ascend.utils" not in sys.modules or not hasattr(sys.modules["vllm_ascend.utils"], "AscendDeviceType"):
|
||||
_ascend_utils = MagicMock()
|
||||
_ascend_utils.AscendDeviceType = MagicMock()
|
||||
_ascend_utils.get_ascend_device_type = MagicMock()
|
||||
sys.modules["vllm_ascend.utils"] = _ascend_utils
|
||||
|
||||
# NOTE: vllm_ascend.{ascend_config, memcache_comm_fence} and their helpers
|
||||
# (get_ascend_config, AttentionComputeStartGate, ...) are intentionally NOT
|
||||
# mocked here. Doing so by mutating these real modules leaks into every other
|
||||
# UT in the same pytest session (breaking test_ascend_config / test_platform,
|
||||
# which collect after ascend_store and bind the polluted symbols at import).
|
||||
# These helpers are mocked per-test, scoped to the ascend_store tests only,
|
||||
# via the autouse fixture in tests/ut/conftest.py.
|
||||
459
tests/ut/distributed/ascend_store/test_ascend_store_connector.py
Normal file
459
tests/ut/distributed/ascend_store/test_ascend_store_connector.py
Normal file
@@ -0,0 +1,459 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import types
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
# isort: off
|
||||
import tests.ut.distributed.ascend_store._mock_deps # noqa: F401, E402
|
||||
from vllm.distributed.kv_events import KVCacheEvent
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector import (
|
||||
AscendStoreConnector,
|
||||
AscendStoreKVEvents,
|
||||
)
|
||||
|
||||
# isort: on
|
||||
|
||||
|
||||
def _mock_events(num_workers=1):
|
||||
events = AscendStoreKVEvents(num_workers=num_workers)
|
||||
events._aggregator = MagicMock()
|
||||
return events
|
||||
|
||||
|
||||
class TestAscendStoreKVEvents(unittest.TestCase):
|
||||
def _make_events(self, num_workers=1):
|
||||
return _mock_events(num_workers=num_workers)
|
||||
|
||||
def test_add_and_get_events(self):
|
||||
ev = self._make_events()
|
||||
mock_events = [MagicMock(spec=KVCacheEvent), MagicMock(spec=KVCacheEvent)]
|
||||
ev.add_events(mock_events)
|
||||
ev._aggregator.get_all_events.return_value = mock_events
|
||||
result = ev.get_all_events()
|
||||
self.assertEqual(result, mock_events)
|
||||
|
||||
def test_aggregate(self):
|
||||
ev = self._make_events()
|
||||
common = [MagicMock()]
|
||||
ev._aggregator.get_common_events.return_value = common
|
||||
result = ev.aggregate()
|
||||
self.assertIs(result, ev)
|
||||
ev._aggregator.clear_events.assert_called()
|
||||
ev._aggregator.add_events.assert_called_with(common)
|
||||
ev._aggregator.reset_workers.assert_called()
|
||||
|
||||
def test_increment_workers(self):
|
||||
ev = self._make_events()
|
||||
ev.increment_workers(3)
|
||||
ev._aggregator.increment_workers.assert_called_with(3)
|
||||
|
||||
def test_get_number_of_workers(self):
|
||||
ev = self._make_events()
|
||||
ev._aggregator.get_number_of_workers.return_value = 5
|
||||
self.assertEqual(ev.get_number_of_workers(), 5)
|
||||
|
||||
def test_clear_events(self):
|
||||
ev = self._make_events()
|
||||
ev.clear_events()
|
||||
ev._aggregator.clear_events.assert_called()
|
||||
ev._aggregator.reset_workers.assert_called()
|
||||
|
||||
def test_repr(self):
|
||||
ev = self._make_events()
|
||||
ev._aggregator.get_all_events.return_value = []
|
||||
s = repr(ev)
|
||||
self.assertIn("AscendStoreKVEvents", s)
|
||||
|
||||
|
||||
class TestAscendStoreConnector(unittest.TestCase):
|
||||
def _make_vllm_config(self, kv_role="kv_producer", extra_config=None):
|
||||
config = MagicMock()
|
||||
config.kv_transfer_config.kv_role = kv_role
|
||||
config.kv_transfer_config.kv_connector = "AscendStoreConnector"
|
||||
config.kv_transfer_config.kv_connector_extra_config = extra_config or {}
|
||||
config.parallel_config.rank = 0
|
||||
return config
|
||||
|
||||
def test_pp_handshake_metadata_is_ignored(self):
|
||||
connector = AscendStoreConnector.__new__(AscendStoreConnector)
|
||||
metadata = {
|
||||
(0, 0): MagicMock(),
|
||||
(1, 0): MagicMock(),
|
||||
}
|
||||
original_metadata = metadata.copy()
|
||||
|
||||
result = connector.set_xfer_handshake_metadata_pp_aware(metadata)
|
||||
|
||||
self.assertIsNone(result)
|
||||
self.assertEqual(metadata, original_metadata)
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolScheduler")
|
||||
def test_init_scheduler_role(self, mock_scheduler_cls):
|
||||
config = self._make_vllm_config()
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
_connector = AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.SCHEDULER,
|
||||
kv_cache_config=MagicMock(),
|
||||
)
|
||||
mock_scheduler_cls.assert_called_once()
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
|
||||
def test_init_worker_role(self, mock_worker_cls, mock_lookup_cls):
|
||||
config = self._make_vllm_config()
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
_connector = AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.WORKER,
|
||||
kv_cache_config=None,
|
||||
)
|
||||
mock_worker_cls.assert_called_once()
|
||||
mock_lookup_cls.assert_called_once()
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolScheduler")
|
||||
def test_scheduler_methods_delegate(self, mock_scheduler_cls):
|
||||
config = self._make_vllm_config()
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
connector = AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.SCHEDULER,
|
||||
kv_cache_config=MagicMock(),
|
||||
)
|
||||
mock_sched = mock_scheduler_cls.return_value
|
||||
|
||||
# get_num_new_matched_tokens
|
||||
mock_sched.get_num_new_matched_tokens.return_value = (10, False)
|
||||
result = connector.get_num_new_matched_tokens(MagicMock(), 5)
|
||||
self.assertEqual(result, (10, False))
|
||||
|
||||
# update_state_after_alloc
|
||||
connector.update_state_after_alloc(MagicMock(), MagicMock(), 10)
|
||||
mock_sched.update_state_after_alloc.assert_called_once()
|
||||
|
||||
# build_connector_meta
|
||||
connector.build_connector_meta(MagicMock())
|
||||
mock_sched.build_connector_meta.assert_called_once()
|
||||
|
||||
# request_finished
|
||||
mock_sched.request_finished.return_value = (True, None)
|
||||
result = connector.request_finished(MagicMock(), [1, 2])
|
||||
self.assertEqual(result, (True, None))
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolScheduler")
|
||||
def test_update_connector_output_no_events(self, mock_scheduler_cls):
|
||||
config = self._make_vllm_config()
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
connector = AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.SCHEDULER,
|
||||
kv_cache_config=MagicMock(),
|
||||
)
|
||||
output = MagicMock()
|
||||
output.kv_cache_events = None
|
||||
connector.update_connector_output(output)
|
||||
self.assertIsNone(connector._kv_cache_events)
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolScheduler")
|
||||
def test_update_connector_output_with_events(self, mock_scheduler_cls):
|
||||
config = self._make_vllm_config()
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
connector = AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.SCHEDULER,
|
||||
kv_cache_config=MagicMock(),
|
||||
)
|
||||
events = _mock_events(num_workers=1)
|
||||
mock_kv_events = [MagicMock()]
|
||||
events._aggregator.get_all_events.return_value = mock_kv_events
|
||||
events._aggregator.get_number_of_workers.return_value = 1
|
||||
|
||||
output = MagicMock()
|
||||
output.kv_cache_events = events
|
||||
connector.update_connector_output(output)
|
||||
self.assertIsNotNone(connector._kv_cache_events)
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolScheduler")
|
||||
def test_update_connector_output_accumulate(self, mock_scheduler_cls):
|
||||
config = self._make_vllm_config()
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
connector = AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.SCHEDULER,
|
||||
kv_cache_config=MagicMock(),
|
||||
)
|
||||
# First update
|
||||
events1 = _mock_events(num_workers=1)
|
||||
events1._aggregator.get_all_events.return_value = [MagicMock()]
|
||||
events1._aggregator.get_number_of_workers.return_value = 1
|
||||
output1 = MagicMock()
|
||||
output1.kv_cache_events = events1
|
||||
connector.update_connector_output(output1)
|
||||
|
||||
# Second update
|
||||
events2 = _mock_events(num_workers=1)
|
||||
events2._aggregator.get_all_events.return_value = [MagicMock()]
|
||||
events2._aggregator.get_number_of_workers.return_value = 1
|
||||
output2 = MagicMock()
|
||||
output2.kv_cache_events = events2
|
||||
connector.update_connector_output(output2)
|
||||
self.assertIsNotNone(connector._kv_cache_events)
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolScheduler")
|
||||
def test_take_events(self, mock_scheduler_cls):
|
||||
config = self._make_vllm_config()
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
connector = AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.SCHEDULER,
|
||||
kv_cache_config=MagicMock(),
|
||||
)
|
||||
# No events
|
||||
result = list(connector.take_events())
|
||||
self.assertEqual(result, [])
|
||||
|
||||
# With events
|
||||
events = _mock_events(num_workers=1)
|
||||
mock_event = MagicMock()
|
||||
events._aggregator.get_common_events.return_value = [mock_event]
|
||||
events._aggregator.get_all_events.return_value = [mock_event]
|
||||
connector._kv_cache_events = events
|
||||
result = list(connector.take_events())
|
||||
self.assertEqual(len(result), 1)
|
||||
self.assertIsNone(connector._kv_cache_events)
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
|
||||
def test_worker_methods(self, mock_worker_cls, mock_lookup_cls):
|
||||
config = self._make_vllm_config()
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
connector = AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.WORKER,
|
||||
kv_cache_config=None,
|
||||
)
|
||||
mock_worker = mock_worker_cls.return_value
|
||||
|
||||
# register_kv_caches
|
||||
connector.register_kv_caches({"layer1": MagicMock()})
|
||||
mock_worker.register_kv_caches.assert_called_once()
|
||||
|
||||
# start_load_kv
|
||||
connector._get_connector_metadata = MagicMock(return_value=MagicMock())
|
||||
connector.start_load_kv(MagicMock())
|
||||
mock_worker.start_load_kv.assert_called_once()
|
||||
|
||||
# wait_for_save (non-consumer)
|
||||
connector.kv_role = "kv_producer"
|
||||
connector.use_layerwise = False
|
||||
connector.wait_for_save()
|
||||
mock_worker.wait_for_save.assert_called_once()
|
||||
|
||||
# get_finished
|
||||
mock_worker.get_finished.return_value = ({"r1"}, {"r2"})
|
||||
done_s, done_r = connector.get_finished({"r1"})
|
||||
self.assertEqual(done_s, {"r1"})
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
|
||||
def test_wait_for_layer_load_not_layerwise(self, mock_worker_cls, mock_lookup_cls):
|
||||
config = self._make_vllm_config(extra_config={"use_layerwise": False})
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
connector = AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.WORKER,
|
||||
kv_cache_config=None,
|
||||
)
|
||||
# Should return immediately without calling worker
|
||||
connector.wait_for_layer_load("layer_0")
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
|
||||
def test_save_kv_layer_not_layerwise(self, mock_worker_cls, mock_lookup_cls):
|
||||
config = self._make_vllm_config(extra_config={"use_layerwise": False})
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
connector = AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.WORKER,
|
||||
kv_cache_config=None,
|
||||
)
|
||||
connector.save_kv_layer("layer_0", MagicMock(), MagicMock())
|
||||
# Should return immediately
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
|
||||
def test_save_kv_layer_consumer(self, mock_worker_cls, mock_lookup_cls):
|
||||
config = self._make_vllm_config(kv_role="kv_consumer", extra_config={"use_layerwise": True})
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
connector = AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.WORKER,
|
||||
kv_cache_config=None,
|
||||
)
|
||||
connector.save_kv_layer("layer_0", MagicMock(), MagicMock())
|
||||
# Consumer should not save
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
|
||||
def test_wait_for_save_consumer(self, mock_worker_cls, mock_lookup_cls):
|
||||
config = self._make_vllm_config(kv_role="kv_consumer")
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
connector = AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.WORKER,
|
||||
kv_cache_config=None,
|
||||
)
|
||||
connector.wait_for_save()
|
||||
mock_worker_cls.return_value.wait_for_save.assert_not_called()
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
|
||||
def test_get_kv_connector_kv_cache_events_empty(self, mock_worker_cls, mock_lookup_cls):
|
||||
config = self._make_vllm_config()
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
connector = AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.WORKER,
|
||||
kv_cache_config=None,
|
||||
)
|
||||
mock_worker_cls.return_value.get_kv_events.return_value = []
|
||||
result = connector.get_kv_connector_kv_cache_events()
|
||||
self.assertIsNone(result)
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.LookupKeyServer")
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector.KVPoolWorker")
|
||||
def test_get_kv_connector_kv_cache_events_with_events(self, mock_worker_cls, mock_lookup_cls):
|
||||
config = self._make_vllm_config()
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
connector = AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.WORKER,
|
||||
kv_cache_config=None,
|
||||
)
|
||||
mock_worker_cls.return_value.get_kv_events.return_value = [MagicMock()]
|
||||
result = connector.get_kv_connector_kv_cache_events()
|
||||
self.assertIsNotNone(result)
|
||||
self.assertIsInstance(result, AscendStoreKVEvents)
|
||||
|
||||
|
||||
class TestAscendStoreConnectorLayerwise(unittest.TestCase):
|
||||
"""Test connector methods that are specific to layerwise mode."""
|
||||
|
||||
connector_mod: types.ModuleType
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store import ascend_store_connector
|
||||
|
||||
cls.connector_mod = ascend_store_connector
|
||||
|
||||
def test_requires_piecewise_for_cudagraph_enabled(self):
|
||||
self.assertTrue(
|
||||
self.connector_mod.AscendStoreConnector.requires_piecewise_for_cudagraph({"use_layerwise": True})
|
||||
)
|
||||
|
||||
def test_requires_piecewise_for_cudagraph_disabled(self):
|
||||
self.assertFalse(
|
||||
self.connector_mod.AscendStoreConnector.requires_piecewise_for_cudagraph({"use_layerwise": False})
|
||||
)
|
||||
|
||||
def test_requires_piecewise_for_cudagraph_missing(self):
|
||||
self.assertFalse(self.connector_mod.AscendStoreConnector.requires_piecewise_for_cudagraph({}))
|
||||
|
||||
def test_wait_for_save_layerwise_returns_early(self):
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
with (
|
||||
patch.object(self.connector_mod, "KVPoolWorker") as mock_worker_cls,
|
||||
patch.object(self.connector_mod, "LookupKeyServer") as _mock_lookup_cls,
|
||||
):
|
||||
config = MagicMock()
|
||||
config.kv_transfer_config.kv_role = "kv_producer"
|
||||
config.kv_transfer_config.kv_connector = "AscendStoreConnector"
|
||||
config.kv_transfer_config.kv_connector_extra_config = {"use_layerwise": True}
|
||||
config.parallel_config.rank = 0
|
||||
|
||||
connector = self.connector_mod.AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.WORKER,
|
||||
kv_cache_config=None,
|
||||
)
|
||||
connector.wait_for_save()
|
||||
mock_worker_cls.return_value.wait_for_save.assert_not_called()
|
||||
|
||||
def test_save_kv_layer_layerwise_producer(self):
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
with (
|
||||
patch.object(self.connector_mod, "KVPoolWorker") as mock_worker_cls,
|
||||
patch.object(self.connector_mod, "LookupKeyServer") as _mock_lookup_cls,
|
||||
):
|
||||
config = MagicMock()
|
||||
config.kv_transfer_config.kv_role = "kv_producer"
|
||||
config.kv_transfer_config.kv_connector = "AscendStoreConnector"
|
||||
config.kv_transfer_config.kv_connector_extra_config = {"use_layerwise": True}
|
||||
config.parallel_config.rank = 0
|
||||
|
||||
connector = self.connector_mod.AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.WORKER,
|
||||
kv_cache_config=None,
|
||||
)
|
||||
connector._get_connector_metadata = MagicMock(return_value=MagicMock())
|
||||
connector.save_kv_layer("layer_0", MagicMock(), MagicMock())
|
||||
mock_worker_cls.return_value.save_kv_layer.assert_called_once()
|
||||
|
||||
def test_wait_for_layer_load_layerwise(self):
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole
|
||||
|
||||
with (
|
||||
patch.object(self.connector_mod, "KVPoolWorker") as mock_worker_cls,
|
||||
patch.object(self.connector_mod, "LookupKeyServer") as _mock_lookup_cls,
|
||||
):
|
||||
config = MagicMock()
|
||||
config.kv_transfer_config.kv_role = "kv_consumer"
|
||||
config.kv_transfer_config.kv_connector = "AscendStoreConnector"
|
||||
config.kv_transfer_config.kv_connector_extra_config = {"use_layerwise": True}
|
||||
config.parallel_config.rank = 0
|
||||
|
||||
connector = self.connector_mod.AscendStoreConnector(
|
||||
vllm_config=config,
|
||||
role=KVConnectorRole.WORKER,
|
||||
kv_cache_config=None,
|
||||
)
|
||||
connector.wait_for_layer_load("layer_0")
|
||||
mock_worker_cls.return_value.wait_for_layer_load.assert_called_once()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
665
tests/ut/distributed/ascend_store/test_backend.py
Normal file
665
tests/ut/distributed/ascend_store/test_backend.py
Normal file
@@ -0,0 +1,665 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import json
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import tests.ut.distributed.ascend_store._mock_deps # noqa: F401, E402
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.backend import Backend
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend import (
|
||||
MooncakeStoreConfig,
|
||||
_convert_to_bytes,
|
||||
_parse_global_segment_size,
|
||||
_ssd_setup_kwargs,
|
||||
)
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.yuanrong_backend import (
|
||||
YuanrongConfig,
|
||||
YuanrongHelper,
|
||||
)
|
||||
|
||||
|
||||
def _format_log_call(call):
|
||||
args = call.args
|
||||
return args[0] % args[1:]
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# Backend ABC
|
||||
# =========================================================================
|
||||
class TestBackendABC(unittest.TestCase):
|
||||
def test_cannot_instantiate(self):
|
||||
with self.assertRaises(TypeError):
|
||||
Backend(MagicMock()) # type: ignore[abstract]
|
||||
|
||||
|
||||
def _make_mooncake_store_config(**overrides) -> MooncakeStoreConfig:
|
||||
"""Build MooncakeStoreConfig via from_file(); inherits from_file() defaults."""
|
||||
config = dict(overrides)
|
||||
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
|
||||
json.dump(config, f)
|
||||
f.flush()
|
||||
path = f.name
|
||||
try:
|
||||
return MooncakeStoreConfig.from_file(path)
|
||||
finally:
|
||||
os.unlink(path)
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# MooncakeStoreConfig
|
||||
# =========================================================================
|
||||
class TestMooncakeStoreConfig(unittest.TestCase):
|
||||
def test_from_file(self):
|
||||
config = {
|
||||
"metadata_server": "127.0.0.1:2379",
|
||||
"global_segment_size": "2GB",
|
||||
"local_buffer_size": "1GB",
|
||||
"protocol": "ascend",
|
||||
"device_name": "npu0",
|
||||
"master_server_address": "127.0.0.1:8080",
|
||||
}
|
||||
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
|
||||
json.dump(config, f)
|
||||
f.flush()
|
||||
path = f.name
|
||||
|
||||
try:
|
||||
cfg = MooncakeStoreConfig.from_file(path)
|
||||
self.assertEqual(cfg.metadata_server, "127.0.0.1:2379")
|
||||
self.assertEqual(cfg.global_segment_size, 2 * 1024**3)
|
||||
self.assertEqual(cfg.local_buffer_size, 1 * 1024**3)
|
||||
self.assertEqual(cfg.protocol, "ascend")
|
||||
self.assertEqual(cfg.device_name, "npu0")
|
||||
finally:
|
||||
os.unlink(path)
|
||||
|
||||
def test_from_file_defaults(self):
|
||||
config = {
|
||||
"metadata_server": "localhost:2379",
|
||||
"master_server_address": "localhost:8080",
|
||||
}
|
||||
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
|
||||
json.dump(config, f)
|
||||
f.flush()
|
||||
path = f.name
|
||||
|
||||
try:
|
||||
cfg = MooncakeStoreConfig.from_file(path)
|
||||
self.assertEqual(cfg.protocol, "ascend")
|
||||
self.assertEqual(cfg.device_name, "")
|
||||
self.assertFalse(cfg.enable_ssd_offload)
|
||||
self.assertEqual(cfg.ssd_offload_path, "")
|
||||
finally:
|
||||
os.unlink(path)
|
||||
|
||||
def test_from_file_ssd_offload(self):
|
||||
ssd_path = TestMooncakeStoreConfig._writable_ssd_path()
|
||||
self.addCleanup(lambda: os.rmdir(ssd_path))
|
||||
cfg = _make_mooncake_store_config(
|
||||
enable_ssd_offload=True,
|
||||
ssd_offload_path=ssd_path,
|
||||
)
|
||||
self.assertTrue(cfg.enable_ssd_offload)
|
||||
self.assertEqual(cfg.ssd_offload_path, ssd_path)
|
||||
|
||||
def test_ssd_offload_requires_absolute_path(self):
|
||||
with self.assertRaises(ValueError):
|
||||
_make_mooncake_store_config(
|
||||
enable_ssd_offload=True,
|
||||
ssd_offload_path="relative/path",
|
||||
)
|
||||
|
||||
def test_ssd_offload_requires_path_in_json(self):
|
||||
with self.assertRaises(ValueError):
|
||||
_make_mooncake_store_config(enable_ssd_offload=True)
|
||||
|
||||
@staticmethod
|
||||
def _writable_ssd_path() -> str:
|
||||
return tempfile.mkdtemp(prefix="mooncake_ssd_ut_")
|
||||
|
||||
@patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend."
|
||||
"mooncake_backend._mooncake_setup_supports_ssd_offload",
|
||||
return_value=False,
|
||||
)
|
||||
def test_ssd_setup_kwargs_off_when_disabled(self, _mock_supports):
|
||||
cfg = _make_mooncake_store_config()
|
||||
self.assertEqual(_ssd_setup_kwargs(cfg), {})
|
||||
|
||||
@patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend."
|
||||
"mooncake_backend._mooncake_setup_supports_ssd_offload",
|
||||
return_value=False,
|
||||
)
|
||||
def test_ssd_setup_kwargs_raises_on_old_mooncake(self, _mock_supports):
|
||||
ssd_path = TestMooncakeStoreConfig._writable_ssd_path()
|
||||
self.addCleanup(lambda: os.rmdir(ssd_path))
|
||||
cfg = _make_mooncake_store_config(
|
||||
enable_ssd_offload=True,
|
||||
ssd_offload_path=ssd_path,
|
||||
)
|
||||
with self.assertRaises(RuntimeError):
|
||||
_ssd_setup_kwargs(cfg)
|
||||
|
||||
@patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend."
|
||||
"mooncake_backend._mooncake_setup_supports_ssd_offload",
|
||||
return_value=True,
|
||||
)
|
||||
def test_ssd_setup_kwargs_when_supported(self, _mock_supports):
|
||||
ssd_path = TestMooncakeStoreConfig._writable_ssd_path()
|
||||
self.addCleanup(lambda: os.rmdir(ssd_path))
|
||||
cfg = _make_mooncake_store_config(
|
||||
enable_ssd_offload=True,
|
||||
ssd_offload_path=ssd_path,
|
||||
)
|
||||
self.assertEqual(
|
||||
_ssd_setup_kwargs(cfg),
|
||||
{
|
||||
"enable_ssd_offload": cfg.enable_ssd_offload,
|
||||
"ssd_offload_path": cfg.ssd_offload_path,
|
||||
},
|
||||
)
|
||||
|
||||
def test_load_from_env_missing(self):
|
||||
with patch.dict(os.environ, {}, clear=True):
|
||||
os.environ.pop("MOONCAKE_CONFIG_PATH", None)
|
||||
with self.assertRaises(ValueError):
|
||||
MooncakeStoreConfig.load_from_env()
|
||||
|
||||
def test_load_from_env(self):
|
||||
config = {
|
||||
"metadata_server": "host:1234",
|
||||
"master_server_address": "host:5678",
|
||||
}
|
||||
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
|
||||
json.dump(config, f)
|
||||
f.flush()
|
||||
path = f.name
|
||||
|
||||
try:
|
||||
with patch.dict(os.environ, {"MOONCAKE_CONFIG_PATH": path}):
|
||||
cfg = MooncakeStoreConfig.load_from_env()
|
||||
self.assertEqual(cfg.metadata_server, "host:1234")
|
||||
finally:
|
||||
os.unlink(path)
|
||||
|
||||
|
||||
class TestParseGlobalSegmentSize(unittest.TestCase):
|
||||
def test_int(self):
|
||||
self.assertEqual(_parse_global_segment_size(1024), 1024)
|
||||
|
||||
def test_gb(self):
|
||||
self.assertEqual(_parse_global_segment_size("2GB"), 2 * 1024**3)
|
||||
|
||||
def test_mb(self):
|
||||
self.assertEqual(_parse_global_segment_size("512MB"), 512 * 1024**2)
|
||||
|
||||
def test_kb(self):
|
||||
self.assertEqual(_parse_global_segment_size("256KB"), 256 * 1024)
|
||||
|
||||
def test_b(self):
|
||||
self.assertEqual(_parse_global_segment_size("4096B"), 4096)
|
||||
|
||||
def test_no_unit(self):
|
||||
self.assertEqual(_parse_global_segment_size("2048"), 2048)
|
||||
|
||||
def test_float_input(self):
|
||||
self.assertEqual(_parse_global_segment_size(2048.0), 2048)
|
||||
|
||||
def test_empty_string(self):
|
||||
with self.assertRaises(ValueError):
|
||||
_parse_global_segment_size("")
|
||||
|
||||
def test_invalid_format(self):
|
||||
with self.assertRaises(ValueError):
|
||||
_parse_global_segment_size("abcGB")
|
||||
|
||||
def test_unsupported_type(self):
|
||||
with self.assertRaises(TypeError):
|
||||
_parse_global_segment_size(None) # type: ignore[arg-type]
|
||||
|
||||
|
||||
class TestConvertToBytes(unittest.TestCase):
|
||||
def test_valid(self):
|
||||
self.assertEqual(_convert_to_bytes("10", 1, "10"), 10)
|
||||
self.assertEqual(_convert_to_bytes("1.5", 1024, "1.5KB"), int(1.5 * 1024))
|
||||
|
||||
def test_invalid_number(self):
|
||||
with self.assertRaises(ValueError):
|
||||
_convert_to_bytes("abc", 1, "abc")
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# YuanrongConfig
|
||||
# =========================================================================
|
||||
class TestYuanrongConfig(unittest.TestCase):
|
||||
def test_load_from_env(self):
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"DS_WORKER_ADDR": "host:1234",
|
||||
"DS_ENABLE_EXCLUSIVE_CONNECTION": "1",
|
||||
"DS_ENABLE_REMOTE_H2D": "0",
|
||||
},
|
||||
):
|
||||
cfg = YuanrongConfig.load_from_env()
|
||||
self.assertEqual(cfg.worker_addr, "host:1234")
|
||||
self.assertTrue(cfg.enable_exclusive_connection)
|
||||
self.assertFalse(cfg.enable_remote_h2d)
|
||||
|
||||
def test_load_from_env_missing(self):
|
||||
with patch.dict(os.environ, {}, clear=True):
|
||||
os.environ.pop("DS_WORKER_ADDR", None)
|
||||
with self.assertRaises(ValueError):
|
||||
YuanrongConfig.load_from_env()
|
||||
|
||||
def test_load_from_env_defaults(self):
|
||||
with patch.dict(os.environ, {"DS_WORKER_ADDR": "h:1"}):
|
||||
cfg = YuanrongConfig.load_from_env()
|
||||
self.assertFalse(cfg.enable_exclusive_connection)
|
||||
self.assertFalse(cfg.enable_remote_h2d)
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# YuanrongHelper
|
||||
# =========================================================================
|
||||
class TestYuanrongHelper(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.blob_cls = MagicMock()
|
||||
self.blob_list_cls = MagicMock()
|
||||
self.helper = YuanrongHelper(self.blob_cls, self.blob_list_cls)
|
||||
|
||||
def test_normalize_keys_short_valid(self):
|
||||
keys = ["abc-123", "key_2"]
|
||||
result = self.helper.normalize_keys(keys)
|
||||
self.assertEqual(result, keys)
|
||||
|
||||
def test_normalize_keys_with_invalid_chars(self):
|
||||
keys = ["key with spaces/and.dots"]
|
||||
result = self.helper.normalize_keys(keys)
|
||||
self.assertEqual(len(result), 1)
|
||||
# Should not contain the original invalid chars
|
||||
self.assertNotIn(" ", result[0])
|
||||
self.assertNotIn("/", result[0])
|
||||
# Should have hash suffix
|
||||
self.assertIn("__", result[0])
|
||||
|
||||
def test_normalize_keys_at_max_length(self):
|
||||
max_length_key = "a" * 1024
|
||||
result = self.helper.normalize_keys([max_length_key])
|
||||
self.assertEqual(result, [max_length_key])
|
||||
|
||||
def test_normalize_keys_over_max_length(self):
|
||||
long_key = "a" * 1025
|
||||
result = self.helper.normalize_keys([long_key])
|
||||
self.assertEqual(len(result), 1)
|
||||
self.assertEqual(len(result[0]), 1024)
|
||||
self.assertIn("__", result[0])
|
||||
|
||||
def test_make_blob_lists(self):
|
||||
self.helper._device_id = 0
|
||||
addrs = [[100, 200], [300, 400]]
|
||||
sizes = [[10, 20], [30, 40]]
|
||||
result = self.helper.make_blob_lists(addrs, sizes)
|
||||
self.assertEqual(len(result), 2)
|
||||
self.assertEqual(self.blob_cls.call_count, 4)
|
||||
|
||||
def test_make_blob_lists_length_mismatch(self):
|
||||
self.helper._device_id = 0
|
||||
with self.assertRaises(ValueError):
|
||||
self.helper.make_blob_lists([[1]], [[1, 2], [3, 4]])
|
||||
|
||||
def test_make_blob_lists_inner_length_mismatch(self):
|
||||
self.helper._device_id = 0
|
||||
with self.assertRaises(ValueError):
|
||||
self.helper.make_blob_lists([[1, 2]], [[1]])
|
||||
|
||||
def test_make_blob_lists_no_device(self):
|
||||
self.helper._device_id = None
|
||||
with self.assertRaises(RuntimeError):
|
||||
self.helper.make_blob_lists([[1]], [[1]])
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# MooncakeBackend (mocked store)
|
||||
# =========================================================================
|
||||
class TestMooncakeBackendMethods(unittest.TestCase):
|
||||
def _make_backend(self):
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend import MooncakeBackend
|
||||
|
||||
with (
|
||||
patch.dict(os.environ, {"MOONCAKE_CONFIG_PATH": "/dev/null"}),
|
||||
patch.object(MooncakeBackend, "__init__", lambda self, pc: None),
|
||||
):
|
||||
backend = MooncakeBackend.__new__(MooncakeBackend)
|
||||
backend.store = MagicMock()
|
||||
backend.config = MagicMock()
|
||||
backend.local_seg = "127.0.0.1:1234"
|
||||
backend._lazy_init = False
|
||||
backend._store_initialized = True
|
||||
backend._use_fabric_mem = False
|
||||
backend._store_init_lock = MagicMock()
|
||||
backend.local_seg = None
|
||||
return backend
|
||||
|
||||
def test_exists(self):
|
||||
b = self._make_backend()
|
||||
b.store.batch_is_exist.return_value = [1, 0]
|
||||
result = b.exists(["k1", "k2"])
|
||||
self.assertEqual(result, [1, 0])
|
||||
|
||||
def test_put(self):
|
||||
b = self._make_backend()
|
||||
b.store.batch_put_from_multi_buffers.return_value = [0, 0]
|
||||
b.put(["k1"], [[100]], [[10]])
|
||||
b.store.batch_put_from_multi_buffers.assert_called_once()
|
||||
|
||||
def test_put_error(self):
|
||||
b = self._make_backend()
|
||||
b.store.batch_put_from_multi_buffers.return_value = [-1]
|
||||
b.put(["k1"], [[100]], [[10]]) # Should log error but not raise
|
||||
|
||||
def test_put_exception(self):
|
||||
b = self._make_backend()
|
||||
b.store.batch_put_from_multi_buffers.side_effect = RuntimeError("backend fail")
|
||||
with patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend.logger"
|
||||
) as mock_logger:
|
||||
b.put(["k1"], [[100]], [[10]]) # Should log error but not raise
|
||||
error_log = _format_log_call(mock_logger.error.call_args)
|
||||
self.assertIn("RuntimeError", error_log)
|
||||
self.assertIn("backend fail", error_log)
|
||||
|
||||
def test_get(self):
|
||||
b = self._make_backend()
|
||||
b.store.batch_get_into_multi_buffers.return_value = [0]
|
||||
b.get(["k1"], [[100]], [[10]])
|
||||
b.store.batch_get_into_multi_buffers.assert_called_once()
|
||||
|
||||
def test_get_error(self):
|
||||
b = self._make_backend()
|
||||
b.store.batch_get_into_multi_buffers.return_value = [-1]
|
||||
b.get(["k1"], [[100]], [[10]])
|
||||
|
||||
def test_get_exception(self):
|
||||
b = self._make_backend()
|
||||
b.store.batch_get_into_multi_buffers.side_effect = RuntimeError("backend fail")
|
||||
with patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend.logger"
|
||||
) as mock_logger:
|
||||
b.get(["k1"], [[100]], [[10]])
|
||||
error_log = _format_log_call(mock_logger.error.call_args)
|
||||
self.assertIn("RuntimeError", error_log)
|
||||
self.assertIn("backend fail", error_log)
|
||||
|
||||
def test_register_buffer(self):
|
||||
b = self._make_backend()
|
||||
with (
|
||||
patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend.global_te"
|
||||
) as mock_te,
|
||||
patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend.get_ip"),
|
||||
):
|
||||
b.register_buffer([100], [200])
|
||||
mock_te.register_buffer.assert_called_once()
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# YuanrongBackend (mocked store)
|
||||
# =========================================================================
|
||||
class TestYuanrongBackendMethods(unittest.TestCase):
|
||||
def _make_backend(self):
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.yuanrong_backend import YuanrongBackend
|
||||
|
||||
with patch.object(YuanrongBackend, "__init__", lambda self, pc: None):
|
||||
backend = YuanrongBackend.__new__(YuanrongBackend)
|
||||
backend._helper = MagicMock()
|
||||
backend._helper._device_id = 0
|
||||
backend._helper.normalize_keys = lambda keys: keys
|
||||
backend._helper.make_blob_lists = lambda a, s: [MagicMock() for _ in a]
|
||||
backend._hetero_client = MagicMock()
|
||||
backend._ds_set_param = MagicMock()
|
||||
backend._is_a2 = False
|
||||
backend._registered_buffers = None
|
||||
backend._buffers_registered = False
|
||||
backend.config = YuanrongConfig(
|
||||
worker_addr="127.0.0.1:0",
|
||||
enable_exclusive_connection=False,
|
||||
enable_remote_h2d=False,
|
||||
)
|
||||
backend.rank = 0
|
||||
return backend
|
||||
|
||||
def test_exists_empty(self):
|
||||
b = self._make_backend()
|
||||
result = b.exists([])
|
||||
self.assertEqual(result, [])
|
||||
|
||||
def test_exists(self):
|
||||
b = self._make_backend()
|
||||
b._hetero_client.exist.return_value = [True, False]
|
||||
result = b.exists(["k1", "k2"])
|
||||
self.assertEqual(result, [1, 0])
|
||||
|
||||
def test_exists_exception(self):
|
||||
b = self._make_backend()
|
||||
b._hetero_client.exist.side_effect = Exception("fail")
|
||||
result = b.exists(["k1"])
|
||||
self.assertEqual(result, [0])
|
||||
|
||||
def test_get_empty(self):
|
||||
b = self._make_backend()
|
||||
result = b.get([], [], [])
|
||||
self.assertEqual(result, [])
|
||||
b._hetero_client.mget_h2d.assert_not_called()
|
||||
|
||||
def test_get(self):
|
||||
b = self._make_backend()
|
||||
b._hetero_client.mget_h2d.return_value = []
|
||||
result = b.get(["k1"], [[100]], [[10]])
|
||||
self.assertEqual(result, [0])
|
||||
b._hetero_client.mget_h2d.assert_called_once()
|
||||
|
||||
def test_get_partial_failure(self):
|
||||
b = self._make_backend()
|
||||
b._hetero_client.mget_h2d.return_value = ["k2"]
|
||||
result = b.get(["k1", "k2", "k3"], [[100], [200], [300]], [[10], [20], [30]])
|
||||
self.assertEqual(result, [0, 1, 0])
|
||||
|
||||
def test_get_failed_keys(self):
|
||||
b = self._make_backend()
|
||||
b._hetero_client.mget_h2d.return_value = ["k1"]
|
||||
result = b.get(["k1"], [[100]], [[10]]) # Should log error
|
||||
self.assertEqual(result, [1])
|
||||
|
||||
def test_get_exception(self):
|
||||
b = self._make_backend()
|
||||
b._hetero_client.mget_h2d.side_effect = RuntimeError("backend fail")
|
||||
with patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.yuanrong_backend.logger"
|
||||
) as mock_logger:
|
||||
result = b.get(["k1"], [[100]], [[10]])
|
||||
error_log = _format_log_call(mock_logger.error.call_args)
|
||||
self.assertIsNone(result)
|
||||
self.assertIn("RuntimeError", error_log)
|
||||
self.assertIn("backend fail", error_log)
|
||||
|
||||
def test_put_empty(self):
|
||||
b = self._make_backend()
|
||||
b.put([], [], [])
|
||||
b._hetero_client.mset_d2h.assert_not_called()
|
||||
|
||||
def test_put(self):
|
||||
b = self._make_backend()
|
||||
b.put(["k1"], [[100]], [[10]])
|
||||
b._hetero_client.mset_d2h.assert_called_once()
|
||||
|
||||
def test_put_exception(self):
|
||||
b = self._make_backend()
|
||||
b._hetero_client.mset_d2h.side_effect = RuntimeError("backend fail")
|
||||
with patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.yuanrong_backend.logger"
|
||||
) as mock_logger:
|
||||
b.put(["k1"], [[100]], [[10]])
|
||||
error_log = _format_log_call(mock_logger.error.call_args)
|
||||
self.assertIn("RuntimeError", error_log)
|
||||
self.assertIn("backend fail", error_log)
|
||||
|
||||
def test_register_buffer_noop_when_remote_h2d_disabled(self):
|
||||
b = self._make_backend()
|
||||
b.register_buffer([100], [200])
|
||||
b._hetero_client.pre_register_device_memory.assert_not_called()
|
||||
|
||||
def test_register_buffer_when_remote_h2d_enabled(self):
|
||||
b = self._make_backend()
|
||||
b.config.enable_remote_h2d = True
|
||||
b.register_buffer([100], [200])
|
||||
b._hetero_client.pre_register_device_memory.assert_called_once_with([100], [200])
|
||||
|
||||
def test_register_buffer_noop_on_a2(self):
|
||||
# A2 must not register (opposite of memcache_backend's _is_a2 gating).
|
||||
b = self._make_backend()
|
||||
b._is_a2 = True
|
||||
b.config.enable_remote_h2d = True
|
||||
b.register_buffer([100], [200])
|
||||
b._hetero_client.pre_register_device_memory.assert_not_called()
|
||||
|
||||
def test_register_buffer_idempotent(self):
|
||||
b = self._make_backend()
|
||||
b.config.enable_remote_h2d = True
|
||||
b.register_buffer([100], [200])
|
||||
b.register_buffer([300], [400])
|
||||
b._hetero_client.pre_register_device_memory.assert_called_once_with([100], [200])
|
||||
|
||||
def test_register_buffers_if_needed_no_buffers(self):
|
||||
b = self._make_backend()
|
||||
b.config.enable_remote_h2d = True
|
||||
b._registered_buffers = None
|
||||
b._register_buffers_if_needed()
|
||||
b._hetero_client.pre_register_device_memory.assert_not_called()
|
||||
|
||||
def test_register_buffers_if_needed_already_registered(self):
|
||||
b = self._make_backend()
|
||||
b.config.enable_remote_h2d = True
|
||||
b._registered_buffers = ([100], [200])
|
||||
b._buffers_registered = True
|
||||
b._register_buffers_if_needed()
|
||||
b._hetero_client.pre_register_device_memory.assert_not_called()
|
||||
|
||||
def test_register_buffers_if_needed_disabled(self):
|
||||
b = self._make_backend()
|
||||
b.config.enable_remote_h2d = False
|
||||
b._registered_buffers = ([100], [200])
|
||||
b._register_buffers_if_needed()
|
||||
b._hetero_client.pre_register_device_memory.assert_not_called()
|
||||
|
||||
def test_ensure_device_ready(self):
|
||||
b = self._make_backend()
|
||||
b._helper._device_id = None
|
||||
b.set_device = MagicMock()
|
||||
b._ensure_device_ready()
|
||||
b.set_device.assert_called_once()
|
||||
|
||||
def test_ensure_device_ready_already_set(self):
|
||||
b = self._make_backend()
|
||||
b._helper._device_id = 0
|
||||
b.set_device = MagicMock()
|
||||
b._ensure_device_ready()
|
||||
b.set_device.assert_not_called()
|
||||
|
||||
|
||||
# =========================================================================
|
||||
# MemcacheBackend (mocked store)
|
||||
# =========================================================================
|
||||
class TestMemcacheBackendMethods(unittest.TestCase):
|
||||
def _make_backend(self):
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.memcache_backend import MemcacheBackend
|
||||
|
||||
with patch.object(MemcacheBackend, "__init__", lambda self, pc: None):
|
||||
backend = MemcacheBackend.__new__(MemcacheBackend)
|
||||
backend.store = MagicMock()
|
||||
backend.local_rank = 0
|
||||
# Set internal state to avoid lazy init logic during tests
|
||||
backend._lazy_init = False
|
||||
backend._store_initialized = True
|
||||
backend._is_a2 = False
|
||||
backend._registered_buffers = None
|
||||
backend._buffers_registered = False
|
||||
return backend
|
||||
|
||||
def test_exists(self):
|
||||
b = self._make_backend()
|
||||
b.store.batch_is_exist.return_value = [1]
|
||||
self.assertEqual(b.exists(["k1"]), [1])
|
||||
|
||||
def test_register_buffer(self):
|
||||
b = self._make_backend()
|
||||
b._is_a2 = True
|
||||
b.register_buffer([100], [200])
|
||||
b.store.register_buffer.assert_called_once()
|
||||
|
||||
def test_get(self):
|
||||
b = self._make_backend()
|
||||
b.store.batch_get_into_layers.return_value = [0]
|
||||
b.get(["k1"], [[100]], [[10]])
|
||||
b.store.batch_get_into_layers.assert_called_once()
|
||||
|
||||
def test_get_error(self):
|
||||
b = self._make_backend()
|
||||
b.store.batch_get_into_layers.return_value = [1] # non-zero = error
|
||||
b.get(["k1"], [[100]], [[10]])
|
||||
|
||||
def test_get_exception(self):
|
||||
b = self._make_backend()
|
||||
b.store.batch_get_into_layers.side_effect = RuntimeError("backend fail")
|
||||
with patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.memcache_backend.logger"
|
||||
) as mock_logger:
|
||||
b.get(["k1"], [[100]], [[10]])
|
||||
error_log = _format_log_call(mock_logger.error.call_args)
|
||||
self.assertIn("RuntimeError", error_log)
|
||||
self.assertIn("backend fail", error_log)
|
||||
|
||||
def test_put(self):
|
||||
b = self._make_backend()
|
||||
b.store.batch_put_from_layers.return_value = [0]
|
||||
b.put(["k1"], [[100]], [[10]])
|
||||
b.store.batch_put_from_layers.assert_called_once()
|
||||
|
||||
def test_put_error(self):
|
||||
b = self._make_backend()
|
||||
b.store.batch_put_from_layers.return_value = [1]
|
||||
b.put(["k1"], [[100]], [[10]])
|
||||
|
||||
def test_put_exception(self):
|
||||
b = self._make_backend()
|
||||
b.store.batch_put_from_layers.side_effect = RuntimeError("backend fail")
|
||||
with patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.memcache_backend.logger"
|
||||
) as mock_logger:
|
||||
b.put(["k1"], [[100]], [[10]])
|
||||
error_log = _format_log_call(mock_logger.error.call_args)
|
||||
self.assertIn("RuntimeError", error_log)
|
||||
self.assertIn("backend fail", error_log)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
581
tests/ut/distributed/ascend_store/test_config_data.py
Normal file
581
tests/ut/distributed/ascend_store/test_config_data.py
Normal file
@@ -0,0 +1,581 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import hashlib
|
||||
import unittest
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import tests.ut.distributed.ascend_store._mock_deps # noqa: F401, E402
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.config_data import (
|
||||
AscendConnectorMetadata,
|
||||
ChunkedTokenDatabase,
|
||||
KeyMetadata,
|
||||
LayerMultiBlockReqMeta,
|
||||
LayerPoolKey,
|
||||
LoadSpec,
|
||||
PoolKey,
|
||||
ReqMeta,
|
||||
RequestTracker,
|
||||
get_block_hashes,
|
||||
)
|
||||
|
||||
_GROUPED_BLOCK_HASH_DOMAIN = b"vllm-ascend-grouped-block-hash-v1\0"
|
||||
_GROUPED_BLOCK_HASH_LENGTH_PREFIX_BYTES = 4
|
||||
|
||||
|
||||
def _expected_grouped_hash(*block_hashes):
|
||||
hasher = hashlib.sha256()
|
||||
hasher.update(_GROUPED_BLOCK_HASH_DOMAIN)
|
||||
hasher.update(len(block_hashes).to_bytes(_GROUPED_BLOCK_HASH_LENGTH_PREFIX_BYTES, "big"))
|
||||
for block_hash in block_hashes:
|
||||
hash_bytes = block_hash.encode("utf-8") if isinstance(block_hash, str) else bytes(block_hash)
|
||||
hasher.update(len(hash_bytes).to_bytes(_GROUPED_BLOCK_HASH_LENGTH_PREFIX_BYTES, "big"))
|
||||
hasher.update(hash_bytes)
|
||||
return hasher.digest()
|
||||
|
||||
|
||||
class TestKeyMetadata(unittest.TestCase):
|
||||
def test_fields(self):
|
||||
meta = KeyMetadata(
|
||||
model_name="llama",
|
||||
head_or_tp_rank=0,
|
||||
pcp_rank=0,
|
||||
dcp_rank=0,
|
||||
pp_rank=0,
|
||||
)
|
||||
self.assertEqual(meta.model_name, "llama")
|
||||
self.assertEqual(meta.head_or_tp_rank, 0)
|
||||
self.assertEqual(meta.pcp_rank, 0)
|
||||
self.assertEqual(meta.dcp_rank, 0)
|
||||
self.assertEqual(meta.pp_rank, 0)
|
||||
|
||||
|
||||
class TestPoolKey(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.meta = KeyMetadata("llama", 1, 2, 3, 0)
|
||||
|
||||
def test_hash_equal(self):
|
||||
k1 = PoolKey(self.meta, "abc123")
|
||||
k2 = PoolKey(self.meta, "abc123")
|
||||
self.assertEqual(hash(k1), hash(k2))
|
||||
|
||||
def test_hash_diff(self):
|
||||
k1 = PoolKey(self.meta, "abc123")
|
||||
k2 = PoolKey(self.meta, "def456")
|
||||
self.assertNotEqual(hash(k1), hash(k2))
|
||||
|
||||
def test_to_string(self):
|
||||
k = PoolKey(self.meta, "hash1")
|
||||
s = k.to_string()
|
||||
self.assertIn("llama", s)
|
||||
self.assertIn("@pcp2", s)
|
||||
self.assertIn("@dcp3", s)
|
||||
self.assertIn("@head_or_tp_rank:1", s)
|
||||
self.assertIn("@pp_rank:0", s)
|
||||
self.assertIn("hash1", s)
|
||||
|
||||
def test_pp_ranks_use_distinct_keys(self):
|
||||
other_pp_meta = KeyMetadata("llama", 1, 2, 3, 1)
|
||||
pp0_key = PoolKey(self.meta, "hash1")
|
||||
pp1_key = PoolKey(other_pp_meta, "hash1")
|
||||
|
||||
self.assertNotEqual(pp0_key.to_string(), pp1_key.to_string())
|
||||
self.assertIn("@pp_rank:0", pp0_key.to_string())
|
||||
self.assertIn("@pp_rank:1", pp1_key.to_string())
|
||||
|
||||
def test_split_layers(self):
|
||||
k = PoolKey(self.meta, "hash1")
|
||||
layers = k.split_layers(3)
|
||||
self.assertEqual(len(layers), 3)
|
||||
for i, lk in enumerate(layers):
|
||||
self.assertIsInstance(lk, LayerPoolKey)
|
||||
self.assertEqual(lk.layer_id, i)
|
||||
self.assertEqual(lk.chunk_hash, "hash1")
|
||||
|
||||
|
||||
class TestLayerPoolKey(unittest.TestCase):
|
||||
def test_hash(self):
|
||||
meta = KeyMetadata("model", 0, 0, 0, 0)
|
||||
k1 = LayerPoolKey(meta, "h1", 0)
|
||||
k2 = LayerPoolKey(meta, "h1", 1)
|
||||
self.assertNotEqual(hash(k1), hash(k2))
|
||||
|
||||
def test_to_string_contains_layer_id(self):
|
||||
meta = KeyMetadata("model", 0, 0, 0, 0)
|
||||
k = LayerPoolKey(meta, "h1", 5)
|
||||
s = k.to_string()
|
||||
self.assertIn("@layer_id:5", s)
|
||||
self.assertIn("model", s)
|
||||
self.assertTrue(s.endswith("@h1"))
|
||||
|
||||
|
||||
class TestChunkedTokenDatabase(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.meta = KeyMetadata("llama", 0, 0, 0, 0)
|
||||
self.db = ChunkedTokenDatabase([self.meta], block_size=[16], partitions=None)
|
||||
self.db.set_group_buffers({0: [1000, 2000]}, {0: [160, 320]}, group_num_layers={0: 1})
|
||||
|
||||
def test_make_key_by_hash(self):
|
||||
key = self.db._make_key_by_hash("abc")
|
||||
self.assertIsInstance(key, PoolKey)
|
||||
self.assertEqual(key.chunk_hash, "abc")
|
||||
|
||||
def test_process_tokens_empty(self):
|
||||
result = list(self.db.process_tokens(32, []))
|
||||
self.assertEqual(result, [])
|
||||
|
||||
def test_process_tokens_with_str_hashes(self):
|
||||
hashes = ["aaa", "bbb"]
|
||||
result = list(self.db.process_tokens(32, hashes))
|
||||
self.assertEqual(len(result), 2)
|
||||
self.assertEqual(result[0][0], 0) # start
|
||||
self.assertEqual(result[0][1], 16) # end
|
||||
self.assertEqual(result[1][0], 16)
|
||||
self.assertEqual(result[1][1], 32)
|
||||
|
||||
def test_process_tokens_with_bytes_hashes(self):
|
||||
hashes = [b"\xaa\xbb", b"\xcc\xdd"]
|
||||
result = list(self.db.process_tokens(32, hashes))
|
||||
self.assertEqual(len(result), 2)
|
||||
|
||||
def test_process_tokens_with_mask(self):
|
||||
hashes = ["a", "b", "c"]
|
||||
result = list(self.db.process_tokens(48, hashes, mask_num=16))
|
||||
# first chunk (start=0 < mask_num=16) should be skipped
|
||||
self.assertEqual(len(result), 2)
|
||||
self.assertEqual(result[0][0], 16)
|
||||
|
||||
def test_process_tokens_with_tail_clipped_block_ids_maps_tail_chunks(self):
|
||||
db = ChunkedTokenDatabase([self.meta], block_size=[128], partitions=None)
|
||||
hashes = [bytes([idx % 251]) * 32 for idx in range(128)]
|
||||
|
||||
result = list(
|
||||
db.process_token_key_strings_with_block_ids(
|
||||
128 * 128,
|
||||
hashes,
|
||||
[1000, 1001, 1002, 1003],
|
||||
)
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
[start for start, _, _, _, _ in result],
|
||||
[124 * 128, 125 * 128, 126 * 128, 127 * 128],
|
||||
)
|
||||
self.assertEqual(
|
||||
[block_id for _, _, _, _, block_id in result],
|
||||
[1000, 1001, 1002, 1003],
|
||||
)
|
||||
|
||||
def test_process_tokens_token_len_shorter_than_all_blocks(self):
|
||||
hashes = ["a", "b", "c", "d"]
|
||||
# token_len=32 means only first 2 blocks valid
|
||||
result = list(self.db.process_tokens(32, hashes))
|
||||
self.assertEqual(len(result), 2)
|
||||
|
||||
def test_process_tokens_rehashes_grouped_hashes(self):
|
||||
db = ChunkedTokenDatabase([self.meta], block_size=[16], partitions=None, hash_block_size=8)
|
||||
result = list(db.process_tokens(32, ["a", "b", "c", "d"]))
|
||||
self.assertEqual(len(result), 2)
|
||||
self.assertEqual(result[0][2].chunk_hash, _expected_grouped_hash("a", "b").hex())
|
||||
self.assertEqual(len(result[0][2].chunk_hash), 64)
|
||||
|
||||
def test_key_strings_match_pool_keys(self):
|
||||
hashes = ["aaa", "bbb", "ccc"]
|
||||
pool_keys = list(self.db.process_tokens(40, hashes))
|
||||
self.assertEqual(
|
||||
list(self.db.process_token_key_strings(40, hashes)),
|
||||
[
|
||||
(start, end, key.to_string(), hash_val)
|
||||
for (start, end, key), hash_val in zip(pool_keys, hashes, strict=True)
|
||||
],
|
||||
)
|
||||
|
||||
block_ids = [5, 6]
|
||||
self.assertEqual(
|
||||
list(self.db.process_token_key_strings_with_block_ids(32, hashes, block_ids)),
|
||||
[
|
||||
(start, end, key.to_string(), hash_val, block_id)
|
||||
for (start, end, key), hash_val, block_id in zip(pool_keys[:2], hashes[:2], block_ids, strict=True)
|
||||
],
|
||||
)
|
||||
|
||||
def test_direct_keys_preserve_multigroup_layerwise_key_semantics(self):
|
||||
group_metadata = [
|
||||
KeyMetadata("llama", 0, 0, 0, 0),
|
||||
KeyMetadata("llama", 1, 0, 0, 0),
|
||||
]
|
||||
db = ChunkedTokenDatabase(group_metadata, block_size=[16, 32], partitions=None, hash_block_size=16)
|
||||
db.set_group_buffers(
|
||||
{0: [1000], 1: [2000]},
|
||||
{0: [160], 1: [320]},
|
||||
group_cache_families={0: "c1", 1: "c2"},
|
||||
group_num_layers={0: 2, 1: 2},
|
||||
)
|
||||
hashes = ["a", "b", "c", "d"]
|
||||
|
||||
pool_key_result = list(db.process_tokens(64, hashes, kv_cache_group_id=1))
|
||||
direct_key_result = list(db.process_token_key_strings(64, hashes, kv_cache_group_id=1))
|
||||
|
||||
self.assertEqual(len(pool_key_result), 1)
|
||||
self.assertEqual(
|
||||
direct_key_result[0][:3],
|
||||
(pool_key_result[0][0], pool_key_result[0][1], pool_key_result[0][2].to_string()),
|
||||
)
|
||||
layer_key = pool_key_result[0][2].split_layers(2)[1]
|
||||
self.assertIn("@group:1@cache_role:kv@cache_family:c2@layer_id:1", layer_key.to_string())
|
||||
|
||||
def test_key_strings_pre_shard_after_filtering(self):
|
||||
hashes = ["a", "b", "c", "d"]
|
||||
store_mask = [True, False, True, True]
|
||||
result = list(
|
||||
self.db.process_token_key_strings_with_block_ids(
|
||||
64,
|
||||
hashes,
|
||||
[10, 11, 12, 13],
|
||||
chunk_filter=lambda start: store_mask[start // 16],
|
||||
shard_rank=1,
|
||||
shard_size=2,
|
||||
)
|
||||
)
|
||||
self.assertEqual([(start, end, block_id) for start, end, _, _, block_id in result], [(32, 48, 12)])
|
||||
|
||||
def test_get_block_hashes_rehashes_groups(self):
|
||||
for hashes in (["a", "b", "c", "d"], [b"a", b"b", b"c", b"d"]):
|
||||
with self.subTest(hash_type=type(hashes[0])):
|
||||
result = get_block_hashes(hashes, group_block_size=32, hash_block_size=16)
|
||||
expected = [
|
||||
_expected_grouped_hash(hashes[0], hashes[1]),
|
||||
_expected_grouped_hash(hashes[2], hashes[3]),
|
||||
]
|
||||
self.assertEqual(list(result), expected)
|
||||
|
||||
def test_prepare_value(self):
|
||||
addr, size, block_id = self.db.prepare_value(0, 16, [5, 6, 7])
|
||||
self.assertEqual(block_id, 5)
|
||||
self.assertEqual(len(addr), 2)
|
||||
self.assertEqual(addr[0], 1000 + 5 * 160)
|
||||
self.assertEqual(addr[1], 2000 + 5 * 320)
|
||||
self.assertEqual(size[0], 160)
|
||||
self.assertEqual(size[1], 320)
|
||||
|
||||
def test_prepare_value_partial_block(self):
|
||||
addr, size, block_id = self.db.prepare_value(0, 8, [5])
|
||||
self.assertEqual(size[0], 80) # 160/16*8
|
||||
self.assertEqual(size[1], 160) # 320/16*8
|
||||
|
||||
def test_prepare_value_uses_block_id_override(self):
|
||||
addr, size, block_id = self.db.prepare_value(64, 80, [5], block_id=99)
|
||||
self.assertEqual(block_id, 99)
|
||||
self.assertEqual(addr[0], 1000 + 99 * 160)
|
||||
self.assertEqual(addr[1], 2000 + 99 * 320)
|
||||
self.assertEqual(size[0], 160)
|
||||
self.assertEqual(size[1], 320)
|
||||
|
||||
def test_prepare_value_layer(self):
|
||||
addr, size, block_id = self.db.prepare_value_layer(0, 16, [5, 6], layer_id=0)
|
||||
self.assertEqual(block_id, 5)
|
||||
self.assertEqual(len(addr), 2)
|
||||
# layer_id=0, entries_per_layers=2 => group_addrs[0] and group_addrs[1]
|
||||
self.assertEqual(addr[0], 1000 + 5 * 160)
|
||||
self.assertEqual(addr[1], 2000 + 5 * 320)
|
||||
|
||||
def test_decode_adaptor_prefill_pp_no_partitions(self):
|
||||
key, addr, size = self.db.decode_adaptor_prefill_pp(["k1"], [[1, 2]], [[10, 20]])
|
||||
self.assertEqual(key, ["k1"])
|
||||
|
||||
def test_decode_adaptor_prefill_pp_single_partition(self):
|
||||
db = ChunkedTokenDatabase([self.meta], [16], partitions=[4])
|
||||
key, addr, size = db.decode_adaptor_prefill_pp(["k1"], [[1, 2]], [[10, 20]])
|
||||
self.assertEqual(key, ["k1"])
|
||||
|
||||
def test_decode_adaptor_prefill_pp_multi_partition(self):
|
||||
db = ChunkedTokenDatabase([self.meta], [16], partitions=[2, 2])
|
||||
db.set_group_buffers({0: [1000, 2000]}, {0: [160, 320]})
|
||||
keys = ["k1@pp_rank:0"]
|
||||
addrs = [[1, 2, 3, 4, 5, 6, 7, 8]]
|
||||
sizes = [[10, 20, 30, 40, 50, 60, 70, 80]]
|
||||
new_keys, new_addrs, new_sizes = db.decode_adaptor_prefill_pp(keys, addrs, sizes)
|
||||
self.assertEqual(len(new_keys), 2)
|
||||
self.assertIn("@pp_rank:0", new_keys[0])
|
||||
self.assertIn("@pp_rank:1", new_keys[1])
|
||||
|
||||
|
||||
class TestLoadSpec(unittest.TestCase):
|
||||
def test_fields(self):
|
||||
spec = LoadSpec(vllm_cached_tokens=10, kvpool_cached_tokens=20, can_load=True)
|
||||
self.assertEqual(spec.vllm_cached_tokens, 10)
|
||||
self.assertEqual(spec.kvpool_cached_tokens, 20)
|
||||
self.assertTrue(spec.can_load)
|
||||
self.assertEqual(spec.token_len, 0)
|
||||
|
||||
def test_token_len_default(self):
|
||||
spec = LoadSpec(0, 0, False, token_len=128)
|
||||
self.assertEqual(spec.token_len, 128)
|
||||
|
||||
|
||||
class TestRequestTracker(unittest.TestCase):
|
||||
def test_from_new_request(self):
|
||||
new_req = MagicMock()
|
||||
new_req.req_id = "req-1"
|
||||
new_req.block_ids = [10, 20, 30]
|
||||
new_req.prompt_token_ids = list(range(100))
|
||||
|
||||
tracker = RequestTracker.from_new_request(new_req, num_tokens_to_compute=48)
|
||||
self.assertEqual(tracker.req_id, "req-1")
|
||||
self.assertEqual(tracker.token_len, 48)
|
||||
self.assertEqual(tracker.allocated_block_ids, [10, 20, 30])
|
||||
self.assertEqual(len(tracker.token_ids), 48)
|
||||
self.assertEqual(tracker.num_saved_tokens, 0)
|
||||
|
||||
def test_from_new_request_nested_block_ids(self):
|
||||
new_req = MagicMock()
|
||||
new_req.req_id = "req-2"
|
||||
new_req.block_ids = [[10, 20], [30, 40]]
|
||||
new_req.prompt_token_ids = list(range(32))
|
||||
|
||||
tracker = RequestTracker.from_new_request(new_req, num_tokens_to_compute=32)
|
||||
self.assertEqual(tracker.allocated_block_ids, [10, 20])
|
||||
|
||||
def test_update_with_list(self):
|
||||
tracker = RequestTracker(req_id="r1", token_len=16, allocated_block_ids=[1, 2])
|
||||
tracker.update([3, 4])
|
||||
self.assertEqual(tracker.allocated_block_ids, [1, 2, 3, 4])
|
||||
|
||||
def test_update_with_tuple(self):
|
||||
tracker = RequestTracker(req_id="r1", token_len=16, allocated_block_ids=[1])
|
||||
tracker.update(([5, 6], [7, 8]))
|
||||
self.assertEqual(tracker.allocated_block_ids, [1, 5, 6])
|
||||
|
||||
def test_update_with_empty(self):
|
||||
tracker = RequestTracker(req_id="r1", token_len=16, allocated_block_ids=[1])
|
||||
tracker.update([])
|
||||
self.assertEqual(tracker.allocated_block_ids, [1])
|
||||
|
||||
def test_update_invalid_type(self):
|
||||
tracker = RequestTracker(req_id="r1", token_len=16, allocated_block_ids=[1])
|
||||
with self.assertRaises(ValueError):
|
||||
tracker.update("invalid") # type: ignore[arg-type]
|
||||
|
||||
def test_update_mamba_with_tuple(self):
|
||||
tracker = RequestTracker(
|
||||
req_id="r1", token_len=16, allocated_block_ids_by_group=[[1], [2], [3], [4]], block_sizes=[16] * 4
|
||||
)
|
||||
tracker.update(([5, 6], [0, 7], [0, 8], [0, 9]))
|
||||
self.assertEqual(tracker.allocated_block_ids_by_group[0], [1, 5, 6])
|
||||
self.assertEqual(tracker.allocated_block_ids_by_group[1], [2, 0, 7])
|
||||
self.assertEqual(tracker.allocated_block_ids_by_group[2], [3, 0, 8])
|
||||
self.assertEqual(tracker.allocated_block_ids_by_group[3], [4, 0, 9])
|
||||
|
||||
def test_update_mamba_mtp_with_tuple_chunk2(self):
|
||||
tracker = RequestTracker(
|
||||
req_id="r1",
|
||||
token_len=32,
|
||||
allocated_block_ids_by_group=[
|
||||
[1, 2],
|
||||
[0, 3, 4, 5, 6],
|
||||
[0, 7, 8, 9, 10],
|
||||
[0, 11, 12, 13, 14],
|
||||
],
|
||||
mamba_group_ids=[1, 2, 3],
|
||||
num_speculative_blocks=3,
|
||||
block_sizes=[16] * 4,
|
||||
)
|
||||
|
||||
tracker.update(([15, 16], [4, 17], [8, 18], [12, 19]), 32)
|
||||
self.assertEqual(tracker.allocated_block_ids_by_group[0], [1, 2, 15, 16])
|
||||
self.assertEqual(tracker.allocated_block_ids_by_group[1], [0, 3, 0, 5, 6, 4, 17])
|
||||
self.assertEqual(tracker.allocated_block_ids_by_group[2], [0, 7, 0, 9, 10, 8, 18])
|
||||
self.assertEqual(tracker.allocated_block_ids_by_group[3], [0, 11, 0, 13, 14, 12, 19])
|
||||
|
||||
def test_update_mamba_mtp_with_tuple_chunk8(self):
|
||||
tracker = RequestTracker(
|
||||
req_id="r1",
|
||||
token_len=128,
|
||||
allocated_block_ids_by_group=[
|
||||
[1, 2, 3, 4, 5, 6, 7, 8],
|
||||
[0, 0, 0, 0, 0, 0, 0, 9, 10, 11, 12],
|
||||
[0, 0, 0, 0, 0, 0, 0, 13, 14, 15, 16],
|
||||
[0, 0, 0, 0, 0, 0, 0, 17, 18, 19, 20],
|
||||
],
|
||||
mamba_group_ids=[1, 2, 3],
|
||||
num_speculative_blocks=3,
|
||||
block_sizes=[16] * 4,
|
||||
)
|
||||
|
||||
tracker.update(
|
||||
(
|
||||
[21, 22, 23, 24, 25, 26, 27, 28],
|
||||
[0, 0, 0, 0, 10, 11, 12, 29],
|
||||
[0, 0, 0, 0, 14, 15, 16, 30],
|
||||
[0, 0, 0, 0, 18, 19, 20, 31],
|
||||
),
|
||||
128,
|
||||
)
|
||||
self.assertEqual(
|
||||
tracker.allocated_block_ids_by_group[0], [1, 2, 3, 4, 5, 6, 7, 8, 21, 22, 23, 24, 25, 26, 27, 28]
|
||||
)
|
||||
self.assertEqual(
|
||||
tracker.allocated_block_ids_by_group[1], [0, 0, 0, 0, 0, 0, 0, 9, 0, 0, 0, 0, 0, 0, 0, 10, 11, 12, 29]
|
||||
)
|
||||
self.assertEqual(
|
||||
tracker.allocated_block_ids_by_group[2], [0, 0, 0, 0, 0, 0, 0, 13, 0, 0, 0, 0, 0, 0, 0, 14, 15, 16, 30]
|
||||
)
|
||||
self.assertEqual(
|
||||
tracker.allocated_block_ids_by_group[3], [0, 0, 0, 0, 0, 0, 0, 17, 0, 0, 0, 0, 0, 0, 0, 18, 19, 20, 31]
|
||||
)
|
||||
|
||||
|
||||
class TestReqMeta(unittest.TestCase):
|
||||
def test_from_request_tracker_basic_save(self):
|
||||
tracker = RequestTracker(
|
||||
req_id="r1",
|
||||
token_len=32,
|
||||
allocated_block_ids=[0, 1],
|
||||
num_saved_tokens=0,
|
||||
token_ids=list(range(32)),
|
||||
)
|
||||
meta = ReqMeta.from_request_tracker(tracker, cache_transfer_granularity=16, block_hashes=[b"h1", b"h2"])
|
||||
self.assertIsNotNone(meta)
|
||||
self.assertEqual(meta.req_id, "r1")
|
||||
self.assertTrue(meta.can_save)
|
||||
self.assertEqual(meta.token_len_chunk, 32)
|
||||
self.assertIsNone(meta.load_spec)
|
||||
|
||||
def test_from_request_tracker_skip_save(self):
|
||||
tracker = RequestTracker(
|
||||
req_id="r1",
|
||||
token_len=32,
|
||||
allocated_block_ids=[0, 1],
|
||||
num_saved_tokens=0,
|
||||
)
|
||||
meta = ReqMeta.from_request_tracker(tracker, cache_transfer_granularity=16, skip_save=True)
|
||||
self.assertIsNone(meta)
|
||||
|
||||
def test_from_request_tracker_with_load_spec(self):
|
||||
tracker = RequestTracker(
|
||||
req_id="r1",
|
||||
token_len=32,
|
||||
allocated_block_ids=[0, 1],
|
||||
num_saved_tokens=0,
|
||||
)
|
||||
load_spec = LoadSpec(vllm_cached_tokens=0, kvpool_cached_tokens=32, can_load=True)
|
||||
meta = ReqMeta.from_request_tracker(tracker, cache_transfer_granularity=16, load_spec=load_spec, skip_save=True)
|
||||
self.assertIsNotNone(meta)
|
||||
self.assertIsNotNone(meta.load_spec)
|
||||
|
||||
def test_from_request_tracker_load_spec_cannot_load(self):
|
||||
tracker = RequestTracker(
|
||||
req_id="r1",
|
||||
token_len=32,
|
||||
allocated_block_ids=[0, 1],
|
||||
num_saved_tokens=32,
|
||||
)
|
||||
load_spec = LoadSpec(vllm_cached_tokens=0, kvpool_cached_tokens=32, can_load=False)
|
||||
meta = ReqMeta.from_request_tracker(tracker, cache_transfer_granularity=16, load_spec=load_spec, skip_save=True)
|
||||
# can_load=False => load_spec set to None in meta,
|
||||
# but skip_save+load_spec input is not None, so meta is still created
|
||||
self.assertIsNotNone(meta)
|
||||
self.assertIsNone(meta.load_spec)
|
||||
self.assertFalse(meta.can_save)
|
||||
|
||||
def test_from_request_tracker_partial_tokens_discarded(self):
|
||||
tracker = RequestTracker(
|
||||
req_id="r1",
|
||||
token_len=20,
|
||||
allocated_block_ids=[0, 1],
|
||||
num_saved_tokens=0,
|
||||
)
|
||||
meta = ReqMeta.from_request_tracker(tracker, cache_transfer_granularity=16, discard_partial_chunks=True)
|
||||
self.assertIsNotNone(meta)
|
||||
self.assertEqual(meta.token_len_chunk, 16)
|
||||
|
||||
def test_from_request_tracker_no_discard(self):
|
||||
tracker = RequestTracker(
|
||||
req_id="r1",
|
||||
token_len=20,
|
||||
allocated_block_ids=[0, 1],
|
||||
num_saved_tokens=0,
|
||||
)
|
||||
meta = ReqMeta.from_request_tracker(tracker, cache_transfer_granularity=16, discard_partial_chunks=False)
|
||||
self.assertIsNotNone(meta)
|
||||
self.assertEqual(meta.token_len_chunk, 20)
|
||||
|
||||
def test_from_request_tracker_already_saved(self):
|
||||
tracker = RequestTracker(
|
||||
req_id="r1",
|
||||
token_len=32,
|
||||
allocated_block_ids=[0, 1],
|
||||
num_saved_tokens=32,
|
||||
)
|
||||
meta = ReqMeta.from_request_tracker(tracker, cache_transfer_granularity=16)
|
||||
|
||||
# num_saved_tokens=32, chunk_boundary=ceil(33/16)*16=48 > 32
|
||||
# so skip_save, and no load_spec => None
|
||||
self.assertIsNone(meta)
|
||||
|
||||
def test_from_request_tracker_with_original_block_size(self):
|
||||
tracker = RequestTracker(
|
||||
req_id="r1",
|
||||
token_len=32,
|
||||
allocated_block_ids=[0, 1],
|
||||
num_saved_tokens=0,
|
||||
)
|
||||
# Provide block_hashes (2 full blocks for token_len=32 / granularity=16)
|
||||
# so the boundary_without_hash short-circuit does not zero out the save
|
||||
# length and skip; this exercises the original_block_size propagation.
|
||||
meta = ReqMeta.from_request_tracker(
|
||||
tracker,
|
||||
cache_transfer_granularity=16,
|
||||
original_block_size=8,
|
||||
block_hashes=[b"h0", b"h1"],
|
||||
)
|
||||
self.assertIsNotNone(meta)
|
||||
self.assertEqual(meta.original_block_size, 8)
|
||||
|
||||
|
||||
class TestAscendConnectorMetadata(unittest.TestCase):
|
||||
def test_add_request(self):
|
||||
meta = AscendConnectorMetadata(unfinished_request_ids=set(), preempted_req_ids=set())
|
||||
req = ReqMeta(
|
||||
req_id="r1",
|
||||
token_len_chunk=16,
|
||||
block_ids=[0],
|
||||
block_hashes=[],
|
||||
)
|
||||
meta.add_request(req)
|
||||
self.assertEqual(len(meta.requests), 1)
|
||||
self.assertEqual(meta.requests[0].req_id, "r1")
|
||||
|
||||
|
||||
class TestLayerMultiBlockReqMeta(unittest.TestCase):
|
||||
def test_fields(self):
|
||||
meta = LayerMultiBlockReqMeta(
|
||||
req_id="r1",
|
||||
keys=[],
|
||||
starts=[0, 16],
|
||||
ends=[16, 32],
|
||||
block_ids=[0, 1],
|
||||
layer_id=2,
|
||||
)
|
||||
self.assertEqual(meta.req_id, "r1")
|
||||
self.assertEqual(meta.layer_id, 2)
|
||||
self.assertTrue(meta.is_last_chunk)
|
||||
self.assertIsNone(meta.current_event)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
291
tests/ut/distributed/ascend_store/test_coordinator.py
Normal file
291
tests/ut/distributed/ascend_store/test_coordinator.py
Normal file
@@ -0,0 +1,291 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import unittest
|
||||
from dataclasses import dataclass, replace
|
||||
from unittest.mock import patch
|
||||
|
||||
# isort: off
|
||||
import tests.ut.distributed.ascend_store._mock_deps # noqa: F401, E402
|
||||
import torch
|
||||
from vllm.v1.kv_cache_interface import FullAttentionSpec, KVCacheGroupSpec, SlidingWindowSpec
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store import config_data
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.config_data import (
|
||||
ChunkedTokenDatabase,
|
||||
KeyMetadata,
|
||||
get_block_hashes,
|
||||
)
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.coordinator import (
|
||||
AscendStoreCoordinator,
|
||||
ExternalCachedBlockPool,
|
||||
)
|
||||
|
||||
# isort: on
|
||||
|
||||
|
||||
def _hashes(num_blocks: int) -> list[bytes]:
|
||||
return [bytes([idx % 251]) * 32 for idx in range(num_blocks)]
|
||||
|
||||
|
||||
def _full_spec(block_size: int) -> FullAttentionSpec:
|
||||
return FullAttentionSpec(
|
||||
block_size=block_size,
|
||||
num_kv_heads=1,
|
||||
head_size=1,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
|
||||
def _sliding_spec(block_size: int, sliding_window: int) -> SlidingWindowSpec:
|
||||
return SlidingWindowSpec(
|
||||
block_size=block_size,
|
||||
num_kv_heads=1,
|
||||
head_size=1,
|
||||
dtype=torch.float32,
|
||||
sliding_window=sliding_window,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _FakeCompressedSpec:
|
||||
block_size: int
|
||||
compress_ratio: int
|
||||
|
||||
def copy_with_new_block_size(self, block_size):
|
||||
return replace(self, block_size=block_size)
|
||||
|
||||
|
||||
class _FakeCompressedManager:
|
||||
@classmethod
|
||||
def find_longest_cache_hit(
|
||||
cls,
|
||||
block_hashes,
|
||||
max_length,
|
||||
kv_cache_group_ids,
|
||||
block_pool,
|
||||
kv_cache_spec,
|
||||
drop_eagle_block=False,
|
||||
alignment_tokens=16,
|
||||
**kwargs,
|
||||
):
|
||||
computed: tuple[list[object], ...] = tuple([] for _ in kv_cache_group_ids)
|
||||
logical_block_size = kv_cache_spec.block_size * kv_cache_spec.compress_ratio
|
||||
max_blocks = max_length // logical_block_size
|
||||
for block_hash in list(block_hashes)[:max_blocks]:
|
||||
cached = block_pool.get_cached_block(block_hash, kv_cache_group_ids)
|
||||
if not cached:
|
||||
break
|
||||
for blocks, block in zip(computed, cached):
|
||||
blocks.append(block)
|
||||
return computed
|
||||
|
||||
|
||||
class TestAscendStoreCoordinator(unittest.TestCase):
|
||||
def test_load_mask_grouped_hashes_are_reused_by_key_build(self):
|
||||
block_hashes = _hashes(4)
|
||||
coord = AscendStoreCoordinator(
|
||||
[KVCacheGroupSpec(["layer.0"], _full_spec(16))],
|
||||
scheduler_block_size=16,
|
||||
hash_block_size=8,
|
||||
group_block_sizes=[16],
|
||||
group_cache_families=["c1"],
|
||||
)
|
||||
db = ChunkedTokenDatabase(
|
||||
[KeyMetadata("model", 0, 0, 0, 0)],
|
||||
block_size=[16],
|
||||
partitions=None,
|
||||
hash_block_size=8,
|
||||
)
|
||||
db.cache_coordinator = coord
|
||||
grouped_hash_cache: config_data.GroupedBlockHashCache = {}
|
||||
|
||||
with patch.object(
|
||||
config_data,
|
||||
"_rehash_block_hash_group",
|
||||
wraps=config_data._rehash_block_hash_group,
|
||||
) as rehash:
|
||||
self.assertEqual(
|
||||
db.load_mask(
|
||||
block_hashes,
|
||||
32,
|
||||
grouped_hash_cache=grouped_hash_cache,
|
||||
),
|
||||
([True, True],),
|
||||
)
|
||||
self.assertEqual(rehash.call_count, 2)
|
||||
|
||||
keys = list(
|
||||
db.process_token_key_strings_with_block_ids(
|
||||
32,
|
||||
block_hashes,
|
||||
[10, 11],
|
||||
grouped_hash_cache=grouped_hash_cache,
|
||||
)
|
||||
)
|
||||
|
||||
self.assertEqual(len(keys), 2)
|
||||
self.assertEqual(rehash.call_count, 2)
|
||||
|
||||
def test_compressed_group_hits_on_effective_granularity(self):
|
||||
block_hashes = _hashes(128)
|
||||
grouped_hash = get_block_hashes(block_hashes, group_block_size=128 * 128, hash_block_size=128)[0]
|
||||
coord = AscendStoreCoordinator(
|
||||
[KVCacheGroupSpec(["layer.0"], _full_spec(128))],
|
||||
scheduler_block_size=128 * 128,
|
||||
hash_block_size=128,
|
||||
group_block_sizes=[128],
|
||||
group_cache_families=["c128"],
|
||||
)
|
||||
|
||||
_, hit_length = coord.find_longest_cache_hit(
|
||||
block_hashes,
|
||||
128 * 128,
|
||||
ExternalCachedBlockPool({(0, bytes(grouped_hash))}),
|
||||
)
|
||||
|
||||
self.assertEqual(hit_length, 128 * 128)
|
||||
|
||||
def test_compressed_spec_does_not_apply_ratio_twice(self):
|
||||
block_hashes = _hashes(128)
|
||||
grouped_hash = get_block_hashes(block_hashes, group_block_size=128 * 128, hash_block_size=128)[0]
|
||||
|
||||
with patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.coordinator._get_manager_class",
|
||||
return_value=_FakeCompressedManager,
|
||||
):
|
||||
coord = AscendStoreCoordinator(
|
||||
[KVCacheGroupSpec(["layer.0"], _FakeCompressedSpec(block_size=128, compress_ratio=128))],
|
||||
scheduler_block_size=128 * 128,
|
||||
hash_block_size=128,
|
||||
group_block_sizes=[128],
|
||||
group_cache_families=["c128"],
|
||||
)
|
||||
|
||||
_, hit_length = coord.find_longest_cache_hit(
|
||||
block_hashes,
|
||||
128 * 128,
|
||||
ExternalCachedBlockPool({(0, bytes(grouped_hash))}),
|
||||
)
|
||||
|
||||
self.assertEqual(coord.group_effective_specs[0].compress_ratio, 1)
|
||||
self.assertEqual(hit_length, 128 * 128)
|
||||
|
||||
def test_missing_required_group_returns_zero(self):
|
||||
block_hashes = _hashes(128)
|
||||
c1_exists = {(0, block_hash) for block_hash in block_hashes}
|
||||
coord = AscendStoreCoordinator(
|
||||
[
|
||||
KVCacheGroupSpec(["layer.0"], _full_spec(128)),
|
||||
KVCacheGroupSpec(["layer.1"], _full_spec(128)),
|
||||
],
|
||||
scheduler_block_size=128 * 128,
|
||||
hash_block_size=128,
|
||||
group_block_sizes=[128, 128],
|
||||
group_cache_families=["c1", "c128"],
|
||||
)
|
||||
|
||||
_, hit_length = coord.find_longest_cache_hit(
|
||||
block_hashes,
|
||||
128 * 128,
|
||||
ExternalCachedBlockPool(c1_exists),
|
||||
)
|
||||
|
||||
self.assertEqual(hit_length, 0)
|
||||
|
||||
def test_store_mask_uses_manager_reachability(self):
|
||||
coord = AscendStoreCoordinator(
|
||||
[KVCacheGroupSpec(["layer.0"], _sliding_spec(block_size=128, sliding_window=256))],
|
||||
scheduler_block_size=512,
|
||||
hash_block_size=128,
|
||||
group_block_sizes=[128],
|
||||
group_cache_families=["c1"],
|
||||
)
|
||||
|
||||
with patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.coordinator._reachable_block_mask",
|
||||
return_value=[False, False, False, True],
|
||||
):
|
||||
masks = coord.store_mask(512)
|
||||
|
||||
self.assertEqual(masks, ([False, False, False, True],))
|
||||
|
||||
def test_lookup_mask_uses_reachability_without_retention(self):
|
||||
coord = AscendStoreCoordinator(
|
||||
[KVCacheGroupSpec(["layer.0"], _sliding_spec(block_size=128, sliding_window=256))],
|
||||
scheduler_block_size=512,
|
||||
hash_block_size=128,
|
||||
group_block_sizes=[128],
|
||||
group_cache_families=["c1"],
|
||||
retention_interval=256,
|
||||
)
|
||||
with patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.coordinator._reachable_block_mask",
|
||||
return_value=[False, False, False, True],
|
||||
) as reachable:
|
||||
masks = coord.lookup_mask(512)
|
||||
|
||||
self.assertEqual(masks, ([False, False, False, True],))
|
||||
self.assertIsNone(reachable.call_args.kwargs["retention_interval"])
|
||||
|
||||
def test_store_mask_propagates_eagle_to_same_spec_siblings(self):
|
||||
calls = []
|
||||
|
||||
def fake_reachable_block_mask(*args, **kwargs):
|
||||
calls.append(kwargs["use_eagle"])
|
||||
return [True, False, True, False]
|
||||
|
||||
shared_spec = _sliding_spec(block_size=128, sliding_window=256)
|
||||
coord = AscendStoreCoordinator(
|
||||
[
|
||||
KVCacheGroupSpec(["layer.0"], shared_spec),
|
||||
KVCacheGroupSpec(["layer.mtp"], shared_spec, is_eagle_group=True),
|
||||
],
|
||||
scheduler_block_size=512,
|
||||
hash_block_size=128,
|
||||
group_block_sizes=[128, 128],
|
||||
group_cache_families=["c1", "c1"],
|
||||
)
|
||||
|
||||
with patch(
|
||||
"vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.coordinator._reachable_block_mask",
|
||||
side_effect=fake_reachable_block_mask,
|
||||
):
|
||||
masks = coord.store_mask(512)
|
||||
|
||||
self.assertEqual(calls, [True, True])
|
||||
self.assertEqual(masks, ([True, False, True, False], [True, False, True, False]))
|
||||
|
||||
def test_compressed_masks_stay_unmasked(self):
|
||||
coord = AscendStoreCoordinator(
|
||||
[KVCacheGroupSpec(["layer.0"], _sliding_spec(block_size=128, sliding_window=512))],
|
||||
scheduler_block_size=2048,
|
||||
hash_block_size=128,
|
||||
group_block_sizes=[128],
|
||||
group_cache_families=["c4"],
|
||||
)
|
||||
|
||||
self.assertEqual(coord.store_mask(2048, num_prompt_tokens=2048), ([True] * 4,))
|
||||
with patch.object(
|
||||
coord,
|
||||
"find_longest_cache_hit",
|
||||
return_value=(([False, False, False, True],), 2048),
|
||||
):
|
||||
self.assertEqual(coord.load_mask(_hashes(16), 2048), ([True] * 4,))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
664
tests/ut/distributed/ascend_store/test_kv_transfer.py
Normal file
664
tests/ut/distributed/ascend_store/test_kv_transfer.py
Normal file
@@ -0,0 +1,664 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import threading
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
# isort: off
|
||||
import tests.ut.distributed.ascend_store._mock_deps # noqa: F401, E402
|
||||
from vllm.distributed.kv_events import BlockStored
|
||||
from vllm.v1.core.kv_cache_utils import maybe_convert_block_hash
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store import config_data
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.config_data import (
|
||||
ChunkedTokenDatabase,
|
||||
KeyMetadata,
|
||||
LayerMultiBlockReqMeta,
|
||||
LayerPoolKey,
|
||||
LoadSpec,
|
||||
ReqMeta,
|
||||
)
|
||||
|
||||
# isort: on
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer import (
|
||||
KVCacheStoreLayerRecvingThread,
|
||||
KVCacheStoreLayerSendingThread,
|
||||
KVCacheStoreRecvingThread,
|
||||
KVCacheStoreSendingThread,
|
||||
KVTransferThread,
|
||||
)
|
||||
|
||||
|
||||
class FakeStore:
|
||||
def __init__(self, exists_result=None):
|
||||
self.exists_result = exists_result or []
|
||||
self.put_calls = []
|
||||
self.get_calls = []
|
||||
|
||||
def set_device(self):
|
||||
pass
|
||||
|
||||
def exists(self, keys):
|
||||
return self.exists_result[: len(keys)]
|
||||
|
||||
def put(self, keys, addrs, sizes):
|
||||
self.put_calls.append((list(keys), list(addrs), list(sizes)))
|
||||
|
||||
def get(self, keys, addrs, sizes):
|
||||
self.get_calls.append((list(keys), list(addrs), list(sizes)))
|
||||
|
||||
|
||||
class FakeTokenDatabase(ChunkedTokenDatabase):
|
||||
def __init__(self, block_size=16):
|
||||
super().__init__([KeyMetadata("m", 0, 0, 0, 0)], [block_size], None)
|
||||
self.set_group_buffers({0: [1000]}, {0: [block_size]}, {0: [1]}, group_num_layers={0: 1})
|
||||
|
||||
|
||||
class MaskedFakeTokenDatabase(FakeTokenDatabase):
|
||||
def __init__(self, block_size=16, masks=([True],)):
|
||||
super().__init__(block_size)
|
||||
self.masks = masks
|
||||
|
||||
def store_mask(self, token_len, num_prompt_tokens=None):
|
||||
return self.masks
|
||||
|
||||
def load_mask(self, block_hashes, token_len, grouped_hash_cache=None):
|
||||
return self.masks
|
||||
|
||||
def mask_allows_chunk(self, masks, kv_cache_group_id, start):
|
||||
if masks is None:
|
||||
return True
|
||||
block_idx = start // self.get_block_size(kv_cache_group_id)
|
||||
return block_idx < len(masks[kv_cache_group_id]) and masks[kv_cache_group_id][block_idx]
|
||||
|
||||
|
||||
class TestKVTransferThread(unittest.TestCase):
|
||||
def _make_thread(self, exists_result=None):
|
||||
store = FakeStore(exists_result or [])
|
||||
db = FakeTokenDatabase()
|
||||
t = KVTransferThread(
|
||||
m_store=store,
|
||||
token_database=db,
|
||||
block_size=16,
|
||||
tp_rank=0,
|
||||
dcp_size=1,
|
||||
ready_event=threading.Event(),
|
||||
name="test",
|
||||
)
|
||||
return t, store
|
||||
|
||||
def test_add_request(self):
|
||||
t, _ = self._make_thread()
|
||||
req = MagicMock()
|
||||
t.add_request(req)
|
||||
self.assertFalse(t.request_queue.empty())
|
||||
|
||||
def test_get_and_clear_finished_requests(self):
|
||||
t, _ = self._make_thread()
|
||||
t.set_finished_request("r1")
|
||||
t.set_finished_request("r2")
|
||||
finished = t.get_and_clear_finished_requests()
|
||||
self.assertEqual(finished, {"r1", "r2"})
|
||||
self.assertEqual(t.get_and_clear_finished_requests(), set())
|
||||
|
||||
def test_lookup_all_exist(self):
|
||||
t, _ = self._make_thread([1, 1, 1])
|
||||
result = t.lookup(["k1", "k2", "k3"])
|
||||
self.assertEqual(result, [True, True, True])
|
||||
|
||||
def test_lookup_partial(self):
|
||||
t, _ = self._make_thread([1, 0, 1])
|
||||
result = t.lookup(["k1", "k2", "k3"])
|
||||
self.assertEqual(result, [True, False, True])
|
||||
|
||||
def test_lookup_exception(self):
|
||||
t, store = self._make_thread()
|
||||
store.exists = MagicMock(side_effect=Exception("conn fail"))
|
||||
result = t.lookup(["k1"])
|
||||
self.assertEqual(result, [False])
|
||||
|
||||
def test_update_and_get_kv_events(self):
|
||||
t, _ = self._make_thread()
|
||||
event1 = BlockStored(
|
||||
block_hashes=["h1"],
|
||||
parent_block_hash=None,
|
||||
token_ids=[1, 2, 3],
|
||||
block_size=16,
|
||||
lora_id=None,
|
||||
medium="cpu",
|
||||
lora_name=None,
|
||||
)
|
||||
event2 = BlockStored(
|
||||
block_hashes=["h2"],
|
||||
parent_block_hash="h1",
|
||||
token_ids=[4, 5, 6],
|
||||
block_size=16,
|
||||
lora_id=None,
|
||||
medium="cpu",
|
||||
lora_name=None,
|
||||
)
|
||||
t.update_kv_event([event1, event2])
|
||||
events = t.get_kv_events()
|
||||
self.assertEqual(len(events), 2)
|
||||
# After get, events should be cleared
|
||||
self.assertEqual(len(t.get_kv_events()), 0)
|
||||
|
||||
def test_handle_request_base_noop(self):
|
||||
t, _ = self._make_thread()
|
||||
# Base class _handle_request does nothing
|
||||
t._handle_request(MagicMock())
|
||||
|
||||
|
||||
class TestKVCacheStoreSendingThread(unittest.TestCase):
|
||||
def _make_thread(self, exists_result=None, kv_role="kv_producer", enable_kv_event=False):
|
||||
store = FakeStore(exists_result or [0, 0, 0, 0])
|
||||
db = FakeTokenDatabase()
|
||||
t = KVCacheStoreSendingThread(
|
||||
m_store=store,
|
||||
token_database=db,
|
||||
block_size=16,
|
||||
tp_rank=0,
|
||||
dcp_size=1,
|
||||
put_step=1,
|
||||
kv_role=kv_role,
|
||||
ready_event=threading.Event(),
|
||||
group_uses_align_state=[False],
|
||||
enable_kv_event=enable_kv_event,
|
||||
)
|
||||
return t, store
|
||||
|
||||
def test_handle_request_puts_missing_keys(self):
|
||||
t, store = self._make_thread([1, 0, 1, 0])
|
||||
req = ReqMeta(
|
||||
req_id="r1",
|
||||
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,
|
||||
)
|
||||
t.add_stored_request("r1")
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
self.assertEqual(len(store.put_calls), 1)
|
||||
keys, _, _ = store.put_calls[0]
|
||||
self.assertEqual(len(keys), 2)
|
||||
|
||||
def test_handle_request_all_exist_no_put(self):
|
||||
t, store = self._make_thread([1, 1])
|
||||
req = ReqMeta(
|
||||
req_id="r1",
|
||||
token_len_chunk=32,
|
||||
block_ids=[0, 1],
|
||||
block_hashes=[b"h0", b"h1"], # type: ignore[arg-type]
|
||||
current_event=None,
|
||||
)
|
||||
t.add_stored_request("r1")
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
self.assertEqual(len(store.put_calls), 0)
|
||||
|
||||
def test_handle_request_not_in_stored(self):
|
||||
t, store = self._make_thread([0])
|
||||
req = ReqMeta(
|
||||
req_id="r1",
|
||||
token_len_chunk=16,
|
||||
block_ids=[0],
|
||||
block_hashes=[b"h0"], # type: ignore[arg-type]
|
||||
current_event=None,
|
||||
)
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
self.assertEqual(len(store.put_calls), 0)
|
||||
|
||||
def test_handle_request_with_kv_event(self):
|
||||
t, store = self._make_thread([1, 0, 1], enable_kv_event=True)
|
||||
req = ReqMeta(
|
||||
req_id="r1",
|
||||
token_len_chunk=48,
|
||||
block_ids=[0, 1, 2],
|
||||
block_hashes=[b"h0", b"h1", b"h2"], # type: ignore[arg-type]
|
||||
current_event=None,
|
||||
token_ids=list(range(48)),
|
||||
original_block_size=16,
|
||||
)
|
||||
t.add_stored_request("r1")
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
events = t.get_kv_events()
|
||||
self.assertEqual(len(events), 1)
|
||||
self.assertEqual(events[0].block_hashes, [maybe_convert_block_hash(b"h1")])
|
||||
self.assertEqual(events[0].parent_block_hash, maybe_convert_block_hash(b"h0"))
|
||||
|
||||
def test_save_reuses_grouped_hashes_for_kv_events(self):
|
||||
store = FakeStore([0, 0])
|
||||
db = ChunkedTokenDatabase(
|
||||
[KeyMetadata("m", 0, 0, 0, 0)],
|
||||
[16],
|
||||
None,
|
||||
hash_block_size=8,
|
||||
)
|
||||
db.set_group_buffers({0: [1000]}, {0: [16]}, {0: [1]}, group_num_layers={0: 1})
|
||||
thread = KVCacheStoreSendingThread(
|
||||
m_store=store,
|
||||
token_database=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=True,
|
||||
)
|
||||
request = ReqMeta(
|
||||
req_id="r1",
|
||||
token_len_chunk=32,
|
||||
block_ids=[0, 1],
|
||||
block_hashes=[b"h0", b"h1", b"h2", b"h3"], # type: ignore[arg-type]
|
||||
current_event=None,
|
||||
token_ids=list(range(32)),
|
||||
original_block_size=16,
|
||||
)
|
||||
thread.add_stored_request("r1")
|
||||
thread.request_queue.put(request)
|
||||
|
||||
with patch.object(
|
||||
config_data,
|
||||
"_rehash_block_hash_group",
|
||||
wraps=config_data._rehash_block_hash_group,
|
||||
) as rehash:
|
||||
thread._handle_request(request)
|
||||
|
||||
self.assertEqual(rehash.call_count, 2)
|
||||
self.assertEqual(len(thread.get_kv_events()), 2)
|
||||
|
||||
def test_handle_request_consumer_role(self):
|
||||
t, store = self._make_thread([0], kv_role="kv_consumer")
|
||||
req = ReqMeta(
|
||||
req_id="r1",
|
||||
token_len_chunk=16,
|
||||
block_ids=[0],
|
||||
block_hashes=[b"h0"], # type: ignore[arg-type]
|
||||
current_event=None,
|
||||
)
|
||||
t.add_stored_request("r1")
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
self.assertEqual(len(store.put_calls), 1)
|
||||
|
||||
def test_add_dec_delete_stored_request(self):
|
||||
t, _ = self._make_thread()
|
||||
t.add_stored_request("r1")
|
||||
t.add_stored_request("r1")
|
||||
self.assertEqual(t.stored_requests["r1"], 2)
|
||||
t.dec_stored_request("r1")
|
||||
self.assertEqual(t.stored_requests["r1"], 1)
|
||||
t.delete_finished_stored_request("r1")
|
||||
self.assertNotIn("r1", t.stored_requests)
|
||||
|
||||
def test_dec_nonexistent_request(self):
|
||||
t, _ = self._make_thread()
|
||||
t.dec_stored_request("nonexist") # should not raise
|
||||
|
||||
def test_delete_nonexistent_request(self):
|
||||
t, _ = self._make_thread()
|
||||
t.delete_finished_stored_request("nonexist") # should not raise
|
||||
|
||||
def test_handle_request_with_current_event(self):
|
||||
t, store = self._make_thread([0])
|
||||
event = MagicMock()
|
||||
req = ReqMeta(
|
||||
req_id="r1",
|
||||
token_len_chunk=16,
|
||||
block_ids=[0],
|
||||
block_hashes=[b"h0"], # type: ignore[arg-type]
|
||||
current_event=event,
|
||||
)
|
||||
t.add_stored_request("r1")
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
event.synchronize.assert_called_once()
|
||||
|
||||
def test_handle_request_dcp_size_gt_1(self):
|
||||
store = FakeStore([0, 0])
|
||||
db = FakeTokenDatabase()
|
||||
t = KVCacheStoreSendingThread(
|
||||
m_store=store,
|
||||
token_database=db,
|
||||
block_size=16,
|
||||
tp_rank=0,
|
||||
dcp_size=2,
|
||||
put_step=1,
|
||||
kv_role="kv_producer",
|
||||
ready_event=threading.Event(),
|
||||
group_uses_align_state=[False],
|
||||
)
|
||||
req = ReqMeta(
|
||||
req_id="r1",
|
||||
token_len_chunk=32,
|
||||
block_ids=[0, 1],
|
||||
block_hashes=[b"h0", b"h1"], # type: ignore[arg-type]
|
||||
current_event=None,
|
||||
)
|
||||
t.add_stored_request("r1")
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
# dcp_size > 1 means no slicing
|
||||
self.assertEqual(len(store.put_calls), 1)
|
||||
|
||||
def test_handle_request_applies_store_mask(self):
|
||||
store = FakeStore([0, 0])
|
||||
db = MaskedFakeTokenDatabase(masks=([True, False],))
|
||||
t = KVCacheStoreSendingThread(
|
||||
m_store=store,
|
||||
token_database=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],
|
||||
)
|
||||
req = ReqMeta(
|
||||
req_id="r1",
|
||||
token_len_chunk=32,
|
||||
block_ids=[0, 1],
|
||||
block_hashes=[b"h0", b"h1"], # type: ignore[arg-type]
|
||||
current_event=None,
|
||||
)
|
||||
t.add_stored_request("r1")
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
keys, _, _ = store.put_calls[0]
|
||||
self.assertEqual(len(keys), 1)
|
||||
|
||||
def test_handle_request_skips_compressed_hit_in_raw_token_domain(self):
|
||||
t, store = self._make_thread([0, 0])
|
||||
t.token_database.group_cache_families["kv"][0] = "c4"
|
||||
req = ReqMeta(
|
||||
req_id="r1",
|
||||
token_len_chunk=128,
|
||||
block_ids=[0, 1],
|
||||
block_hashes=[f"h{i}" for i in range(8)],
|
||||
load_spec=LoadSpec(
|
||||
vllm_cached_tokens=0,
|
||||
kvpool_cached_tokens=63,
|
||||
kvpool_store_skip_tokens=64,
|
||||
can_load=True,
|
||||
),
|
||||
)
|
||||
t.add_stored_request("r1")
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
keys, addrs, _ = store.put_calls[0]
|
||||
self.assertEqual(len(keys), 1)
|
||||
self.assertEqual(addrs, [[1001]])
|
||||
|
||||
def test_save_exception_cleans_queue_lifecycle(self):
|
||||
t, store = self._make_thread([0])
|
||||
store.put = MagicMock(side_effect=RuntimeError("put failed"))
|
||||
req = ReqMeta(
|
||||
req_id="r1",
|
||||
token_len_chunk=16,
|
||||
block_ids=[0],
|
||||
block_hashes=[b"h0"], # type: ignore[arg-type]
|
||||
)
|
||||
t.add_stored_request("r1")
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
self.assertEqual(t.request_queue.unfinished_tasks, 0)
|
||||
self.assertNotIn("r1", t.stored_requests)
|
||||
|
||||
|
||||
class TestKVCacheStoreRecvingThread(unittest.TestCase):
|
||||
def test_handle_request(self):
|
||||
store = FakeStore()
|
||||
db = FakeTokenDatabase()
|
||||
t = KVCacheStoreRecvingThread(
|
||||
m_store=store,
|
||||
token_database=db,
|
||||
block_size=16,
|
||||
tp_rank=0,
|
||||
dcp_size=1,
|
||||
ready_event=threading.Event(),
|
||||
invalid_block_ids=set(),
|
||||
invalid_block_ids_lock=threading.Lock(),
|
||||
)
|
||||
load_spec = LoadSpec(vllm_cached_tokens=0, kvpool_cached_tokens=32, can_load=True, token_len=32)
|
||||
req = ReqMeta(
|
||||
req_id="r1",
|
||||
token_len_chunk=32,
|
||||
block_ids=[0, 1],
|
||||
block_hashes=[b"h0", b"h1"], # type: ignore[arg-type]
|
||||
load_spec=load_spec,
|
||||
)
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
self.assertEqual(len(store.get_calls), 1)
|
||||
finished = t.get_and_clear_finished_requests()
|
||||
self.assertIn("r1", finished)
|
||||
|
||||
def test_handle_request_applies_load_mask(self):
|
||||
store = FakeStore()
|
||||
db = MaskedFakeTokenDatabase(masks=([True, False],))
|
||||
t = KVCacheStoreRecvingThread(
|
||||
m_store=store,
|
||||
token_database=db,
|
||||
block_size=16,
|
||||
tp_rank=0,
|
||||
dcp_size=1,
|
||||
ready_event=threading.Event(),
|
||||
invalid_block_ids=set(),
|
||||
invalid_block_ids_lock=threading.Lock(),
|
||||
)
|
||||
load_spec = LoadSpec(vllm_cached_tokens=0, kvpool_cached_tokens=32, can_load=True, token_len=32)
|
||||
req = ReqMeta(
|
||||
req_id="r1",
|
||||
token_len_chunk=32,
|
||||
block_ids=[0, 1],
|
||||
block_hashes=[b"h0", b"h1"], # type: ignore[arg-type]
|
||||
load_spec=load_spec,
|
||||
)
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
keys, _, _ = store.get_calls[0]
|
||||
self.assertEqual(len(keys), 1)
|
||||
|
||||
|
||||
@unittest.skip("LayerMultiBlockReqMeta API is deprecated, tests need update for LayerTransferTask")
|
||||
class TestKVCacheStoreLayerSendingThread(unittest.TestCase):
|
||||
def _make_thread(self, exists_result=None, num_layers=2):
|
||||
store = FakeStore(exists_result or [0, 0])
|
||||
db = FakeTokenDatabase()
|
||||
t = KVCacheStoreLayerSendingThread(
|
||||
m_store=store,
|
||||
token_database=db,
|
||||
block_size=16,
|
||||
tp_rank=0,
|
||||
tp_size=1,
|
||||
dcp_size=1,
|
||||
put_step=1,
|
||||
my_key_index=0,
|
||||
num_ranks_per_layer=1,
|
||||
page_size_bytes=32,
|
||||
ready_event=threading.Event(),
|
||||
num_layers=num_layers,
|
||||
layer_save_finished_events=[threading.Event() for _ in range(num_layers)],
|
||||
sync_save_events=[],
|
||||
)
|
||||
return t, store
|
||||
|
||||
def _make_layer_req(self, layer_id=0, is_last_chunk=False, num_keys=2):
|
||||
meta = KeyMetadata("m", 0, 0, 0, 0)
|
||||
keys = [LayerPoolKey(meta, f"h{i}", layer_id) for i in range(num_keys)]
|
||||
return LayerMultiBlockReqMeta(
|
||||
req_id="r1",
|
||||
keys=keys,
|
||||
starts=[i * 16 for i in range(num_keys)],
|
||||
ends=[(i + 1) * 16 for i in range(num_keys)],
|
||||
block_ids=list(range(num_keys)),
|
||||
layer_id=layer_id,
|
||||
is_last_chunk=is_last_chunk,
|
||||
current_event=None,
|
||||
token_ids=list(range(num_keys * 16)),
|
||||
original_block_size=16,
|
||||
block_hashes=[f"h{i}".encode() for i in range(num_keys)],
|
||||
)
|
||||
|
||||
def test_handle_request_puts_missing(self):
|
||||
t, store = self._make_thread([1, 0])
|
||||
req = self._make_layer_req(layer_id=0)
|
||||
t.add_stored_request(req.req_id)
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
self.assertEqual(len(store.put_calls), 1)
|
||||
keys, _, _ = store.put_calls[0]
|
||||
self.assertEqual(len(keys), 1)
|
||||
|
||||
def test_handle_request_all_exist_not_last(self):
|
||||
t, store = self._make_thread([1, 1])
|
||||
req = self._make_layer_req(layer_id=0, is_last_chunk=False)
|
||||
t.add_stored_request(req.req_id)
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
self.assertEqual(len(store.put_calls), 0)
|
||||
|
||||
def test_handle_request_all_exist_last_chunk_final_layer(self):
|
||||
t, store = self._make_thread([1, 1], num_layers=2)
|
||||
req = self._make_layer_req(layer_id=1, is_last_chunk=True)
|
||||
t.add_stored_request(req.req_id)
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
finished = t.get_and_clear_finished_requests()
|
||||
self.assertIn("r1", finished)
|
||||
|
||||
def test_handle_request_empty_keys(self):
|
||||
t, store = self._make_thread()
|
||||
_meta = KeyMetadata("m", 0, 0, 0, 0)
|
||||
req = LayerMultiBlockReqMeta(
|
||||
req_id="r1",
|
||||
keys=[],
|
||||
starts=[],
|
||||
ends=[],
|
||||
block_ids=[],
|
||||
layer_id=0,
|
||||
is_last_chunk=True,
|
||||
)
|
||||
t.add_stored_request(req.req_id)
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
finished = t.get_and_clear_finished_requests()
|
||||
self.assertNotIn("r1", finished)
|
||||
|
||||
def test_handle_request_with_current_event(self):
|
||||
t, store = self._make_thread([0])
|
||||
event = MagicMock()
|
||||
meta = KeyMetadata("m", 0, 0, 0, 0)
|
||||
req = LayerMultiBlockReqMeta(
|
||||
req_id="r1",
|
||||
keys=[LayerPoolKey(meta, "h0", 0)],
|
||||
starts=[0],
|
||||
ends=[16],
|
||||
block_ids=[0],
|
||||
layer_id=0,
|
||||
is_last_chunk=False,
|
||||
current_event=event,
|
||||
)
|
||||
t.add_stored_request(req.req_id)
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
event.synchronize.assert_called_once()
|
||||
|
||||
def test_handle_request_last_chunk_final_layer_with_missing(self):
|
||||
t, store = self._make_thread([0], num_layers=2)
|
||||
req = self._make_layer_req(layer_id=1, is_last_chunk=True, num_keys=1)
|
||||
t.add_stored_request(req.req_id)
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
finished = t.get_and_clear_finished_requests()
|
||||
self.assertIn("r1", finished)
|
||||
|
||||
def test_layerwise_kv_event_published_on_final_layer(self):
|
||||
t, store = self._make_thread([0], num_layers=2)
|
||||
req = self._make_layer_req(layer_id=1, is_last_chunk=True, num_keys=1)
|
||||
t.add_stored_request(req.req_id)
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
events = t.get_kv_events()
|
||||
self.assertEqual(len(events), 1)
|
||||
self.assertEqual(events[0].block_hashes, [maybe_convert_block_hash(b"h0")])
|
||||
self.assertEqual(events[0].token_ids, list(range(16)))
|
||||
self.assertEqual(events[0].block_size, 16)
|
||||
|
||||
def test_layerwise_kv_event_not_published_before_final_layer(self):
|
||||
t, store = self._make_thread([0], num_layers=2)
|
||||
req = self._make_layer_req(layer_id=0, is_last_chunk=False, num_keys=1)
|
||||
t.add_stored_request(req.req_id)
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
self.assertEqual(t.get_kv_events(), [])
|
||||
|
||||
def test_layerwise_kv_event_uses_missing_blocks_from_previous_layers(self):
|
||||
t, store = self._make_thread([0], num_layers=2)
|
||||
first_layer_req = self._make_layer_req(layer_id=0, is_last_chunk=True, num_keys=1)
|
||||
t.add_stored_request(first_layer_req.req_id)
|
||||
t.request_queue.put(first_layer_req)
|
||||
t._handle_request(first_layer_req)
|
||||
t.m_store.exists_result = [1]
|
||||
final_layer_req = self._make_layer_req(layer_id=1, is_last_chunk=True, num_keys=1)
|
||||
t.request_queue.put(final_layer_req)
|
||||
t._handle_request(final_layer_req)
|
||||
events = t.get_kv_events()
|
||||
self.assertEqual(len(events), 1)
|
||||
self.assertEqual(events[0].block_hashes, [maybe_convert_block_hash(b"h0")])
|
||||
|
||||
|
||||
@unittest.skip("LayerMultiBlockReqMeta API is deprecated, tests need update for LayerTransferTask")
|
||||
class TestKVCacheStoreLayerRecvingThread(unittest.TestCase):
|
||||
def test_handle_request(self):
|
||||
store = FakeStore()
|
||||
db = FakeTokenDatabase()
|
||||
get_event = threading.Event()
|
||||
t = KVCacheStoreLayerRecvingThread(
|
||||
m_store=store,
|
||||
token_database=db,
|
||||
block_size=16,
|
||||
tp_rank=0,
|
||||
dcp_size=1,
|
||||
ready_event=threading.Event(),
|
||||
get_event=get_event,
|
||||
invalid_block_ids=set(),
|
||||
invalid_block_ids_lock=threading.Lock(),
|
||||
)
|
||||
meta = KeyMetadata("m", 0, 0, 0, 0)
|
||||
req = LayerMultiBlockReqMeta(
|
||||
req_id="r1",
|
||||
keys=[LayerPoolKey(meta, "h0", 0)],
|
||||
starts=[0],
|
||||
ends=[16],
|
||||
block_ids=[0],
|
||||
layer_id=0,
|
||||
)
|
||||
t.request_queue.put(req)
|
||||
t._handle_request(req)
|
||||
self.assertEqual(len(store.get_calls), 1)
|
||||
self.assertTrue(get_event.is_set())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
1267
tests/ut/distributed/ascend_store/test_pool_scheduler.py
Normal file
1267
tests/ut/distributed/ascend_store/test_pool_scheduler.py
Normal file
File diff suppressed because it is too large
Load Diff
1545
tests/ut/distributed/ascend_store/test_pool_worker.py
Normal file
1545
tests/ut/distributed/ascend_store/test_pool_worker.py
Normal file
File diff suppressed because it is too large
Load Diff
@@ -4,8 +4,7 @@ from unittest.mock import MagicMock, patch
|
||||
from vllm.distributed.utils import StatelessProcessGroup
|
||||
|
||||
from tests.ut.base import TestBase
|
||||
from vllm_ascend.distributed.device_communicators.pyhccl import \
|
||||
PyHcclCommunicator
|
||||
from vllm_ascend.distributed.device_communicators.pyhccl import PyHcclCommunicator
|
||||
|
||||
|
||||
class MockHcclLib:
|
||||
@@ -17,7 +16,6 @@ class MockUniqueId:
|
||||
|
||||
|
||||
class TestPyHcclCommunicator(TestBase):
|
||||
|
||||
@patch.dict(os.environ, {"RANK": "0", "WORLD_SIZE": "1"})
|
||||
def test_world_size_1_return_early(self):
|
||||
comm = PyHcclCommunicator(
|
||||
@@ -29,26 +27,17 @@ class TestPyHcclCommunicator(TestBase):
|
||||
|
||||
@patch.dict(os.environ, {"RANK": "0", "WORLD_SIZE": "2"})
|
||||
def test_load_hccl_fail(self):
|
||||
comm = PyHcclCommunicator(group=StatelessProcessGroup(
|
||||
0, 2, None, None),
|
||||
device="npu:0",
|
||||
library_path="/not/exist/path/libhccl.so")
|
||||
comm = PyHcclCommunicator(
|
||||
group=StatelessProcessGroup(0, 2, None, None), device="npu:0", library_path="/not/exist/path/libhccl.so"
|
||||
)
|
||||
self.assertTrue(comm.disabled)
|
||||
|
||||
@patch(
|
||||
"vllm_ascend.distributed.device_communicators.pyhccl_wrapper.HCCLLibrary",
|
||||
MockHcclLib)
|
||||
@patch(
|
||||
"vllm_ascend.distributed.device_communicators.pyhccl_wrapper.hcclUniqueId",
|
||||
MockUniqueId)
|
||||
@patch("vllm_ascend.distributed.device_communicators.pyhccl_wrapper.HCCLLibrary", MockHcclLib)
|
||||
@patch("vllm_ascend.distributed.device_communicators.pyhccl_wrapper.hcclUniqueId", MockUniqueId)
|
||||
@patch("torch.npu.device")
|
||||
@patch("vllm_ascend.utils.current_stream",
|
||||
return_value=MagicMock(npu_stream=5678))
|
||||
@patch("vllm_ascend.utils.current_stream", return_value=MagicMock(npu_stream=5678))
|
||||
def test_stateless_group(self, *_):
|
||||
group = StatelessProcessGroup(rank=3,
|
||||
world_size=4,
|
||||
store=None,
|
||||
socket=None)
|
||||
group = StatelessProcessGroup(rank=3, world_size=4, store=None)
|
||||
|
||||
comm = PyHcclCommunicator(group=group, device=3)
|
||||
|
||||
@@ -56,21 +45,17 @@ class TestPyHcclCommunicator(TestBase):
|
||||
self.assertEqual(comm.world_size, 4)
|
||||
|
||||
@patch.dict(os.environ, {"RANK": "1", "WORLD_SIZE": "2"})
|
||||
@patch(
|
||||
"vllm_ascend.distributed.device_communicators.pyhccl_wrapper.HCCLLibrary",
|
||||
MockHcclLib)
|
||||
@patch(
|
||||
"vllm_ascend.distributed.device_communicators.pyhccl_wrapper.hcclUniqueId",
|
||||
MockUniqueId)
|
||||
@patch("vllm_ascend.distributed.device_communicators.pyhccl_wrapper.HCCLLibrary", MockHcclLib)
|
||||
@patch("vllm_ascend.distributed.device_communicators.pyhccl_wrapper.hcclUniqueId", MockUniqueId)
|
||||
@patch("torch.distributed.is_initialized", return_value=True)
|
||||
@patch("torch.distributed.get_backend", return_value="nccl")
|
||||
@patch("torch.distributed.Backend.HCCL", "hccl", create=True)
|
||||
@patch("torch.distributed.get_rank", return_value=1)
|
||||
@patch("torch.distributed.get_world_size", return_value=2)
|
||||
@patch("torch.distributed.get_process_group_ranks", return_value=[0, 1])
|
||||
@patch("torch.distributed.broadcast")
|
||||
@patch("torch.npu.device")
|
||||
@patch("vllm_ascend.utils.current_stream",
|
||||
return_value=MagicMock(npu_stream=1234))
|
||||
@patch("vllm_ascend.utils.current_stream", return_value=MagicMock(npu_stream=1234))
|
||||
def test_multi_gpu_pg_torch(
|
||||
self,
|
||||
*_,
|
||||
|
||||
@@ -5,13 +5,21 @@ from torch.distributed import ReduceOp
|
||||
|
||||
from tests.ut.base import TestBase
|
||||
from vllm_ascend.distributed.device_communicators.pyhccl_wrapper import (
|
||||
Function, HCCLLibrary, aclrtStream_t, buffer_type, hcclComm_t,
|
||||
hcclDataType_t, hcclDataTypeEnum, hcclRedOp_t, hcclRedOpTypeEnum,
|
||||
hcclResult_t, hcclUniqueId)
|
||||
Function,
|
||||
HCCLLibrary,
|
||||
aclrtStream_t,
|
||||
buffer_type,
|
||||
hcclComm_t,
|
||||
hcclDataType_t,
|
||||
hcclDataTypeEnum,
|
||||
hcclRedOp_t,
|
||||
hcclRedOpTypeEnum,
|
||||
hcclResult_t,
|
||||
hcclUniqueId,
|
||||
)
|
||||
|
||||
|
||||
class TestHcclUniqueId(TestBase):
|
||||
|
||||
def test_construct(self):
|
||||
uid = hcclUniqueId()
|
||||
uid.internal[0] = 12
|
||||
@@ -20,7 +28,6 @@ class TestHcclUniqueId(TestBase):
|
||||
|
||||
|
||||
class TestHcclDataTypeEnum(TestBase):
|
||||
|
||||
def test_torch_dtype_mapping(self):
|
||||
expected = {
|
||||
torch.int8: hcclDataTypeEnum.hcclInt8,
|
||||
@@ -35,8 +42,7 @@ class TestHcclDataTypeEnum(TestBase):
|
||||
|
||||
for torch_dtype, expected_enum in expected.items():
|
||||
with self.subTest(torch_dtype=torch_dtype):
|
||||
self.assertEqual(hcclDataTypeEnum.from_torch(torch_dtype),
|
||||
expected_enum)
|
||||
self.assertEqual(hcclDataTypeEnum.from_torch(torch_dtype), expected_enum)
|
||||
|
||||
def test_unsupported_dtype_raises(self):
|
||||
with self.assertRaises(ValueError):
|
||||
@@ -44,7 +50,6 @@ class TestHcclDataTypeEnum(TestBase):
|
||||
|
||||
|
||||
class TestHcclRedOpTypeEnum(TestBase):
|
||||
|
||||
def test_torch_reduce_op_mapping(self):
|
||||
expected = {
|
||||
ReduceOp.SUM: hcclRedOpTypeEnum.hcclSum,
|
||||
@@ -55,8 +60,7 @@ class TestHcclRedOpTypeEnum(TestBase):
|
||||
|
||||
for torch_op, expected_enum in expected.items():
|
||||
with self.subTest(torch_op=torch_op):
|
||||
self.assertEqual(hcclRedOpTypeEnum.from_torch(torch_op),
|
||||
expected_enum)
|
||||
self.assertEqual(hcclRedOpTypeEnum.from_torch(torch_op), expected_enum)
|
||||
|
||||
def test_unsupported_op_raises(self):
|
||||
unsupported_op = "NOT_EXIST"
|
||||
@@ -65,7 +69,6 @@ class TestHcclRedOpTypeEnum(TestBase):
|
||||
|
||||
|
||||
class TestFunction(TestBase):
|
||||
|
||||
def test_construct_with_valid_args(self):
|
||||
func = Function(name="foo", restype=int, argtypes=[int, str, float])
|
||||
self.assertEqual(func.name, "foo")
|
||||
@@ -74,7 +77,6 @@ class TestFunction(TestBase):
|
||||
|
||||
|
||||
class TestHCLLLibrary(TestBase):
|
||||
|
||||
def test_init_with_nonexistent_so(self):
|
||||
fake_path = "/definitely/not/exist/libhccl.so"
|
||||
with self.assertRaises(OSError):
|
||||
@@ -127,7 +129,6 @@ class TestHCLLLibrary(TestBase):
|
||||
|
||||
@patch.object(HCCLLibrary, "HCCL_CHECK")
|
||||
def test_hccl_all_reduce(self, mock_hccl_check):
|
||||
|
||||
lib = HCCLLibrary.__new__(HCCLLibrary)
|
||||
lib._funcs = {"HcclAllReduce": MagicMock(return_value=0)}
|
||||
sendbuff = buffer_type()
|
||||
@@ -138,16 +139,13 @@ class TestHCLLLibrary(TestBase):
|
||||
comm = hcclComm_t()
|
||||
stream = aclrtStream_t()
|
||||
|
||||
lib.hcclAllReduce(sendbuff, recvbuff, count, datatype, op, comm,
|
||||
stream)
|
||||
lib.hcclAllReduce(sendbuff, recvbuff, count, datatype, op, comm, stream)
|
||||
|
||||
lib._funcs["HcclAllReduce"].assert_called_once_with(
|
||||
sendbuff, recvbuff, count, datatype, op, comm, stream)
|
||||
lib._funcs["HcclAllReduce"].assert_called_once_with(sendbuff, recvbuff, count, datatype, op, comm, stream)
|
||||
mock_hccl_check.assert_called_once_with(0)
|
||||
|
||||
@patch.object(HCCLLibrary, "HCCL_CHECK")
|
||||
def test_hccl_broad_cast(self, mock_hccl_check):
|
||||
|
||||
lib = HCCLLibrary.__new__(HCCLLibrary)
|
||||
lib._funcs = {"HcclBroadcast": MagicMock(return_value=0)}
|
||||
buff = buffer_type()
|
||||
@@ -159,8 +157,7 @@ class TestHCLLLibrary(TestBase):
|
||||
|
||||
lib.hcclBroadcast(buff, count, datatype, root, comm, stream)
|
||||
|
||||
lib._funcs["HcclBroadcast"].assert_called_once_with(
|
||||
buff, count, datatype, root, comm, stream)
|
||||
lib._funcs["HcclBroadcast"].assert_called_once_with(buff, count, datatype, root, comm, stream)
|
||||
mock_hccl_check.assert_called_once_with(0)
|
||||
|
||||
@patch.object(HCCLLibrary, "HCCL_CHECK")
|
||||
|
||||
240
tests/ut/distributed/kv_transfer/test_kv_transfer_failures.py
Normal file
240
tests/ut/distributed/kv_transfer/test_kv_transfer_failures.py
Normal file
@@ -0,0 +1,240 @@
|
||||
# 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.
|
||||
#
|
||||
"""Unit tests for KV transfer failure handling in ascend_store.
|
||||
|
||||
This module tests the record_failed_blocks function which handles KV transfer
|
||||
failures by recording which blocks failed to load during the transfer process.
|
||||
"""
|
||||
|
||||
import types
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
if not hasattr(torch, "npu"):
|
||||
torch.npu = types.SimpleNamespace(Event=type("Event", (), {})) # type: ignore[attr-defined]
|
||||
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.ascend_store_connector import AscendStoreConnector
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer import record_failed_blocks
|
||||
|
||||
|
||||
class TestRecordFailedBlocks(unittest.TestCase):
|
||||
"""Test cases for the record_failed_blocks function.
|
||||
|
||||
The record_failed_blocks function takes a list of block IDs and their corresponding
|
||||
return codes from a KV transfer operation, and returns a set of block IDs that failed
|
||||
(i.e., those with non-zero return codes).
|
||||
"""
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer.logger")
|
||||
def test_all_blocks_succeed(self, mock_logger: MagicMock):
|
||||
"""Test when all blocks are transferred successfully (all return codes are 0)."""
|
||||
block_ids: list[int] = [1, 2, 3, 4, 5]
|
||||
ret_codes: list[int] = [0, 0, 0, 0, 0]
|
||||
|
||||
result = record_failed_blocks(block_ids, ret_codes)
|
||||
|
||||
self.assertEqual(result, set())
|
||||
self.assertEqual(len(result), 0)
|
||||
mock_logger.error.assert_not_called()
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer.logger")
|
||||
def test_all_blocks_fail(self, mock_logger: MagicMock):
|
||||
"""Test when all blocks fail to transfer (all return codes are non-zero)."""
|
||||
block_ids: list[int] = [1, 2, 3, 4, 5]
|
||||
ret_codes: list[int] = [1, 2, 3, 4, 5]
|
||||
|
||||
result = record_failed_blocks(block_ids, ret_codes)
|
||||
|
||||
self.assertEqual(result, {1, 2, 3, 4, 5})
|
||||
self.assertEqual(len(result), 5)
|
||||
mock_logger.error.assert_called_once()
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer.logger")
|
||||
def test_partial_blocks_fail(self, mock_logger: MagicMock):
|
||||
"""Test when some blocks fail and some succeed."""
|
||||
block_ids: list[int] = [1, 2, 3, 4, 5]
|
||||
ret_codes: list[int] = [0, 1, 0, 2, 0]
|
||||
|
||||
result = record_failed_blocks(block_ids, ret_codes)
|
||||
|
||||
self.assertEqual(result, {2, 4})
|
||||
self.assertEqual(len(result), 2)
|
||||
mock_logger.error.assert_called_once()
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer.logger")
|
||||
def test_empty_lists(self, mock_logger: MagicMock):
|
||||
"""Test with empty block_ids and ret_codes."""
|
||||
block_ids: list[int] = []
|
||||
ret_codes: list[int] = []
|
||||
|
||||
result = record_failed_blocks(block_ids, ret_codes)
|
||||
|
||||
self.assertEqual(result, set())
|
||||
mock_logger.error.assert_not_called()
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer.logger")
|
||||
def test_single_block_succeed(self, mock_logger: MagicMock):
|
||||
"""Test with a single block that succeeds."""
|
||||
block_ids: list[int] = [42]
|
||||
ret_codes: list[int] = [0]
|
||||
|
||||
result = record_failed_blocks(block_ids, ret_codes)
|
||||
|
||||
self.assertEqual(result, set())
|
||||
mock_logger.error.assert_not_called()
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer.logger")
|
||||
def test_single_block_fail(self, mock_logger: MagicMock):
|
||||
"""Test with a single block that fails."""
|
||||
block_ids: list[int] = [42]
|
||||
ret_codes: list[int] = [1]
|
||||
|
||||
result = record_failed_blocks(block_ids, ret_codes)
|
||||
|
||||
self.assertEqual(result, {42})
|
||||
mock_logger.error.assert_called_once()
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer.logger")
|
||||
def test_negative_return_codes(self, mock_logger: MagicMock):
|
||||
"""Test with negative return codes (error conditions)."""
|
||||
block_ids: list[int] = [1, 2, 3]
|
||||
ret_codes: list[int] = [0, -1, -2]
|
||||
|
||||
result = record_failed_blocks(block_ids, ret_codes)
|
||||
|
||||
self.assertEqual(result, {2, 3})
|
||||
mock_logger.error.assert_called_once()
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer.logger")
|
||||
def test_large_block_ids(self, mock_logger: MagicMock):
|
||||
"""Test with large block ID values."""
|
||||
block_ids: list[int] = [1000000, 2000000, 3000000]
|
||||
ret_codes: list[int] = [0, 1, 0]
|
||||
|
||||
result = record_failed_blocks(block_ids, ret_codes)
|
||||
|
||||
self.assertEqual(result, {2000000})
|
||||
mock_logger.error.assert_called_once()
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer.logger")
|
||||
def test_mixed_error_codes(self, mock_logger: MagicMock):
|
||||
"""Test with various non-zero error codes."""
|
||||
block_ids: list[int] = [10, 20, 30, 40, 50]
|
||||
ret_codes: list[int] = [0, -1, 100, 0, 999]
|
||||
|
||||
result = record_failed_blocks(block_ids, ret_codes)
|
||||
|
||||
self.assertEqual(result, {20, 30, 50})
|
||||
mock_logger.error.assert_called_once()
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer.logger")
|
||||
def test_logs_failed_blocks(self, mock_logger: MagicMock):
|
||||
"""Test that failed blocks are logged."""
|
||||
block_ids: list[int] = [1, 2, 3]
|
||||
ret_codes: list[int] = [0, 1, 2]
|
||||
|
||||
result = record_failed_blocks(block_ids, ret_codes)
|
||||
|
||||
self.assertEqual(result, {2, 3})
|
||||
mock_logger.error.assert_called_once()
|
||||
call_args = mock_logger.error.call_args[0]
|
||||
log_msg = call_args[0]
|
||||
self.assertIn("Failed to load blocks", log_msg)
|
||||
# The last argument is the failed blocks set
|
||||
self.assertEqual(call_args[-1], {2, 3})
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer.logger")
|
||||
def test_no_log_when_all_succeed(self, mock_logger: MagicMock):
|
||||
"""Test that no error is logged when all blocks succeed."""
|
||||
block_ids: list[int] = [1, 2, 3]
|
||||
ret_codes: list[int] = [0, 0, 0]
|
||||
|
||||
result = record_failed_blocks(block_ids, ret_codes)
|
||||
|
||||
self.assertEqual(result, set())
|
||||
mock_logger.error.assert_not_called()
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer.logger")
|
||||
def test_non_hybrid_single_block_semantics(self, mock_logger: MagicMock):
|
||||
"""Test non-hybrid callers still map one return code to one block."""
|
||||
block_ids: list[int] = [10, 11, 12]
|
||||
ret_codes: list[int] = [0, 1, 0]
|
||||
|
||||
result = record_failed_blocks(block_ids, ret_codes)
|
||||
|
||||
self.assertEqual(result, {11})
|
||||
mock_logger.error.assert_called_once()
|
||||
|
||||
|
||||
class TestRecordFailedBlocksEdgeCases(unittest.TestCase):
|
||||
"""Additional edge case tests for record_failed_blocks."""
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer.logger")
|
||||
def test_duplicate_block_ids_all_fail(self, mock_logger: MagicMock):
|
||||
"""Test with duplicate block IDs that all fail."""
|
||||
# Note: This tests the behavior with duplicates
|
||||
# The set will deduplicate, but all should be marked as failed
|
||||
block_ids: list[int] = [1, 1, 2, 2]
|
||||
ret_codes: list[int] = [1, 1, 2, 2]
|
||||
|
||||
result = record_failed_blocks(block_ids, ret_codes)
|
||||
|
||||
# Set deduplicates, so we get unique failed block IDs
|
||||
self.assertEqual(result, {1, 2})
|
||||
mock_logger.error.assert_called_once()
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer.logger")
|
||||
def test_zero_block_id_with_failure(self, mock_logger: MagicMock):
|
||||
"""Test with block ID 0 failing."""
|
||||
block_ids: list[int] = [0, 1, 2]
|
||||
ret_codes: list[int] = [1, 0, 0]
|
||||
|
||||
result = record_failed_blocks(block_ids, ret_codes)
|
||||
|
||||
self.assertEqual(result, {0})
|
||||
mock_logger.error.assert_called_once()
|
||||
|
||||
@patch("vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.kv_transfer.logger")
|
||||
def test_consecutive_failures(self, mock_logger: MagicMock):
|
||||
"""Test with consecutive block failures."""
|
||||
block_ids: list[int] = [100, 101, 102, 103, 104]
|
||||
ret_codes: list[int] = [1, 1, 1, 0, 0]
|
||||
|
||||
result = record_failed_blocks(block_ids, ret_codes)
|
||||
|
||||
self.assertEqual(result, {100, 101, 102})
|
||||
mock_logger.error.assert_called_once()
|
||||
|
||||
|
||||
class TestAscendStoreConnector(unittest.TestCase):
|
||||
"""Regression tests for connector-level load failure reporting."""
|
||||
|
||||
def test_get_block_ids_with_load_errors_forwards_to_worker(self):
|
||||
connector = AscendStoreConnector.__new__(AscendStoreConnector)
|
||||
connector.connector_worker = MagicMock()
|
||||
connector.connector_worker.get_block_ids_with_load_errors.return_value = {3, 7}
|
||||
|
||||
result = connector.get_block_ids_with_load_errors()
|
||||
|
||||
self.assertEqual(result, {3, 7})
|
||||
connector.connector_worker.get_block_ids_with_load_errors.assert_called_once_with()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
0
tests/ut/distributed/mooncake/__init__.py
Normal file
0
tests/ut/distributed/mooncake/__init__.py
Normal file
74
tests/ut/distributed/mooncake/test_mooncake_config_data.py
Normal file
74
tests/ut/distributed/mooncake/test_mooncake_config_data.py
Normal file
@@ -0,0 +1,74 @@
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
fake_engine = types.ModuleType("mooncake.engine")
|
||||
fake_engine.TransferEngine = MagicMock() # type: ignore[attr-defined]
|
||||
sys.modules["mooncake.engine"] = fake_engine
|
||||
fake_store = types.ModuleType("mooncake.store")
|
||||
fake_store.ReplicateConfig = MagicMock() # type: ignore[attr-defined]
|
||||
sys.modules["mooncake.store"] = fake_store
|
||||
|
||||
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.backend.mooncake_backend import ( # noqa: E402
|
||||
_convert_to_bytes,
|
||||
_parse_global_segment_size,
|
||||
)
|
||||
|
||||
|
||||
class TestParseGlobalSegmentSize(unittest.TestCase):
|
||||
def test_int_input(self):
|
||||
self.assertEqual(_parse_global_segment_size(1024), 1024)
|
||||
self.assertEqual(_parse_global_segment_size(0), 0)
|
||||
|
||||
def test_gb_unit(self):
|
||||
self.assertEqual(_parse_global_segment_size("2GB"), 2 * 1024**3)
|
||||
self.assertEqual(_parse_global_segment_size("1.5GB"), int(1.5 * 1024**3))
|
||||
self.assertEqual(_parse_global_segment_size(" 2 GB "), 2 * 1024**3)
|
||||
|
||||
def test_gb_unit_edge_cases(self):
|
||||
with self.assertRaises(ValueError):
|
||||
_parse_global_segment_size("GB")
|
||||
with self.assertRaises(ValueError):
|
||||
_parse_global_segment_size("abcGB")
|
||||
|
||||
def test_mb_unit(self):
|
||||
self.assertEqual(_parse_global_segment_size("512MB"), 512 * 1024**2)
|
||||
self.assertEqual(_parse_global_segment_size("0.5MB"), int(0.5 * 1024**2))
|
||||
self.assertEqual(_parse_global_segment_size("1024MB"), 1024 * 1024**2)
|
||||
|
||||
def test_kb_unit(self):
|
||||
self.assertEqual(_parse_global_segment_size("256KB"), 256 * 1024)
|
||||
self.assertEqual(_parse_global_segment_size("1.25KB"), int(1.25 * 1024))
|
||||
|
||||
def test_b_unit(self):
|
||||
self.assertEqual(_parse_global_segment_size("4096B"), 4096)
|
||||
self.assertEqual(_parse_global_segment_size("1024b"), 1024)
|
||||
|
||||
def test_no_unit(self):
|
||||
self.assertEqual(_parse_global_segment_size("2048"), 2048)
|
||||
self.assertEqual(_parse_global_segment_size("0"), 0)
|
||||
|
||||
def test_non_string_non_int_input(self):
|
||||
self.assertEqual(_parse_global_segment_size(2048.0), 2048)
|
||||
self.assertEqual(_parse_global_segment_size(True), 1)
|
||||
|
||||
with self.assertRaises(TypeError):
|
||||
_parse_global_segment_size(None)
|
||||
|
||||
with self.assertRaises(TypeError):
|
||||
_parse_global_segment_size({"size": 1024})
|
||||
|
||||
|
||||
class TestConvertToBytes(unittest.TestCase):
|
||||
def test_valid_conversion(self):
|
||||
self.assertEqual(_convert_to_bytes("10", 1, "10"), 10)
|
||||
self.assertEqual(_convert_to_bytes("1.5", 1024, "1.5KB"), int(1.5 * 1024))
|
||||
self.assertEqual(_convert_to_bytes("0", 1024**3, "0GB"), 0)
|
||||
|
||||
def test_invalid_numbers(self):
|
||||
with self.assertRaises(ValueError):
|
||||
_convert_to_bytes("abc", 1, "abc")
|
||||
|
||||
with self.assertRaises(ValueError):
|
||||
_convert_to_bytes("1.2.3", 1024, "1.2.3KB")
|
||||
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()
|
||||
@@ -4,19 +4,14 @@ from unittest.mock import MagicMock, patch
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from vllm_ascend.distributed.communicator import NPUCommunicator
|
||||
from vllm_ascend.distributed.device_communicators.npu_communicator import NPUCommunicator
|
||||
|
||||
|
||||
class TestNPUCommunicator(unittest.TestCase):
|
||||
|
||||
@patch("vllm.config.get_current_vllm_config", return_value=None)
|
||||
@patch("torch.npu.current_device", return_value=MagicMock())
|
||||
@patch("torch.npu.set_device", return_value=MagicMock())
|
||||
@patch("torch.distributed.get_process_group_ranks",
|
||||
return_value={
|
||||
0: 0,
|
||||
1: 1
|
||||
})
|
||||
@patch("torch.distributed.get_process_group_ranks", return_value={0: 0, 1: 1})
|
||||
@patch("torch.distributed.get_group_rank", return_value={0: 0, 1: 1})
|
||||
@patch("torch.distributed.is_initialized", return_value=True)
|
||||
@patch("torch.distributed.get_rank", return_value=1)
|
||||
@@ -27,15 +22,8 @@ class TestNPUCommunicator(unittest.TestCase):
|
||||
@patch("torch.distributed.get_process_group_ranks", return_value=[0, 1])
|
||||
@patch("torch.npu.device")
|
||||
def test_all_to_all_with_sizes(self, *_):
|
||||
|
||||
def patched_all_to_all(output_tensor_list,
|
||||
input_tensor_list,
|
||||
group=None,
|
||||
async_op=False):
|
||||
output_tensor_list[:] = ([
|
||||
torch.tensor([10, 20]),
|
||||
torch.tensor([50, 60])
|
||||
])
|
||||
def patched_all_to_all(output_tensor_list, input_tensor_list, group=None, async_op=False):
|
||||
output_tensor_list[:] = [torch.tensor([10, 20]), torch.tensor([50, 60])]
|
||||
|
||||
torch.distributed.all_to_all = patched_all_to_all
|
||||
|
||||
@@ -43,22 +31,17 @@ class TestNPUCommunicator(unittest.TestCase):
|
||||
gather_sizes = [2, 2]
|
||||
input_ = torch.tensor([10, 20, 30, 40])
|
||||
|
||||
comm = NPUCommunicator(cpu_group=dist.group.WORLD)
|
||||
with patch.dict(dist.distributed_c10d._world.pg_map, {dist.group.WORLD: MagicMock()}, clear=False):
|
||||
comm = NPUCommunicator(cpu_group=dist.group.WORLD)
|
||||
|
||||
output = comm.all_to_all(input_,
|
||||
scatter_sizes=scatter_sizes,
|
||||
gather_sizes=gather_sizes)
|
||||
output = comm.all_to_all(input_, scatter_sizes=scatter_sizes, gather_sizes=gather_sizes)
|
||||
|
||||
assert output.tolist() == [10, 20, 50, 60]
|
||||
|
||||
@patch("vllm.config.get_current_vllm_config", return_value=None)
|
||||
@patch("torch.npu.current_device", return_value=MagicMock())
|
||||
@patch("torch.npu.set_device", return_value=MagicMock())
|
||||
@patch("torch.distributed.get_process_group_ranks",
|
||||
return_value={
|
||||
0: 0,
|
||||
1: 1
|
||||
})
|
||||
@patch("torch.distributed.get_process_group_ranks", return_value={0: 0, 1: 1})
|
||||
@patch("torch.distributed.get_group_rank", return_value={0: 0, 1: 1})
|
||||
@patch("torch.distributed.is_initialized", return_value=True)
|
||||
@patch("torch.distributed.get_rank", return_value=1)
|
||||
@@ -69,21 +52,15 @@ class TestNPUCommunicator(unittest.TestCase):
|
||||
@patch("torch.distributed.get_process_group_ranks", return_value=[0, 1])
|
||||
@patch("torch.npu.device")
|
||||
def test_all_to_all_without_sizes(self, *_):
|
||||
|
||||
def patched_all_to_all(output_tensor_list,
|
||||
input_tensor_list,
|
||||
group=None,
|
||||
async_op=False):
|
||||
output_tensor_list[:] = ([
|
||||
torch.tensor([[10, 20]]),
|
||||
torch.tensor([[50, 60]])
|
||||
])
|
||||
def patched_all_to_all(output_tensor_list, input_tensor_list, group=None, async_op=False):
|
||||
output_tensor_list[:] = [torch.tensor([[10, 20]]), torch.tensor([[50, 60]])]
|
||||
|
||||
torch.distributed.all_to_all = patched_all_to_all
|
||||
|
||||
input_ = torch.tensor([[10, 20], [30, 40]])
|
||||
|
||||
comm = NPUCommunicator(cpu_group=dist.group.WORLD)
|
||||
output = comm.all_to_all(input_, scatter_dim=0, gather_dim=0)
|
||||
with patch.dict(dist.distributed_c10d._world.pg_map, {dist.group.WORLD: MagicMock()}, clear=False):
|
||||
comm = NPUCommunicator(cpu_group=dist.group.WORLD)
|
||||
output = comm.all_to_all(input_, scatter_dim=0, gather_dim=0)
|
||||
|
||||
assert output.tolist() == [[10, 20], [50, 60]]
|
||||
|
||||
@@ -1,48 +1,157 @@
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from vllm.config import ParallelConfig
|
||||
|
||||
from vllm_ascend.distributed.parallel_state import (
|
||||
_LMTP, _MC2, _OTP, destroy_ascend_model_parallel, get_lmhead_tp_group,
|
||||
get_mc2_group, get_otp_group, init_ascend_model_parallel)
|
||||
_FLASHCOMM2_ODP,
|
||||
_FLASHCOMM2_OTP,
|
||||
_LMTP,
|
||||
_MC2,
|
||||
_OTP,
|
||||
_P_TP,
|
||||
destroy_ascend_model_parallel,
|
||||
get_flashcomm2_odp_group,
|
||||
get_flashcomm2_otp_group,
|
||||
get_global_rank,
|
||||
get_lmhead_tp_group,
|
||||
get_mc2_group,
|
||||
get_otp_group,
|
||||
get_p_tp_group,
|
||||
init_ascend_model_parallel,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def parallel_config():
|
||||
return ParallelConfig(data_parallel_size=2,
|
||||
tensor_parallel_size=2,
|
||||
pipeline_parallel_size=2)
|
||||
return ParallelConfig(
|
||||
data_parallel_size=2,
|
||||
tensor_parallel_size=4,
|
||||
pipeline_parallel_size=2,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_distributed():
|
||||
with patch('torch.distributed.is_initialized', return_value=True), \
|
||||
patch('torch.distributed.get_world_size', return_value=8), \
|
||||
patch('torch.distributed.get_backend', return_value='nccl'), \
|
||||
patch('vllm_ascend.distributed.parallel_state.get_world_group') as mock_group:
|
||||
with (
|
||||
patch("torch.distributed.is_initialized", return_value=True),
|
||||
patch("torch.distributed.get_world_size", return_value=16),
|
||||
patch("torch.distributed.get_backend", return_value="nccl"),
|
||||
patch("vllm_ascend.distributed.parallel_state.get_world_group") as mock_group,
|
||||
patch("vllm_ascend.distributed.parallel_state.get_tp_group") as mock_tp_group,
|
||||
):
|
||||
mock_group.return_value.local_rank = 0
|
||||
mock_group.return_value.device_group = MagicMock()
|
||||
mock_tp_group.return_value.world_size = 4
|
||||
yield
|
||||
|
||||
|
||||
def test_init_ascend_model_parallel(mock_distributed, parallel_config):
|
||||
mock_ascend_config = MagicMock()
|
||||
mock_ascend_config.lmhead_tensor_parallel_size = 2
|
||||
mock_ascend_config.oproj_tensor_parallel_size = 2
|
||||
with patch('vllm_ascend.distributed.parallel_state.model_parallel_initialized', return_value=False), \
|
||||
patch('vllm_ascend.distributed.parallel_state.init_model_parallel_group'), \
|
||||
patch('vllm_ascend.distributed.parallel_state.get_ascend_config', return_value=mock_ascend_config):
|
||||
mock_ascend_config.finegrained_tp_config.lmhead_tensor_parallel_size = 2
|
||||
mock_ascend_config.finegrained_tp_config.oproj_tensor_parallel_size = 2
|
||||
mock_ascend_config.finegrained_tp_config.embedding_tensor_parallel_size = 2
|
||||
mock_ascend_config.finegrained_tp_config.mlp_tensor_parallel_size = 2
|
||||
mock_ascend_config.flashcomm2_oproj_tensor_parallel_size = 2
|
||||
mock_ascend_config.pd_tp_ratio = 2
|
||||
mock_ascend_config.num_head_replica = 0
|
||||
mock_ascend_config.pd_head_ratio = 2
|
||||
mock_ascend_config.enable_flashcomm2_parallel_size = 2
|
||||
mock_ascend_config.enable_context_parallel = False
|
||||
mock_vllm_config = MagicMock()
|
||||
mock_vllm_config.kv_transfer_config.is_kv_producer = True
|
||||
with (
|
||||
patch("vllm_ascend.distributed.parallel_state.model_parallel_initialized", return_value=False),
|
||||
patch("vllm_ascend.distributed.parallel_state.init_model_parallel_group"),
|
||||
patch("vllm_ascend.distributed.parallel_state.get_current_vllm_config", return_value=mock_vllm_config),
|
||||
patch("vllm_ascend.distributed.parallel_state.get_ascend_config", return_value=mock_ascend_config),
|
||||
patch("vllm_ascend.utils.get_ascend_config", return_value=mock_ascend_config),
|
||||
):
|
||||
init_ascend_model_parallel(parallel_config)
|
||||
|
||||
mc2_group = get_mc2_group()
|
||||
lmheadtp_group = get_lmhead_tp_group()
|
||||
otp_group = get_otp_group()
|
||||
flashcomm2_otp_group = get_flashcomm2_otp_group()
|
||||
flashcomm2_odp_group = get_flashcomm2_odp_group()
|
||||
p_tp_group = get_p_tp_group()
|
||||
assert mc2_group is not None
|
||||
assert otp_group is not None
|
||||
assert flashcomm2_otp_group is not None
|
||||
assert flashcomm2_odp_group is not None
|
||||
assert lmheadtp_group is not None
|
||||
assert p_tp_group is not None
|
||||
|
||||
destroy_ascend_model_parallel()
|
||||
assert _MC2 is None
|
||||
assert _LMTP is None
|
||||
assert _OTP is None
|
||||
assert _FLASHCOMM2_OTP is None
|
||||
assert _FLASHCOMM2_ODP is None
|
||||
assert _P_TP is None
|
||||
|
||||
|
||||
def _build_parallel_config(
|
||||
tensor_parallel_size=1,
|
||||
pipeline_parallel_size=1,
|
||||
prefill_context_parallel_size=1,
|
||||
data_parallel_index=0,
|
||||
):
|
||||
return SimpleNamespace(
|
||||
tensor_parallel_size=tensor_parallel_size,
|
||||
pipeline_parallel_size=pipeline_parallel_size,
|
||||
prefill_context_parallel_size=prefill_context_parallel_size,
|
||||
data_parallel_index=data_parallel_index,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"parallel_config_kwargs, rank_in_group, expected",
|
||||
[
|
||||
# No parallelism at all (single card): replica_size == 1.
|
||||
(dict(tensor_parallel_size=1), 0, 0),
|
||||
# TP only: rank_in_group is the local rank within the single replica.
|
||||
(dict(tensor_parallel_size=4), 0, 0),
|
||||
(dict(tensor_parallel_size=4), 3, 3),
|
||||
# Dense DP: world group spans one replica, rank_in_group is local and
|
||||
# data_parallel_index supplies the DP offset.
|
||||
(dict(tensor_parallel_size=4, data_parallel_index=0), 2, 2),
|
||||
(dict(tensor_parallel_size=4, data_parallel_index=1), 2, 6),
|
||||
# MoE DP / external_launcher: world group spans all DP ranks, so
|
||||
# rank_in_group is already global; the modulo strips the DP offset and
|
||||
# data_parallel_index re-adds it (result equals rank_in_group).
|
||||
(dict(tensor_parallel_size=4, data_parallel_index=1), 6, 6),
|
||||
(dict(tensor_parallel_size=4, data_parallel_index=1), 7, 7),
|
||||
# TP * PP * prefill-CP all contribute to replica_size; DCP/EP do not.
|
||||
(dict(tensor_parallel_size=2, pipeline_parallel_size=2, data_parallel_index=1), 1, 5),
|
||||
(
|
||||
dict(
|
||||
tensor_parallel_size=2, pipeline_parallel_size=2, prefill_context_parallel_size=2, data_parallel_index=1
|
||||
),
|
||||
3,
|
||||
11,
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_get_global_rank(parallel_config_kwargs, rank_in_group, expected):
|
||||
parallel_config = _build_parallel_config(**parallel_config_kwargs)
|
||||
with patch("vllm_ascend.distributed.parallel_state.get_world_group") as mock_group:
|
||||
mock_group.return_value.rank_in_group = rank_in_group
|
||||
assert get_global_rank(parallel_config) == expected
|
||||
|
||||
|
||||
def test_get_global_rank_defaults_to_current_config():
|
||||
parallel_config = _build_parallel_config(tensor_parallel_size=4, data_parallel_index=1)
|
||||
mock_vllm_config = MagicMock()
|
||||
mock_vllm_config.parallel_config = parallel_config
|
||||
with (
|
||||
patch(
|
||||
"vllm_ascend.distributed.parallel_state.get_current_vllm_config",
|
||||
return_value=mock_vllm_config,
|
||||
),
|
||||
patch("vllm_ascend.distributed.parallel_state.get_world_group") as mock_group,
|
||||
):
|
||||
mock_group.return_value.rank_in_group = 3
|
||||
# data_parallel_index(1) * replica_size(4) + 3 == 7
|
||||
assert get_global_rank() == 7
|
||||
|
||||
0
tests/ut/distributed/weight_transfer/__init__.py
Normal file
0
tests/ut/distributed/weight_transfer/__init__.py
Normal file
159
tests/ut/distributed/weight_transfer/test_npu_ipc_engine.py
Normal file
159
tests/ut/distributed/weight_transfer/test_npu_ipc_engine.py
Normal file
@@ -0,0 +1,159 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
"""Regression tests for the NPU IPC weight transfer engine.
|
||||
|
||||
These cover two bugs that broke ``examples/rl/rlhf_http_npu_ipc.py``:
|
||||
|
||||
1. ``NPUIPCWeightTransferEngine.__init__`` did not accept the ``model``
|
||||
argument that ``WeightTransferEngineFactory.create_engine`` passes,
|
||||
raising ``TypeError: __init__() takes 3 positional arguments but 4
|
||||
were given`` at engine construction.
|
||||
2. ``receive_weights`` / ``packed_npu_ipc_consumer`` unpacked the stored
|
||||
IPC handle as ``func, args`` even though the producer stored only the
|
||||
``reduce_tensor`` *args*, raising ``ValueError: too many values to
|
||||
unpack (expected 2)``. Aligned with upstream vLLM's CUDA IPC engine:
|
||||
the producer stores args only and the consumer rebuilds with the
|
||||
well-known ``rebuild_npu_tensor``.
|
||||
"""
|
||||
|
||||
import inspect
|
||||
import sys
|
||||
import types
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from vllm_ascend.distributed.weight_transfer import npu_ipc_engine
|
||||
from vllm_ascend.distributed.weight_transfer.npu_ipc_engine import (
|
||||
NPUIPCWeightTransferEngine,
|
||||
)
|
||||
|
||||
_MODULE = "vllm_ascend.distributed.weight_transfer.npu_ipc_engine"
|
||||
|
||||
|
||||
def _patch_rebuild_npu_tensor(rebuild_func):
|
||||
"""Install a fake ``torch_npu.multiprocessing.reductions`` module.
|
||||
|
||||
The engine imports ``rebuild_npu_tensor`` lazily from ``torch_npu``,
|
||||
which is only a stub on CPU CI runners, so provide a fake submodule.
|
||||
"""
|
||||
fake_mod = types.ModuleType("torch_npu.multiprocessing.reductions")
|
||||
fake_mod.rebuild_npu_tensor = rebuild_func # type: ignore[attr-defined]
|
||||
return patch.dict(
|
||||
sys.modules,
|
||||
{
|
||||
"torch_npu.multiprocessing": types.ModuleType("torch_npu.multiprocessing"),
|
||||
"torch_npu.multiprocessing.reductions": fake_mod,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_init_accepts_model_argument():
|
||||
"""Bug 1: __init__ must accept the optional ``model`` argument."""
|
||||
params = inspect.signature(NPUIPCWeightTransferEngine.__init__).parameters
|
||||
assert "model" in params
|
||||
|
||||
|
||||
def test_init_passes_model_to_super():
|
||||
"""Bug 1: the ``model`` argument must be forwarded to the base engine."""
|
||||
captured = {}
|
||||
|
||||
def fake_init(self, config, parallel_config, model=None):
|
||||
captured["args"] = (config, parallel_config, model)
|
||||
|
||||
with patch.object(npu_ipc_engine.WeightTransferEngine, "__init__", fake_init):
|
||||
NPUIPCWeightTransferEngine("config", "parallel_config", "model")
|
||||
|
||||
assert captured["args"] == ("config", "parallel_config", "model")
|
||||
|
||||
|
||||
def test_unpacked_send_stores_reduce_tensor_args_only():
|
||||
"""Bug 2 (producer): the handle stores only the ``reduce_tensor`` args.
|
||||
|
||||
This matches upstream vLLM's CUDA IPC engine, which drops the rebuild
|
||||
func and relies on the consumer using the well-known rebuild function.
|
||||
"""
|
||||
npu_uuid = "node-0"
|
||||
|
||||
rebuild_args = (None, None, None, None, None, None, 999, None)
|
||||
fake_reduce = MagicMock(return_value=("rebuild_func_sentinel", rebuild_args))
|
||||
|
||||
captured = {}
|
||||
|
||||
def send_mode(update_info):
|
||||
captured["update_info"] = update_info
|
||||
|
||||
trainer_args = MagicMock()
|
||||
trainer_args.send_mode = send_mode
|
||||
trainer_args.packed = False
|
||||
|
||||
iterator = iter([("model.weight", torch.zeros(3))])
|
||||
|
||||
with patch(f"{_MODULE}.reduce_tensor", fake_reduce):
|
||||
NPUIPCWeightTransferEngine._send_unpacked(iterator, trainer_args, npu_uuid)
|
||||
|
||||
update_info = captured["update_info"]
|
||||
assert isinstance(update_info.ipc_handles, list)
|
||||
stored = update_info.ipc_handles[0][npu_uuid]
|
||||
# Only the args tuple is stored, not a (func, args) pair.
|
||||
assert stored == rebuild_args
|
||||
|
||||
|
||||
def test_receive_weights_rebuilds_with_rebuild_npu_tensor():
|
||||
"""Bug 2 (consumer): receive_weights rebuilds via ``rebuild_npu_tensor``.
|
||||
|
||||
Verifies the args-only handle is consumed without unpacking errors and
|
||||
that the receiver's device index is written into the rebuild args.
|
||||
"""
|
||||
npu_uuid = "node-0"
|
||||
device_index = 0
|
||||
|
||||
rebuilt_weight = torch.tensor([1.0, 2.0, 3.0])
|
||||
seen = {}
|
||||
|
||||
def fake_rebuild(*args):
|
||||
seen["args"] = args
|
||||
return rebuilt_weight
|
||||
|
||||
# Sender stores 999 at index 6; the receiver must overwrite it.
|
||||
rebuild_args = (None, None, None, None, None, None, 999, None)
|
||||
|
||||
update_info = NPUIPCWeightTransferEngine.update_info_cls(
|
||||
names=["model.weight"],
|
||||
dtype_names=["float32"],
|
||||
shapes=[[3]],
|
||||
ipc_handles=[{npu_uuid: rebuild_args}],
|
||||
packed=False,
|
||||
)
|
||||
|
||||
engine = object.__new__(NPUIPCWeightTransferEngine)
|
||||
received = {}
|
||||
|
||||
def load_weights(weights):
|
||||
received["weights"] = weights
|
||||
|
||||
with (
|
||||
_patch_rebuild_npu_tensor(fake_rebuild),
|
||||
patch(f"{_MODULE}.npu_generate_uuid", return_value=npu_uuid),
|
||||
patch("torch.accelerator.current_device_index", return_value=device_index),
|
||||
):
|
||||
engine.receive_weights(update_info, load_weights)
|
||||
|
||||
assert received["weights"][0][0] == "model.weight"
|
||||
assert torch.equal(received["weights"][0][1], rebuilt_weight)
|
||||
# Index 6 (device index) overwritten with the receiver's device.
|
||||
assert seen["args"][6] == device_index
|
||||
Reference in New Issue
Block a user