init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

View 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.

View 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()

View 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()

View 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()

View 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()

View 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()

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

View File

@@ -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,
*_,

View File

@@ -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")

View 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()

View 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")

View 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()

View File

@@ -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]]

View File

@@ -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

View 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