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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions tests/e2e/pull_request/one_card/test_gdn_layerwise_kv.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,10 @@ def chunk_attention(**kwargs):
patch("vllm_ascend.attention.utils.has_kv_transfer_group", return_value=True),
patch("vllm_ascend.attention.utils.is_v1_kv_transfer_group", return_value=True),
patch("vllm_ascend.attention.utils.get_kv_transfer_group", return_value=connector),
# NPU-only side effects live in the compiled forward on device; stub them
# so the inductor graph stays fullgraph (no graph break).
patch("vllm_ascend.ops.gdn.wait_for_kv_layer_from_connector", lambda *a, **k: None),
patch("vllm_ascend.ops.gdn.record_attention_compute_start", lambda: None),
):
eager_output = _run_gdn_forward(model, hidden_states, output)
torch.testing.assert_close(eager_output, hidden_states + 1)
Expand Down
1 change: 1 addition & 0 deletions tests/ut/distributed/ascend_store/_mock_deps.py
Original file line number Diff line number Diff line change
Expand Up @@ -99,6 +99,7 @@
"vllm.v1.outputs",
"vllm.v1.request",
"vllm.v1.serial_utils",
"vllm.v1.worker",
]
if _MOCK_VLLM_DEPS:
for _mod_name in _vllm_mock_modules:
Expand Down
67 changes: 67 additions & 0 deletions tests/ut/distributed/ascend_store/test_ascend_store_connector.py
Original file line number Diff line number Diff line change
Expand Up @@ -392,6 +392,73 @@ def test_layerwise_worker_paths(self):
connector.wait_for_layer_load("layer_0")
mock_worker_cls.return_value.wait_for_layer_load.assert_called_once()

def test_mamba_state_copy_runs_after_layer_load(self):
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole

call_order = []
with (
patch.object(self.connector_mod, "KVPoolWorker") as mock_worker_cls,
patch.object(self.connector_mod, "LookupKeyServer"),
patch.object(
self.connector_mod.mamba_utils,
"do_mamba_copy_block_for_layer",
side_effect=lambda *_: call_order.append("copy"),
create=True,
),
patch.object(
self.connector_mod.mamba_utils,
"prepare_mamba_copy_by_layer",
create=True,
) as prepare_copy,
patch.object(
self.connector_mod.mamba_utils,
"finish_mamba_copy_by_layer",
create=True,
) as finish_copy,
):
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
mock_worker_cls.return_value.wait_for_layer_load.side_effect = lambda: call_order.append("load")

connector = self.connector_mod.AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.WORKER,
kv_cache_config=None,
)
copy_bufs = MagicMock()
self.assertTrue(connector.prepare_mamba_state_copy(copy_bufs))

connector.wait_for_layer_load("layers.0.linear_attn")
connector.finish_mamba_state_copy()

self.assertEqual(call_order, ["load", "copy"])
prepare_copy.assert_called_once_with(copy_bufs)
finish_copy.assert_called_once_with(copy_bufs)
self.assertIsNone(connector._mamba_copy_bufs)

def test_non_layerwise_connector_keeps_batched_mamba_copy(self):
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole

with (
patch.object(self.connector_mod, "KVPoolWorker"),
patch.object(self.connector_mod, "LookupKeyServer"),
):
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": False}
config.parallel_config.rank = 0
connector = self.connector_mod.AscendStoreConnector(
vllm_config=config,
role=KVConnectorRole.WORKER,
kv_cache_config=None,
)

self.assertFalse(connector.prepare_mamba_state_copy(MagicMock()))


if __name__ == "__main__":
unittest.main()
23 changes: 22 additions & 1 deletion tests/ut/kv_offload/test_ascend_multi_connector.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
"""Tests for Ascend-specific MultiConnector allocation fan-out."""

from types import SimpleNamespace
from unittest.mock import MagicMock
from unittest.mock import MagicMock, patch

import pytest

Expand Down Expand Up @@ -125,3 +125,24 @@ def test_layerwise_reuse_without_sink_keeps_provider_layer_entry_wait():
connector.wait_for_layer_load("model.layers.7.self_attn")

assert call_order == ["provider", "sibling"]


def test_mamba_state_copy_runs_after_all_connector_loads():
call_order = []
first = SimpleNamespace(wait_for_layer_load=MagicMock(side_effect=lambda *_: call_order.append("first-load")))
second = SimpleNamespace(wait_for_layer_load=MagicMock(side_effect=lambda *_: call_order.append("second-load")))
connector = AscendMultiConnector.__new__(AscendMultiConnector)
connector._connectors = [first, second]
connector._layerwise_slot_release_providers = []
connector._non_slot_release_connectors = [first, second]
connector._external_slot_release_sink_configured = False
connector._mamba_copy_bufs = object()

with patch(
"vllm_ascend.distributed.kv_transfer.ascend_multi_connector.mamba_utils.do_mamba_copy_block_for_layer",
side_effect=lambda *_: call_order.append("copy"),
create=True,
):
connector.wait_for_layer_load("model.layers.7.linear_attn")

assert call_order == ["first-load", "second-load", "copy"]
4 changes: 4 additions & 0 deletions tests/ut/ops/test_gdn_layerwise_kv.py
Original file line number Diff line number Diff line change
Expand Up @@ -183,6 +183,10 @@ def chunk_attention(**kwargs):
patch("vllm_ascend.attention.utils.has_kv_transfer_group", return_value=True),
patch("vllm_ascend.attention.utils.is_v1_kv_transfer_group", return_value=True),
patch("vllm_ascend.attention.utils.get_kv_transfer_group", return_value=connector),
# NPU-only side effects live in the compiled forward on device; stub them
# so the CPU inductor graph stays fullgraph (no graph break).
patch("vllm_ascend.ops.gdn.wait_for_kv_layer_from_connector", lambda *a, **k: None),
patch("vllm_ascend.ops.gdn.record_attention_compute_start", lambda: None),
):
eager_output = _run_gdn_forward(model, hidden_states, output)
torch.testing.assert_close(eager_output, hidden_states + 1)
Expand Down
62 changes: 62 additions & 0 deletions tests/ut/patch/worker/test_patch_mamba_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -146,3 +146,65 @@ def collect_metadata(copy_buffers, *_args):
assert collect.call_args.args[4:7] == (63, 64, 0)
stage.assert_called_once_with(copy_bufs)
assert mamba_state_idx["req"] == 64


def test_layerwise_mamba_copy_is_grouped_by_layer():
"""Per-layer scheduling: prepare stages all layers once; each layer's
do_mamba_copy_block_for_layer consumes only its own slice; finish validates
all layers executed."""
from vllm_ascend.patch.worker import patch_mamba_utils as pm

bufs = SimpleNamespace(
src_ptrs=CpuGpuBuffer(8, dtype=torch.int64, device=torch.device("cpu"), pin_memory=False),
dst_ptrs=CpuGpuBuffer(8, dtype=torch.int64, device=torch.device("cpu"), pin_memory=False),
sizes=CpuGpuBuffer(8, dtype=torch.int32, device=torch.device("cpu"), pin_memory=False),
offset=0,
_layer_copy_metadata={
"layers.0.linear_attn": ([11, 12], [21, 22], [31, 32]),
"layers.1.linear_attn": ([13, 14], [23, 24], [33, 34]),
},
_layer_copy_slices={},
_layer_copy_staged=False,
_layer_tensor_copy_pairs={},
_tensor_copy_pairs=[],
)

copy_calls = []
orig = pm._batch_memcpy_triton
pm._batch_memcpy_triton = lambda s, d, z: copy_calls.append((list(s), list(d), list(z)))
try:
pm.prepare_mamba_copy_by_layer(bufs)
assert bufs._layer_copy_staged is True
assert bufs.src_ptrs.np[:4].tolist() == [11, 12, 13, 14]

pm.do_mamba_copy_block_for_layer(bufs, "layers.0.linear_attn")
pm.do_mamba_copy_block_for_layer(bufs, "layers.1.linear_attn")

assert copy_calls == [
([11, 12], [21, 22], [31, 32]),
([13, 14], [23, 24], [33, 34]),
]

pm.finish_mamba_copy_by_layer(bufs)
assert bufs.offset == 0
finally:
pm._batch_memcpy_triton = orig

# finish must raise if a layer was never executed
bufs2 = SimpleNamespace(
src_ptrs=CpuGpuBuffer(8, dtype=torch.int64, device=torch.device("cpu"), pin_memory=False),
dst_ptrs=CpuGpuBuffer(8, dtype=torch.int64, device=torch.device("cpu"), pin_memory=False),
sizes=CpuGpuBuffer(8, dtype=torch.int32, device=torch.device("cpu"), pin_memory=False),
offset=0,
_layer_copy_metadata={"layers.2.linear_attn": ([15], [25], [35])},
_layer_copy_slices={},
_layer_copy_staged=False,
_layer_tensor_copy_pairs={},
_tensor_copy_pairs=[],
)
try:
pm.finish_mamba_copy_by_layer(bufs2)
raised = False
except RuntimeError:
raised = True
assert raised, "finish must raise when a loaded layer never executed its copy"
4 changes: 2 additions & 2 deletions vllm_ascend/attention/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -477,7 +477,7 @@ def wait_for_kv_layer_from_connector(layer_name: str):

forward_context: ForwardContext = get_forward_context()
attn_metadata = forward_context.attn_metadata
if attn_metadata is None:
if attn_metadata is None or not connector.has_connector_metadata():

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

To prevent potential AttributeError or TypeError if connector is None or does not implement has_connector_metadata, we should defensively check if connector is not None and has the attribute before calling it.

Suggested change
if attn_metadata is None or not connector.has_connector_metadata():
if attn_metadata is None or connector is None or not getattr(connector, "has_connector_metadata", lambda: False)():

return
# TODO: assert ascendMetadata
connector.wait_for_layer_load(layer_name)
Expand All @@ -494,7 +494,7 @@ def maybe_save_kv_layer_to_connector(

forward_context: ForwardContext = get_forward_context()
attn_metadata = forward_context.attn_metadata
if attn_metadata is None:
if attn_metadata is None or not connector.has_connector_metadata():

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

To prevent potential AttributeError or TypeError if connector is None or does not implement has_connector_metadata, we should defensively check if connector is not None and has the attribute before calling it.

Suggested change
if attn_metadata is None or not connector.has_connector_metadata():
if attn_metadata is None or connector is None or not getattr(connector, "has_connector_metadata", lambda: False)():

return
# TODO: assert ascendMetadata
connector.save_kv_layer(layer_name, kv_cache_layer, attn_metadata)
Expand Down
30 changes: 30 additions & 0 deletions vllm_ascend/distributed/kv_transfer/ascend_multi_connector.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
supports_hma,
)
from vllm.distributed.kv_transfer.kv_connector.v1.multi_connector import MultiConnector
from vllm.v1.worker import mamba_utils

if TYPE_CHECKING:
from vllm.config import VllmConfig
Expand All @@ -27,6 +28,15 @@ def __init__(self, vllm_config: "VllmConfig", role: KVConnectorRole, kv_cache_co
"HMA should not be enabled unless all sub-connectors support it"
)
self._configure_layerwise_reuse_completion()
self._mamba_copy_bufs = None
self.requires_mamba_state_copy_after_layer_load = any(
getattr(
connector,
"requires_mamba_state_copy_after_layer_load",
False,
)
for connector in self._connectors
)

def _configure_layerwise_reuse_completion(self) -> None:
# Producers that report when a shared physical KV slot is safe to reuse.
Expand Down Expand Up @@ -67,6 +77,26 @@ def wait_for_layer_load(self, layer_name: str) -> None:
connectors = [*self._layerwise_slot_release_providers, *self._non_slot_release_connectors]
for connector in connectors:
connector.wait_for_layer_load(layer_name)
if (copy_bufs := getattr(self, "_mamba_copy_bufs", None)) is not None:
mamba_utils.do_mamba_copy_block_for_layer(
copy_bufs,
layer_name,
)

def prepare_mamba_state_copy(self, copy_bufs) -> bool:
if not self.requires_mamba_state_copy_after_layer_load:
return False
mamba_utils.prepare_mamba_copy_by_layer(copy_bufs)
self._mamba_copy_bufs = copy_bufs
return True

def finish_mamba_state_copy(self) -> None:
if self._mamba_copy_bufs is None:
return
try:
mamba_utils.finish_mamba_copy_by_layer(self._mamba_copy_bufs)
finally:
self._mamba_copy_bufs = None

def save_kv_layer(
self,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
from vllm.v1.outputs import KVConnectorOutput
from vllm.v1.request import Request
from vllm.v1.serial_utils import MsgpackDecoder
from vllm.v1.worker import mamba_utils

from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.metadata import AscendStoreKVConnectorWorkerMetadata
from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.pool_scheduler import (
Expand Down Expand Up @@ -102,6 +103,8 @@ def __init__(self, vllm_config: VllmConfig, role: KVConnectorRole, kv_cache_conf
self._kv_cache_events: AscendStoreKVEvents | None = None

self._current_step_has_real_forward = False
self._mamba_copy_bufs = None
self.requires_mamba_state_copy_after_layer_load = self.use_layerwise

if role == KVConnectorRole.SCHEDULER:
assert kv_cache_config is not None
Expand Down Expand Up @@ -209,6 +212,7 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]):

def start_load_kv(self, forward_context: "ForwardContext", **kwargs) -> None:
assert self.connector_worker is not None
self._mamba_copy_bufs = None
metadata = self._get_connector_metadata()
self._current_step_has_real_forward = forward_context is not None
logger.debug(
Expand All @@ -230,6 +234,26 @@ def wait_for_layer_load(self, layer_name: str) -> None:
if not self.use_layerwise:
return
self.connector_worker.wait_for_layer_load()
if self._mamba_copy_bufs is not None:
mamba_utils.do_mamba_copy_block_for_layer(
self._mamba_copy_bufs,
layer_name,
)

def prepare_mamba_state_copy(self, copy_bufs) -> bool:
if not self.requires_mamba_state_copy_after_layer_load:
return False
mamba_utils.prepare_mamba_copy_by_layer(copy_bufs)
self._mamba_copy_bufs = copy_bufs
return True

def finish_mamba_state_copy(self) -> None:
if self._mamba_copy_bufs is None:
return
try:
mamba_utils.finish_mamba_copy_by_layer(self._mamba_copy_bufs)
finally:
self._mamba_copy_bufs = None

def save_kv_layer(
self, layer_name: str, kv_layer: torch.Tensor, attn_metadata: "AttentionMetadata", **kwargs
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -499,11 +499,21 @@ def get_num_new_matched_tokens(

store_skip_tokens = num_external_hit_tokens
if self.use_layerwise and self.use_eagle:
# TODO(lf): Support loading the trailing block as dirty data.
num_external_hit_tokens = max(
num_computed_tokens,
num_external_hit_tokens - self.lcm_block_size,
)
# Keep the draft model's recomputation zone intact: the
# generation-point hidden states must be freshly computed, and the
# local prefix-cache path already drops its trailing block
# (drop_eagle_block). Only trim the external hit when it reaches
# into the prompt's final granularity block, so that (local +
# external) never covers the last block whose KV the engine will
# rewrite during MTP draft/verify steps. Partial hits that stop on
# an interior block boundary carry a valid mamba state snapshot
# at that boundary and can be loaded as-is.
hit_reaches_final_block = num_external_hit_tokens > (request.num_tokens - self.lcm_block_size)
if hit_reaches_final_block:
num_external_hit_tokens = max(
num_computed_tokens,
num_external_hit_tokens - self.lcm_block_size,
)
if num_external_hit_tokens == request.num_tokens:
num_external_hit_tokens -= 1

Expand Down Expand Up @@ -889,6 +899,11 @@ def touch_sending_mamba_blocks(self, req_meta: ReqMeta):
"""
if not self.use_hybrid or len(self.mamba_group_ids) == 0 or not req_meta.can_save:
return
# Layerwise transfer frees mamba blocks layer by layer on its own
# completion path (see KVCacheStoreSendingThread); bulk-touching them
# here would double-reference the block pool and leak blocks.
if self.use_layerwise:
return
using_event_id = self.get_sending_event_id()
req_meta.event_id = using_event_id
current_step_sending: list[int] = []
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1045,6 +1045,16 @@ def _process_load_for_layer_batch(
cached_tokens = request.load_spec.kvpool_cached_tokens
if not getattr(self, "use_eagle", False) and request.load_spec.kvpool_store_skip_tokens is not None:
cached_tokens = request.load_spec.kvpool_store_skip_tokens
if (
getattr(self, "use_eagle", False)
and request.load_spec.kvpool_cached_tokens == request.target_token_len - 1
):
# Full-hit path: the trailing block is recomputed and will be
# re-stored by the normal save path, so never skip it here.
logger.debug(
"Reqid: %s full-hit tail recompute path, tail block will be re-stored",
request.req_id,
)
group_block_hashes = get_block_hashes(
request.block_hashes,
block_size,
Expand Down Expand Up @@ -1330,6 +1340,16 @@ def _prepare_load_gvas(self, requests: list[ReqMeta]) -> None:
cached_tokens = request.load_spec.kvpool_cached_tokens
if not getattr(self, "use_eagle", False) and request.load_spec.kvpool_store_skip_tokens is not None:
cached_tokens = request.load_spec.kvpool_store_skip_tokens
if (
getattr(self, "use_eagle", False)
and request.load_spec.kvpool_cached_tokens == request.target_token_len - 1
):
# Full-hit path: the trailing block is recomputed and will be
# re-stored by the normal save path, so never skip it here.
logger.debug(
"Reqid: %s full-hit tail recompute path, tail block will be re-stored",
request.req_id,
)
block_hashes = request.block_hashes

all_group_load_gvas: list[np.ndarray] = []
Expand Down
Loading
Loading