diff --git a/tests/e2e/pull_request/one_card/test_gdn_layerwise_kv.py b/tests/e2e/pull_request/one_card/test_gdn_layerwise_kv.py index 9cf4c22a70dc..a1e83a3c4cbe 100644 --- a/tests/e2e/pull_request/one_card/test_gdn_layerwise_kv.py +++ b/tests/e2e/pull_request/one_card/test_gdn_layerwise_kv.py @@ -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) diff --git a/tests/ut/distributed/ascend_store/_mock_deps.py b/tests/ut/distributed/ascend_store/_mock_deps.py index aa963f74434f..898af4847080 100644 --- a/tests/ut/distributed/ascend_store/_mock_deps.py +++ b/tests/ut/distributed/ascend_store/_mock_deps.py @@ -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: diff --git a/tests/ut/distributed/ascend_store/test_ascend_store_connector.py b/tests/ut/distributed/ascend_store/test_ascend_store_connector.py index ed118ada19fd..f4bbefe9b289 100644 --- a/tests/ut/distributed/ascend_store/test_ascend_store_connector.py +++ b/tests/ut/distributed/ascend_store/test_ascend_store_connector.py @@ -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() diff --git a/tests/ut/kv_offload/test_ascend_multi_connector.py b/tests/ut/kv_offload/test_ascend_multi_connector.py index 7bb5a4bb590f..7e54c3ba41f7 100644 --- a/tests/ut/kv_offload/test_ascend_multi_connector.py +++ b/tests/ut/kv_offload/test_ascend_multi_connector.py @@ -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 @@ -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"] diff --git a/tests/ut/ops/test_gdn_layerwise_kv.py b/tests/ut/ops/test_gdn_layerwise_kv.py index 4b7b5de3ff86..6e482f170725 100644 --- a/tests/ut/ops/test_gdn_layerwise_kv.py +++ b/tests/ut/ops/test_gdn_layerwise_kv.py @@ -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) diff --git a/tests/ut/patch/worker/test_patch_mamba_utils.py b/tests/ut/patch/worker/test_patch_mamba_utils.py index cf47c2151c19..da0dbdc70353 100644 --- a/tests/ut/patch/worker/test_patch_mamba_utils.py +++ b/tests/ut/patch/worker/test_patch_mamba_utils.py @@ -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" diff --git a/vllm_ascend/attention/utils.py b/vllm_ascend/attention/utils.py index d6e772779b55..42317dab2058 100644 --- a/vllm_ascend/attention/utils.py +++ b/vllm_ascend/attention/utils.py @@ -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(): return # TODO: assert ascendMetadata connector.wait_for_layer_load(layer_name) @@ -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(): return # TODO: assert ascendMetadata connector.save_kv_layer(layer_name, kv_cache_layer, attn_metadata) diff --git a/vllm_ascend/distributed/kv_transfer/ascend_multi_connector.py b/vllm_ascend/distributed/kv_transfer/ascend_multi_connector.py index 2f95319ee385..02978561e612 100644 --- a/vllm_ascend/distributed/kv_transfer/ascend_multi_connector.py +++ b/vllm_ascend/distributed/kv_transfer/ascend_multi_connector.py @@ -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 @@ -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. @@ -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, diff --git a/vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/ascend_store_connector.py b/vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/ascend_store_connector.py index 5350cdf2b504..fda4929db6b3 100644 --- a/vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/ascend_store_connector.py +++ b/vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/ascend_store_connector.py @@ -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 ( @@ -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 @@ -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( @@ -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 diff --git a/vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/pool_scheduler.py b/vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/pool_scheduler.py index 566c6783f13e..543e34f39b93 100644 --- a/vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/pool_scheduler.py +++ b/vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/pool_scheduler.py @@ -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 @@ -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] = [] diff --git a/vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/pool_worker.py b/vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/pool_worker.py index 90ada9e31086..e2319005f125 100644 --- a/vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/pool_worker.py +++ b/vllm_ascend/distributed/kv_transfer/kv_pool/ascend_store/pool_worker.py @@ -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, @@ -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] = [] diff --git a/vllm_ascend/ops/gdn.py b/vllm_ascend/ops/gdn.py index 956144c52f7b..798353b3e083 100644 --- a/vllm_ascend/ops/gdn.py +++ b/vllm_ascend/ops/gdn.py @@ -28,8 +28,12 @@ from vllm.v1.attention.backends.gdn_attn import GDNAttentionMetadata from vllm.v1.attention.backends.utils import PAD_SLOT_ID -from vllm_ascend.attention.utils import maybe_save_kv_layer_to_connector +from vllm_ascend.attention.utils import ( + maybe_save_kv_layer_to_connector, + wait_for_kv_layer_from_connector, +) from vllm_ascend.device.device_op import DeviceOperator +from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.attention_fence import record_attention_compute_start from vllm_ascend.ops.gdn_attn_builder import AscendGDNAttentionBackend from vllm_ascend.ops.triton.fla.chunk import chunk_gated_delta_rule from vllm_ascend.ops.triton.fla.fused_qkvzba_split_reshape import fused_qkvzba_split_reshape_cat @@ -273,6 +277,14 @@ def _forward_core( # V1 profile run return + # Layerwise KV pool hooks must stay inside the custom op body: the + # forward() caller region is traced by Dynamo in fullgraph mode, and + # these side effects (thread locks, connector waits) would break the + # graph. Waiting here still orders the deferred mamba state copy and + # the layer load before conv/attention kernels touch mamba state. + wait_for_kv_layer_from_connector(self.prefix) + record_attention_compute_start() + assert isinstance(attn_metadata, dict) attn_metadata = attn_metadata[self.prefix] assert isinstance(attn_metadata, GDNAttentionMetadata) diff --git a/vllm_ascend/patch/__init__.py b/vllm_ascend/patch/__init__.py index e153064661c0..c8c2ae6c608b 100644 --- a/vllm_ascend/patch/__init__.py +++ b/vllm_ascend/patch/__init__.py @@ -739,7 +739,9 @@ # 2. preprocess_mamba copy the state of previous step to the last block before kv transfer load # How: # 1. patch to remove assert -# 2. path to only collect copy metadata in preprocess_mamba(and do actual copy after kv transfer load). +# 2. patch to collect per-layer copy metadata in preprocess_mamba. With +# layerwise KV transfer, copy each layer's state only after that +# layer finishes loading; otherwise keep the original batched copy. # Future Plan: # Remove this patch when: # vLLM itself supports kv transfer for mamba diff --git a/vllm_ascend/patch/worker/patch_mamba_utils.py b/vllm_ascend/patch/worker/patch_mamba_utils.py index c6a9bd2d8cb5..da891f6be8f4 100644 --- a/vllm_ascend/patch/worker/patch_mamba_utils.py +++ b/vllm_ascend/patch/worker/patch_mamba_utils.py @@ -254,11 +254,140 @@ def _batch_memcpy_unavailable(src_ptrs, dst_ptrs, sizes): ) +def _reset_layerwise_copy_meta(copy_bufs: mamba_utils.MambaCopyBuffers) -> None: + copy_bufs._layer_copy_metadata = {} + copy_bufs._layer_copy_slices = {} + copy_bufs._layer_copy_staged = False + copy_bufs._layer_tensor_copy_pairs = {} + copy_bufs._tensor_copy_pairs = [] + + +def _collect_mamba_copy_meta_with_layers( + copy_bufs: mamba_utils.MambaCopyBuffers, + kv_cache_config, + mamba_state_copy_funcs, + mamba_group_ids: list[int], + src_block_idx: int, + dest_block_idx: int, + accept_token_bias: int, + req_state, + forward_context: dict[str, Any], +) -> None: + """Collect copy metadata grouped per layer (triton pointer path). + + Same iteration order and buffer rows as upstream collect_mamba_copy_meta, but + records for every layer which rows belong to it so a layerwise KV load can + consume its slice as soon as that layer finishes loading. + """ + if src_block_idx == dest_block_idx and accept_token_bias == 0: + return + + layer_copy_metadata = copy_bufs._layer_copy_metadata + src_ptrs_np = copy_bufs.src_ptrs.np + dst_ptrs_np = copy_bufs.dst_ptrs.np + sizes_np = copy_bufs.sizes.np + offset = copy_bufs.offset + + for mamba_group_id in mamba_group_ids: + block_ids = req_state.block_ids[mamba_group_id] + dest_block_id = block_ids[dest_block_idx] + layer_names = kv_cache_config.kv_cache_groups[mamba_group_id].layer_names + for layer_name in layer_names: + attention = forward_context[layer_name] + kv_caches: list[torch.Tensor] = attention.kv_cache + layer_meta = layer_copy_metadata.setdefault(layer_name, ([], [], [])) + for state, state_copy_func in zip(kv_caches, mamba_state_copy_funcs): + copy_spec = state_copy_func(state, block_ids, src_block_idx, accept_token_bias + 1) + src_ptr = copy_spec.start_addr + dst_ptr = state[dest_block_id].data_ptr() + size = copy_spec.num_elements * state.element_size() + src_ptrs_np[offset] = src_ptr + dst_ptrs_np[offset] = dst_ptr + sizes_np[offset] = size + layer_meta[0].append(src_ptr) + layer_meta[1].append(dst_ptr) + layer_meta[2].append(size) + offset += 1 + + copy_bufs.offset = offset + + +def prepare_mamba_copy_by_layer(copy_bufs: mamba_utils.MambaCopyBuffers) -> None: + """Stage all layer copy metadata before layerwise model execution. + + ``CpuGpuBuffer.copy_to_gpu`` is non-blocking. Repacking the same pinned CPU + buffers for every layer can therefore overwrite metadata that an earlier + H2D copy has not consumed yet. Pack all layers once and keep their GPU + slices immutable for the duration of the forward pass. + """ + layer_copy_metadata = getattr(copy_bufs, "_layer_copy_metadata", None) + if not layer_copy_metadata or getattr(copy_bufs, "_layer_copy_staged", False): + return + + offset = 0 + layer_copy_slices = {} + for layer_name, (src_ptrs, dst_ptrs, sizes) in layer_copy_metadata.items(): + num_copies = len(src_ptrs) + end = offset + num_copies + copy_bufs.src_ptrs.np[offset:end] = src_ptrs + copy_bufs.dst_ptrs.np[offset:end] = dst_ptrs + copy_bufs.sizes.np[offset:end] = sizes + layer_copy_slices[layer_name] = (offset, end) + offset = end + + copy_bufs.offset = offset + copy_bufs._layer_copy_slices = layer_copy_slices + if offset: + copy_bufs.src_ptrs.copy_to_gpu(offset) + copy_bufs.dst_ptrs.copy_to_gpu(offset) + copy_bufs.sizes.copy_to_gpu(offset) + copy_bufs._layer_copy_staged = True + + +def do_mamba_copy_block_for_layer(copy_bufs: mamba_utils.MambaCopyBuffers, layer_name: str) -> None: + """Copy one layer's running state after its layerwise load completes.""" + layer_copy_metadata = getattr(copy_bufs, "_layer_copy_metadata", None) + if not layer_copy_metadata: + return + metadata = layer_copy_metadata.pop(layer_name, None) + if metadata is None: + return + if not _can_launch_triton_batch_memcpy(): + for src_state, dst_state in copy_bufs._layer_tensor_copy_pairs.pop(layer_name, []): + dst_state.copy_(src_state.clone()) + return + layer_slice = copy_bufs._layer_copy_slices.pop(layer_name, None) + if layer_slice is None: + return + start, end = layer_slice + if start == end: + return + _batch_memcpy_triton( + copy_bufs.src_ptrs.gpu[start:end], + copy_bufs.dst_ptrs.gpu[start:end], + copy_bufs.sizes.gpu[start:end], + ) + + +def finish_mamba_copy_by_layer(copy_bufs: mamba_utils.MambaCopyBuffers) -> None: + remaining = getattr(copy_bufs, "_layer_copy_metadata", {}) + if remaining: + raise RuntimeError(f"Mamba state copy was not executed for loaded layers: {sorted(remaining)}") + copy_bufs._layer_copy_slices = {} + copy_bufs._layer_copy_staged = False + copy_bufs._layer_tensor_copy_pairs = {} + copy_bufs._tensor_copy_pairs = [] + copy_bufs.offset = 0 + + if _can_launch_triton_batch_memcpy(): mamba_utils.batch_memcpy_kernel = batch_memcpy_kernel mamba_utils.batch_memcpy = _batch_memcpy_triton mamba_utils.do_mamba_copy_block = _do_mamba_copy_block_npu mamba_utils.postprocess_mamba_fused_kernel = postprocess_mamba_fused_kernel + # Layerwise KV pool: collect copy metadata grouped per layer so each + # layer's state copy can run right after its layer load finishes. + mamba_utils.collect_mamba_copy_meta = _collect_mamba_copy_meta_with_layers else: mamba_utils.batch_memcpy = _batch_memcpy_unavailable mamba_utils.collect_mamba_copy_meta = _collect_mamba_copy_meta_torch @@ -317,6 +446,7 @@ def preprocess_mamba( mamba_state_idx.pop(req_id, None) copy_bufs.offset = 0 + _reset_layerwise_copy_meta(copy_bufs) for i, req_id in enumerate(input_batch.req_ids): req_state = requests[req_id] num_scheduled_tokens = scheduler_output.num_scheduled_tokens[req_id] @@ -365,11 +495,18 @@ def preprocess_mamba( ) input_batch.num_accepted_tokens_cpu[i] = 1 if _can_launch_triton_batch_memcpy(): - # Only stage the pointer table here. This runs inside the existing - # input-preparation event scope, so its pinned CPU buffers cannot be - # reused until the asynchronous H2D copies finish. The state copy must - # remain after KV transfer and is executed by do_mamba_copy_block(). - _stage_mamba_copy_metadata(copy_bufs) + # Layerwise mode: stage per-layer pointer tables here (same protected + # input-preparation scope), then each layer's state copy is executed by + # do_mamba_copy_block_for_layer() right after that layer's KV load + # finishes. Non-layerwise callers still get the bulk staging below. + if getattr(copy_bufs, "_layer_copy_metadata", None): + prepare_mamba_copy_by_layer(copy_bufs) + else: + _stage_mamba_copy_metadata(copy_bufs) mamba_utils.preprocess_mamba = preprocess_mamba + +mamba_utils.prepare_mamba_copy_by_layer = prepare_mamba_copy_by_layer +mamba_utils.do_mamba_copy_block_for_layer = do_mamba_copy_block_for_layer +mamba_utils.finish_mamba_copy_by_layer = finish_mamba_copy_by_layer diff --git a/vllm_ascend/worker/model_runner_v1.py b/vllm_ascend/worker/model_runner_v1.py index 2afbd40e93ed..d0a3ec10c7b3 100644 --- a/vllm_ascend/worker/model_runner_v1.py +++ b/vllm_ascend/worker/model_runner_v1.py @@ -2373,11 +2373,33 @@ def execute_model( ), ) as kv_connector_output, ): + # Mamba state copy must run AFTER the KV transfer load finishes, + # otherwise the copy would race with in-flight layerwise loads and + # read half-loaded state. With a layerwise-capable connector we + # defer the copy: prepare_mamba_state_copy stages per-layer copy + # metadata here, and each layer's copy is executed right after its + # own KV load completes (wait_for_layer_load), overlapping the copy + # with the remaining layers' loads. Non-layerwise connectors keep + # the batched copy after all loads finish (do_mamba_copy_block). + mamba_copy_connector = None if self.cache_config.mamba_cache_mode == "align": - mamba_utils.do_mamba_copy_block(preprocess_bufs) + if has_kv_transfer_group(): + connector = get_kv_transfer_group() + prepare_mamba_state_copy = getattr( + connector, + "prepare_mamba_state_copy", + None, + ) + if callable(prepare_mamba_state_copy) and prepare_mamba_state_copy(preprocess_bufs): + mamba_copy_connector = connector + if mamba_copy_connector is None: + mamba_utils.do_mamba_copy_block(preprocess_bufs) hidden_states = self._model_forward( num_tokens_padded, input_ids, positions, intermediate_tensors, inputs_embeds, **model_kwargs ) + # Verify every scheduled layer executed its deferred copy. + if self.cache_config.mamba_cache_mode == "align" and mamba_copy_connector is not None: + mamba_copy_connector.finish_mamba_state_copy() with record_function_or_nullcontext("post process"): aux_hidden_states = None if self.use_aux_hidden_state_outputs: