From a81a0811314e23bc6b2baa69f34247aacea71e23 Mon Sep 17 00:00:00 2001 From: FujitsuPolycom <87842395+FujitsuPolycom@users.noreply.github.com> Date: Mon, 20 Jul 2026 22:50:25 -0500 Subject: [PATCH] feat(dcp): bulk-prefetch shared CKV layers over CE Group three eligible shared-layer selected-record prefetches into one copy-engine exchange while preserving direct, partial, and overflow fallbacks. Design informed by koush's sparse CKV prefetch work. Co-authored-by: OpenAI Codex Signed-off-by: FujitsuPolycom <87842395+FujitsuPolycom@users.noreply.github.com> --- .../test_b12x_ckv_prefetch_policy.py | 42 +- .../test_b12x_sparse_ckv_decode_policy.py | 474 ++++++++++++++++++ vllm/envs.py | 37 ++ .../attention/backends/mla/b12x_mla_sparse.py | 212 ++++++-- .../backends/mla/b12x_sparse_ckv_decode.py | 110 +++- 5 files changed, 833 insertions(+), 42 deletions(-) diff --git a/tests/v1/attention/test_b12x_ckv_prefetch_policy.py b/tests/v1/attention/test_b12x_ckv_prefetch_policy.py index d9e22918c1c2..98c5a5e39451 100644 --- a/tests/v1/attention/test_b12x_ckv_prefetch_policy.py +++ b/tests/v1/attention/test_b12x_ckv_prefetch_policy.py @@ -94,13 +94,40 @@ def wait(self): self.wait_calls += 1 +class _FakeStream: + def __init__(self): + self.waited_events = [] + + def wait_event(self, event): + self.waited_events.append(event) + + +def test_ckv_bulk_prefetch_ticket_waits_once_for_three_shared_layers(): + registry = _CKVPrefetchStateRegistry() + state = registry.for_workspace(torch.empty(16, dtype=torch.uint8)) + event = _FakeEvent() + stream = _FakeStream() + + ticket = state.register_pending_group({3: 0, 4: 1, 5: 2}, event) + pending = [state.pop_pending_layer(layer) for layer in (3, 4, 5)] + + assert all(item is not None for item in pending) + assert all(item.ticket is ticket for item in pending if item is not None) + for item in pending: + assert item is not None + item.ticket.wait_on_stream_once(stream) + + assert stream.waited_events == [event] + assert state.pending_layers == {} + + def test_ckv_prefetch_incomplete_step_recovers_with_stream_wait(): registry = _CKVPrefetchStateRegistry() state = registry.for_workspace(torch.empty(16, dtype=torch.uint8)) cache = torch.empty(0) event = _FakeEvent() state.register_cache(3, cache) - state.pending_layers[3] = (event, 0) + state.register_pending_group({3: 0}, event) state.last_layer_idx = 2 state.enter_layer(0) @@ -151,7 +178,7 @@ def test_ckv_prefetch_target_and_draft_lifecycles_are_isolated(monkeypatch): draft_ring = draft_state.get_ckv_workspace(64) target_state.register_cache(1, target_cache) - target_state.pending_layers[1] = (target_event, 1) + target_pending = target_state.register_pending_group({1: 1}, target_event) draft_state.register_cache(1, draft_cache) assert target_state is not draft_state @@ -159,7 +186,8 @@ def test_ckv_prefetch_target_and_draft_lifecycles_are_isolated(monkeypatch): draft_ring.untyped_storage().data_ptr() ) assert target_state.layer_caches[1] is target_cache - assert target_state.pending_layers[1] == (target_event, 1) + assert target_state.pending_layers[1].ticket is target_pending + assert target_state.pending_layers[1].buf_idx == 1 assert draft_state.layer_caches[1] is draft_cache assert draft_state.pending_layers == {} @@ -226,7 +254,7 @@ def test_ckv_prefetch_ring_resize_drains_pending_generation(): assert state.ckv_workspace_generation == 1 event = _FakeEvent() - state.pending_layers[1] = (event, 0) + state.register_pending_group({1: 0}, event) resized_ring = state.get_ckv_workspace(128) assert resized_ring is not first_ring @@ -243,7 +271,7 @@ def test_ckv_prefetch_workspace_identity_invalidates_changed_geometry(): changed_geometry = workspace_buffer[:8] old_state = registry.for_workspace(workspace_buffer) event = _FakeEvent() - old_state.pending_layers[1] = (event, 0) + old_state.register_pending_group({1: 0}, event) assert registry.for_workspace(same_geometry) is old_state @@ -269,7 +297,7 @@ def test_ckv_prefetch_workspace_identity_tracks_manager_resize(monkeypatch): draft_state = registry.for_workspace(draft_workspace, 0, draft_cache) draft_state.register_cache(0, draft_cache) event = _FakeEvent() - first_state.pending_layers[1] = (event, 0) + first_state.register_pending_group({1: 0}, event) (resized_workspace,) = manager.get_simultaneous(((257,), torch.uint8)) resized_state = registry.for_workspace(resized_workspace, 0, cache) @@ -290,7 +318,7 @@ def test_ckv_prefetch_registry_retires_released_profile_workspace(): profile_workspace = torch.empty(16, dtype=torch.uint8) state = registry.for_workspace(profile_workspace) event = _FakeEvent() - state.pending_layers[1] = (event, 0) + state.register_pending_group({1: 0}, event) del profile_workspace gc.collect() diff --git a/tests/v1/attention/test_b12x_sparse_ckv_decode_policy.py b/tests/v1/attention/test_b12x_sparse_ckv_decode_policy.py index 465cf777fb9b..2498c865a86c 100644 --- a/tests/v1/attention/test_b12x_sparse_ckv_decode_policy.py +++ b/tests/v1/attention/test_b12x_sparse_ckv_decode_policy.py @@ -1,11 +1,16 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import sys +from contextlib import nullcontext +from types import ModuleType + import pytest import torch from vllm.distributed import parallel_state from vllm.v1.attention.backends.mla import b12x_sparse_ckv_decode +from vllm.v1.attention.backends.mla.b12x_mla_sparse import B12xMLASparseImpl from vllm.v1.attention.backends.mla.b12x_sparse_ckv_decode import ( build_dense_union_remap, dense_union_remap_reference, @@ -13,6 +18,7 @@ plan_sparse_ckv_decode, sparse_decode_batch_eligible, sparse_decode_prefetch_targets, + try_exchange_sparse_decode_shared_layers, ) @@ -28,6 +34,474 @@ def _layout(*, dcp: int = 4, requests: int = 8, pool: int = 0): ) +def test_bulk_shared_prefetch_exchanges_three_layers_once_without_new_workspace(): + class Exchange: + def __init__(self): + self.calls = [] + + def exchange_layers(self, **kwargs): + self.calls.append(kwargs) + + layout = _layout() + exchange = Exchange() + records = (object(), object(), object()) + outputs = (object(), object(), object()) + transport_indices = object() + workspace_geometry = (layout.workspace_slots, layout.workspace_bytes) + + used_bulk = try_exchange_sparse_decode_shared_layers( + enabled=True, + transport="ce", + exchange=exchange, + records_by_layer=records, + local_indices_by_destination=transport_indices, + outputs_by_layer=outputs, + active_records=layout.active_records(1), + pool_records=layout.pool_records, + ) + + assert used_bulk + assert exchange.calls == [ + { + "records_by_layer": records, + "local_indices_by_destination": transport_indices, + "outputs_by_layer": outputs, + } + ] + assert (layout.workspace_slots, layout.workspace_bytes) == workspace_geometry + + +def test_bulk_shared_prefetch_is_default_off(monkeypatch): + name = "VLLM_B12X_MLA_SPARSE_DECODE_BULK_PREFETCH" + monkeypatch.delenv(name, raising=False) + + assert b12x_sparse_ckv_decode.envs.environment_variables[name]() is False + + monkeypatch.setenv(name, "1") + assert b12x_sparse_ckv_decode.envs.environment_variables[name]() is True + + +@pytest.mark.parametrize("value", ["enabled", "1 ", ""]) +def test_bulk_shared_prefetch_rejects_invalid_value(monkeypatch, value): + name = "VLLM_B12X_MLA_SPARSE_DECODE_BULK_PREFETCH" + monkeypatch.setenv(name, value) + + with pytest.raises(ValueError, match="Valid options"): + b12x_sparse_ckv_decode.envs.environment_variables[name]() + + +def test_sparse_decode_transport_rejects_unknown_value(monkeypatch): + name = "VLLM_B12X_MLA_SPARSE_DECODE_TRANSPORT" + monkeypatch.setenv(name, "unknown") + + with pytest.raises(ValueError, match="Valid options"): + b12x_sparse_ckv_decode.envs.environment_variables[name]() + + +@pytest.mark.parametrize( + ("enabled", "transport"), + [(False, "ce"), (True, "direct")], +) +def test_bulk_shared_prefetch_preserves_per_layer_fallback(enabled, transport): + class Exchange: + pass + + used_bulk = try_exchange_sparse_decode_shared_layers( + enabled=enabled, + transport=transport, + exchange=Exchange(), + records_by_layer=(object(), object(), object()), + local_indices_by_destination=object(), + outputs_by_layer=(object(), object(), object()), + active_records=8192, + pool_records=65536, + ) + + assert not used_bulk + + +def test_bulk_shared_prefetch_old_b12x_falls_back_cleanly(): + class LegacyCopyExchange: + pass + + used_bulk = try_exchange_sparse_decode_shared_layers( + enabled=True, + transport="ce", + exchange=LegacyCopyExchange(), + records_by_layer=(object(), object(), object()), + local_indices_by_destination=object(), + outputs_by_layer=(object(), object(), object()), + active_records=8192, + pool_records=65536, + ) + + assert not used_bulk + + +def test_bulk_shared_prefetch_requires_exact_depth_three(): + class Exchange: + def exchange_layers(self, **kwargs): + raise AssertionError("partial groups must use the per-layer path") + + used_bulk = try_exchange_sparse_decode_shared_layers( + enabled=True, + transport="ce", + exchange=Exchange(), + records_by_layer=(object(), object()), + local_indices_by_destination=object(), + outputs_by_layer=(object(), object()), + active_records=8192, + pool_records=65536, + ) + + assert not used_bulk + + +def test_bulk_shared_prefetch_falls_back_above_layered_capacity(): + class Exchange: + layered_max_records = 8191 + + def exchange_layers(self, **kwargs): + raise AssertionError("oversized batches must use the per-layer path") + + used_bulk = try_exchange_sparse_decode_shared_layers( + enabled=True, + transport="ce", + exchange=Exchange(), + records_by_layer=(object(), object(), object()), + local_indices_by_destination=object(), + outputs_by_layer=(object(), object(), object()), + active_records=8192, + pool_records=65536, + ) + + assert not used_bulk + + +def test_bulk_shared_prefetch_rejects_pool_overflow(): + class Exchange: + def exchange_layers(self, **kwargs): + raise AssertionError("overflow must fail before exchange") + + with pytest.raises(ValueError, match="exceed pool"): + try_exchange_sparse_decode_shared_layers( + enabled=True, + transport="ce", + exchange=Exchange(), + records_by_layer=(object(), object(), object()), + local_indices_by_destination=object(), + outputs_by_layer=(object(), object(), object()), + active_records=65537, + pool_records=65536, + ) + + +def test_bulk_shared_prefetch_reuses_existing_sparse_state(monkeypatch): + layout = plan_sparse_ckv_decode( + dcp_world_size=4, + topk=2, + rows_per_request=2, + max_requests=1, + pool_records=4, + record_bytes=8, + prefetch_depth=3, + ) + + class Output: + def __init__(self, slot): + self.slot = slot + self.streams = [] + + def record_stream(self, stream): + self.streams.append(stream) + + outputs = [Output(slot) for slot in range(layout.workspace_slots)] + + class Workspace: + def __getitem__(self, key): + slot, records = key + assert records == slice(None, layout.active_records(1), None) + return outputs[slot] + + class Event: + def __init__(self): + self.streams = [] + + def record(self, stream): + self.streams.append(stream) + + class Stream: + def __init__(self): + self.waited_for = [] + + def wait_stream(self, stream): + self.waited_for.append(stream) + + class Exchange: + layered_max_records = layout.pool_records + + def __init__(self): + self.calls = [] + + def exchange_layers(self, **kwargs): + self.calls.append(kwargs) + + class State: + payload_workspace = Workspace() + transport_local_slots = {1: object()} + complete_events = [Event() for _ in range(layout.workspace_slots)] + + class PrefetchState: + def get_sparse_decode_state(self, actual_layout, actual_exchange): + assert actual_layout is layout + assert actual_exchange is exchange + return state + + class Metadata: + num_reqs = 1 + + exchange = Exchange() + state = State() + stream = Stream() + current_stream = object() + monkeypatch.setitem( + b12x_sparse_ckv_decode._SELECTED_RECORD_STREAMS, + id(exchange), + stream, + ) + monkeypatch.setattr( + b12x_sparse_ckv_decode.envs, + "VLLM_B12X_MLA_SPARSE_DECODE_TRANSPORT", + "ce", + ) + monkeypatch.setattr(torch.cuda, "current_stream", lambda: current_stream) + monkeypatch.setattr(torch.cuda, "stream", lambda _stream: nullcontext()) + + impl = object.__new__(B12xMLASparseImpl) + impl._sparse_decode_bulk_prefetch = True + impl._sparse_decode_layout = layout + impl._sparse_decode_exchange = exchange + impl.block_size = 1 + impl.cp_kv_cache_interleave_size = 1 + targets = [] + target_kv_caches = [] + for target_idx in (1, 2, 3): + target_impl = object.__new__(B12xMLASparseImpl) + target_impl._sparse_decode_layout = layout + target_impl._sparse_decode_exchange = exchange + target_impl.block_size = 1 + target_impl.cp_kv_cache_interleave_size = 1 + target_kv = torch.zeros((1, 1, layout.record_bytes), dtype=torch.uint8) + target_kv_caches.append(target_kv) + targets.append((target_idx, target_impl, target_kv)) + + def fail_allocation(*args, **kwargs): + raise AssertionError("bulk prefetch must reuse allocated sparse state") + + for allocator in ("empty", "empty_like", "zeros", "zeros_like"): + monkeypatch.setattr(torch, allocator, fail_allocation) + result = impl._dcp_gather_sparse_decode_shared_layers( + targets, + Metadata(), + PrefetchState(), + ) + + assert result is not None + layer_slots, complete = result + assert layer_slots == {1: 1, 2: 2, 3: 3} + assert stream.waited_for == [current_stream] + assert all(output.streams == [stream] for output in outputs[1:4]) + assert complete is state.complete_events[1] + assert complete.streams == [stream] + assert len(exchange.calls) == 1 + call = exchange.calls[0] + assert all( + actual.data_ptr() == expected.data_ptr() + and actual.dtype == expected.dtype + and actual.shape[-1] == layout.record_bytes + for actual, expected in zip( + call["records_by_layer"], target_kv_caches, strict=True + ) + ) + assert call["outputs_by_layer"] == tuple(outputs[1:4]) + assert call["local_indices_by_destination"] is state.transport_local_slots[1] + + +def _install_fake_b12x_transport(monkeypatch, *, ce_error=None, include_ce=True): + calls = [] + + class InitializationError(RuntimeError): + pass + + class DirectExchange: + @classmethod + def from_process_group(cls, **kwargs): + calls.append("direct") + return cls() + + def close(self): + pass + + class CopyExchange: + @classmethod + def from_process_group(cls, **kwargs): + calls.append("ce") + if ce_error is not None: + raise InitializationError(ce_error) + return cls() + + def close(self): + pass + + package = ModuleType("sparkinfer") + package.__path__ = [] + comm = ModuleType("sparkinfer.comm") + comm.__path__ = [] + pcie = ModuleType("sparkinfer.comm.pcie") + pcie.SelectedRecordExchange = DirectExchange + pcie.SelectedRecordExchangeInitializationError = InitializationError + if include_ce: + pcie.SelectedRecordCopyExchange = CopyExchange + package.comm = comm + comm.pcie = pcie + monkeypatch.setitem(sys.modules, "sparkinfer", package) + monkeypatch.setitem(sys.modules, "sparkinfer.comm", comm) + monkeypatch.setitem(sys.modules, "sparkinfer.comm.pcie", pcie) + return calls, DirectExchange, CopyExchange + + +def _prepare_transport_factory_test(monkeypatch, transport): + monkeypatch.setattr( + b12x_sparse_ckv_decode.envs, + "VLLM_B12X_MLA_SPARSE_DECODE_TRANSPORT", + transport, + ) + monkeypatch.setattr(b12x_sparse_ckv_decode, "_SELECTED_RECORD_EXCHANGES", {}) + monkeypatch.setattr(b12x_sparse_ckv_decode, "_SELECTED_RECORD_STREAMS", {}) + stream = object() + monkeypatch.setattr(torch.cuda, "Stream", lambda **kwargs: stream) + return stream + + +def test_selected_record_transport_auto_prefers_copy_engine(monkeypatch): + stream = _prepare_transport_factory_test(monkeypatch, "auto") + calls, _, CopyExchange = _install_fake_b12x_transport(monkeypatch) + + exchange = b12x_sparse_ckv_decode.get_selected_record_exchange( + process_group=object(), + device=torch.device("cuda", 0), + layout=_layout(), + lane_key=("target", 1), + ) + + assert isinstance(exchange, CopyExchange) + assert calls == ["ce"] + assert b12x_sparse_ckv_decode.get_selected_record_stream(exchange) is stream + + +def test_selected_record_transport_auto_falls_back_to_direct(monkeypatch): + _prepare_transport_factory_test(monkeypatch, "auto") + calls, DirectExchange, _ = _install_fake_b12x_transport( + monkeypatch, ce_error="CE unavailable" + ) + + exchange = b12x_sparse_ckv_decode.get_selected_record_exchange( + process_group=object(), + device=torch.device("cuda", 0), + layout=_layout(), + lane_key=("target", 1), + ) + + assert isinstance(exchange, DirectExchange) + assert calls == ["ce", "direct"] + + +def test_selected_record_transport_auto_supports_older_b12x(monkeypatch): + _prepare_transport_factory_test(monkeypatch, "auto") + calls, DirectExchange, _ = _install_fake_b12x_transport( + monkeypatch, include_ce=False + ) + + exchange = b12x_sparse_ckv_decode.get_selected_record_exchange( + process_group=object(), + device=torch.device("cuda", 0), + layout=_layout(), + lane_key=("target", 1), + ) + + assert isinstance(exchange, DirectExchange) + assert calls == ["direct"] + + +def test_selected_record_transport_direct_never_constructs_copy_engine(monkeypatch): + _prepare_transport_factory_test(monkeypatch, "direct") + calls, DirectExchange, _ = _install_fake_b12x_transport(monkeypatch) + + exchange = b12x_sparse_ckv_decode.get_selected_record_exchange( + process_group=object(), + device=torch.device("cuda", 0), + layout=_layout(), + lane_key=("target", 1), + ) + + assert isinstance(exchange, DirectExchange) + assert calls == ["direct"] + + +def test_selected_record_transport_isolates_target_and_draft_lanes(monkeypatch): + _prepare_transport_factory_test(monkeypatch, "direct") + calls, _, _ = _install_fake_b12x_transport(monkeypatch) + monkeypatch.setattr(torch.cuda, "Stream", lambda **kwargs: object()) + process_group = object() + layout = _layout() + + target = b12x_sparse_ckv_decode.get_selected_record_exchange( + process_group=process_group, + device=torch.device("cuda", 0), + layout=layout, + lane_key=("target", 1), + ) + draft = b12x_sparse_ckv_decode.get_selected_record_exchange( + process_group=process_group, + device=torch.device("cuda", 0), + layout=layout, + lane_key=("draft", 1), + ) + + assert target is not draft + assert b12x_sparse_ckv_decode.get_selected_record_stream(target) is not ( + b12x_sparse_ckv_decode.get_selected_record_stream(draft) + ) + assert calls == ["direct", "direct"] + + +def test_selected_record_transport_strict_ce_does_not_fallback(monkeypatch): + _prepare_transport_factory_test(monkeypatch, "ce") + calls, _, _ = _install_fake_b12x_transport(monkeypatch, ce_error="CE unavailable") + + with pytest.raises(RuntimeError, match="strict copy-engine"): + b12x_sparse_ckv_decode.get_selected_record_exchange( + process_group=object(), + device=torch.device("cuda", 0), + layout=_layout(), + lane_key=("target", 1), + ) + + assert calls == ["ce"] + + +def test_selected_record_transport_rejects_unknown_mode(monkeypatch): + _prepare_transport_factory_test(monkeypatch, "mystery") + + with pytest.raises(ValueError, match="must be auto, ce, or direct"): + b12x_sparse_ckv_decode.get_selected_record_exchange( + process_group=object(), + device=torch.device("cuda", 0), + layout=_layout(), + lane_key=("target", 1), + ) + + @pytest.mark.parametrize("dcp", [2, 3, 4, 5, 6, 7, 8]) def test_sparse_decode_layout_supports_dcp_two_through_eight(dcp): layout = _layout(dcp=dcp) diff --git a/vllm/envs.py b/vllm/envs.py index 1522ff3bb169..0c802dc02858 100755 --- a/vllm/envs.py +++ b/vllm/envs.py @@ -78,6 +78,8 @@ VLLM_DCP_QUERY_SPLIT: bool = False VLLM_B12X_MLA_CKV_GATHER: bool = False VLLM_B12X_MLA_SPARSE_DECODE_CKV_GATHER: bool = False + VLLM_B12X_MLA_SPARSE_DECODE_TRANSPORT: str = "direct" + VLLM_B12X_MLA_SPARSE_DECODE_BULK_PREFETCH: bool = False VLLM_B12X_MLA_SPARSE_DECODE_MAX_SEQS: int = 8 VLLM_B12X_MLA_SPARSE_DECODE_POOL_RECORDS: int = 0 VLLM_B12X_MLA_CKV_GATHER_MIN_TOKENS: int = 16 @@ -443,6 +445,26 @@ def _get_validated_env() -> str | None: return _get_validated_env +def env_bool_with_choices( + env_name: str, + default: bool = False, +) -> Callable[[], bool]: + """Create a case-insensitive, strictly validated boolean env getter.""" + getter = env_with_choices( + env_name, + "1" if default else "0", + ["0", "1", "false", "true", "no", "yes", "off", "on"], + case_sensitive=False, + ) + + def _get_validated_bool() -> bool: + value = getter() + assert value is not None + return value.lower() in ("1", "true", "yes", "on") + + return _get_validated_bool + + def env_list_with_choices( env_name: str, default: list[str], @@ -1183,6 +1205,21 @@ def _resolve_rust_frontend_path() -> str | None: os.getenv("VLLM_B12X_MLA_SPARSE_DECODE_CKV_GATHER", "0").lower() in ("1", "true", "yes", "on") ), + # Transport for sparse selected-record decode. Keep the existing direct + # path as the default because copy-engine staging has additional VRAM cost. + "VLLM_B12X_MLA_SPARSE_DECODE_TRANSPORT": env_with_choices( + "VLLM_B12X_MLA_SPARSE_DECODE_TRANSPORT", + "direct", + ["auto", "ce", "direct"], + ), + # Combine exactly three eligible Shared-layer selected-record prefetches + # into one CE exchange. Sparse decode must also be enabled. This is + # effective only with transport=ce|auto and sufficient layered capacity; + # direct transport, partial groups, and oversized groups retain the + # existing per-layer path. + "VLLM_B12X_MLA_SPARSE_DECODE_BULK_PREFETCH": env_bool_with_choices( + "VLLM_B12X_MLA_SPARSE_DECODE_BULK_PREFETCH" + ), "VLLM_B12X_MLA_SPARSE_DECODE_MAX_SEQS": lambda: int( os.getenv("VLLM_B12X_MLA_SPARSE_DECODE_MAX_SEQS", "8") ), diff --git a/vllm/v1/attention/backends/mla/b12x_mla_sparse.py b/vllm/v1/attention/backends/mla/b12x_mla_sparse.py index 0af5756b914e..21c564579a06 100644 --- a/vllm/v1/attention/backends/mla/b12x_mla_sparse.py +++ b/vllm/v1/attention/backends/mla/b12x_mla_sparse.py @@ -158,7 +158,7 @@ def _ckv_prefetch_target_indices( layer_idx: int, depth: int, layer_caches: list[torch.Tensor | None], - pending_layers: dict[int, tuple[Any, int]], + pending_layers: dict[int, "_CKVPrefetchPendingLayer"], ) -> list[int]: targets: list[int] = [] for distance in range(1, max(0, int(depth)) + 1): @@ -182,6 +182,28 @@ class _CKVWorkspaceIdentity: dtype: torch.dtype +@dataclass +class _CKVPrefetchTicket: + event: Any + wait_scheduled: bool = False + + def wait_once(self) -> None: + if not self.wait_scheduled: + self.event.wait() + self.wait_scheduled = True + + def wait_on_stream_once(self, stream: Any) -> None: + if not self.wait_scheduled: + stream.wait_event(self.event) + self.wait_scheduled = True + + +@dataclass(frozen=True) +class _CKVPrefetchPendingLayer: + ticket: _CKVPrefetchTicket + buf_idx: int + + def _ckv_workspace_identity(workspace: torch.Tensor) -> _CKVWorkspaceIdentity: storage = workspace.untyped_storage() return _CKVWorkspaceIdentity( @@ -209,7 +231,7 @@ def __init__( self.workspace_storage_ref = weakref.ref(workspace.untyped_storage()) self.layer_caches: list[torch.Tensor | None] = [] self.layer_impls: list[B12xMLASparseImpl | None] = [] - self.pending_layers: dict[int, tuple[Any, int]] = {} + self.pending_layers: dict[int, _CKVPrefetchPendingLayer] = {} self.sparse_decode_state: Any | None = None self.gather_stream: torch.cuda.Stream | None = None self.ckv_workspace: torch.Tensor | None = None @@ -217,13 +239,46 @@ def __init__( self.last_layer_idx: int | None = None def begin_step(self) -> None: - for event, _ in self.pending_layers.values(): + tickets = { + id(pending.ticket): pending.ticket + for pending in self.pending_layers.values() + } + for ticket in tickets.values(): # Preserve ring ordering without blocking the host indefinitely. # The next main-stream gather is enqueued after these dependencies. - event.wait() + ticket.wait_once() self.pending_layers.clear() self.last_layer_idx = None + def register_pending_group( + self, + layer_slots: dict[int, int], + event: Any, + ) -> _CKVPrefetchTicket: + if not layer_slots: + raise ValueError("prefetch group must contain at least one layer") + duplicate_layers = self.pending_layers.keys() & layer_slots.keys() + if duplicate_layers: + raise RuntimeError( + f"prefetch group overlaps pending layers {sorted(duplicate_layers)}" + ) + ticket = _CKVPrefetchTicket(event) + self.pending_layers.update( + { + layer_idx: _CKVPrefetchPendingLayer(ticket, int(buf_idx)) + for layer_idx, buf_idx in layer_slots.items() + } + ) + return ticket + + def pop_pending_layer( + self, + layer_idx: int | None, + ) -> _CKVPrefetchPendingLayer | None: + if layer_idx is None: + return None + return self.pending_layers.pop(layer_idx, None) + def enter_layer(self, layer_idx: int) -> None: if self.last_layer_idx is not None and layer_idx <= self.last_layer_idx: self.begin_step() @@ -1445,6 +1500,9 @@ def _make_plan( ) self._sparse_decode_layout = None self._sparse_decode_exchange = None + self._sparse_decode_bulk_prefetch = bool( + envs_mod.VLLM_B12X_MLA_SPARSE_DECODE_BULK_PREFETCH + ) if self._sparse_decode_enabled: self._sparse_decode_layout = plan_sparse_ckv_decode( dcp_world_size=self.dcp_world_size, @@ -2203,6 +2261,88 @@ def _dcp_gather_sparse_decode_ckv( gathered_cache = output.view(-1, self.block_size, layout.record_bytes) return gathered_cache, selected_indices, nsa_cache_seqlens, complete + def _dcp_gather_sparse_decode_shared_layers( + self, + targets: list[tuple[int, "B12xMLASparseImpl", torch.Tensor]], + attn_metadata: B12xMLASparseMetadata, + prefetch_state: _CKVPrefetchState, + ) -> tuple[dict[int, int], torch.cuda.Event] | None: + """Exchange one Full layer's S1/S2/S3 CKV payloads together.""" + layout = self._sparse_decode_layout + exchange = self._sparse_decode_exchange + if ( + not self._sparse_decode_bulk_prefetch + or layout is None + or exchange is None + or len(targets) != 3 + ): + return None + + num_requests = int(attn_metadata.num_reqs) + active_records = layout.active_records(num_requests) + state = prefetch_state.get_sparse_decode_state(layout, exchange) + layer_slots: dict[int, int] = {} + records_by_layer: list[torch.Tensor] = [] + outputs_by_layer: list[torch.Tensor] = [] + for target_idx, target_impl, target_kv in targets: + if ( + target_impl._sparse_decode_layout != layout + or target_impl._sparse_decode_exchange is not exchange + or target_impl.block_size != self.block_size + or target_impl.cp_kv_cache_interleave_size + != self.cp_kv_cache_interleave_size + or target_kv.dtype != torch.uint8 + or target_kv.ndim != 3 + or tuple(target_kv.shape[1:]) != (self.block_size, layout.record_bytes) + or not target_kv.is_contiguous() + ): + return None + buf_idx = target_idx % layout.workspace_slots + if buf_idx in layer_slots.values(): + return None + layer_slots[target_idx] = buf_idx + records_by_layer.append(target_kv.reshape(-1, layout.record_bytes)) + outputs_by_layer.append(state.payload_workspace[buf_idx, :active_records]) + + from vllm import envs as envs_mod + from vllm.v1.attention.backends.mla.b12x_sparse_ckv_decode import ( + get_selected_record_stream, + try_exchange_sparse_decode_shared_layers, + ) + + transport = envs_mod.VLLM_B12X_MLA_SPARSE_DECODE_TRANSPORT + if transport not in {"auto", "ce"}: + return None + if not callable(getattr(exchange, "exchange_layers", None)): + return None + layered_capacity = getattr(exchange, "layered_max_records", None) + if layered_capacity is not None and active_records > int(layered_capacity): + return None + + gather_stream = get_selected_record_stream(exchange) + gather_stream.wait_stream(torch.cuda.current_stream()) + for output in outputs_by_layer: + output.record_stream(gather_stream) + + with torch.cuda.stream(gather_stream): + used_bulk = try_exchange_sparse_decode_shared_layers( + enabled=True, + transport=transport, + exchange=exchange, + records_by_layer=tuple(records_by_layer), + local_indices_by_destination=( + state.transport_local_slots[num_requests] + ), + outputs_by_layer=tuple(outputs_by_layer), + active_records=active_records, + pool_records=layout.pool_records, + ) + if not used_bulk: + return None + complete = state.complete_events[next(iter(layer_slots.values()))] + complete.record(gather_stream) + return layer_slots, complete + def _append_current_token_to_sparse_decode_gathered( self, gathered_cache: torch.Tensor, @@ -2758,14 +2898,10 @@ def forward_mqa( prefetch_state.register_cache(layer_idx, kv_cache) prefetch_state.register_impl(layer_idx, self) sparse_state = prefetch_state.get_sparse_decode_state(layout, exchange) - pending = ( - prefetch_state.pending_layers.pop(layer_idx, None) - if layer_idx is not None - else None - ) + pending = prefetch_state.pop_pending_layer(layer_idx) if pending is not None: - gather_event, current_buf_idx = pending - torch.cuda.current_stream().wait_event(gather_event) + pending.ticket.wait_on_stream_once(torch.cuda.current_stream()) + current_buf_idx = pending.buf_idx active_records = layout.active_records(int(attn_metadata.num_reqs)) kv_cache = sparse_state.payload_workspace[ current_buf_idx, :active_records @@ -2837,6 +2973,7 @@ def forward_mqa( layout.prefetch_depth, emits_topk_by_layer, ) + target_entries: list[tuple[int, B12xMLASparseImpl, torch.Tensor]] = [] for target_idx in targets: if target_idx in prefetch_state.pending_layers: continue @@ -2844,20 +2981,37 @@ def forward_mqa( target_kv = prefetch_state.layer_caches[target_idx] if target_impl is None or target_kv is None: break - target_buf_idx = target_idx % layout.workspace_slots - _, _, _, target_event = target_impl._dcp_gather_sparse_decode_ckv( - target_kv, - attn_metadata, - topk_indices, - prefetch_state, - buf_idx=target_buf_idx, - build_union=False, - wait_for_completion=False, - ) - prefetch_state.pending_layers[target_idx] = ( + target_entries.append((target_idx, target_impl, target_kv)) + + bulk_pending = self._dcp_gather_sparse_decode_shared_layers( + target_entries, + attn_metadata, + prefetch_state, + ) + if bulk_pending is not None: + layer_slots, target_event = bulk_pending + prefetch_state.register_pending_group( + layer_slots, target_event, - target_buf_idx, ) + else: + for target_idx, target_impl, target_kv in target_entries: + target_buf_idx = target_idx % layout.workspace_slots + _, _, _, target_event = ( + target_impl._dcp_gather_sparse_decode_ckv( + target_kv, + attn_metadata, + topk_indices, + prefetch_state, + buf_idx=target_buf_idx, + build_union=False, + wait_for_completion=False, + ) + ) + prefetch_state.register_pending_group( + {target_idx: target_buf_idx}, + target_event, + ) if use_ckv_gather: layer_idx = self._resolve_layer_index(layer) prefetch_registry = attn_metadata.ckv_prefetch_registry @@ -2871,13 +3025,13 @@ def forward_mqa( prefetch_state.enter_layer(layer_idx) prefetch_state.register_cache(layer_idx, kv_cache) pending = ( - prefetch_state.pending_layers.pop(layer_idx, None) - if layer_idx is not None and self._ckv_prefetch_depth > 0 + prefetch_state.pop_pending_layer(layer_idx) + if self._ckv_prefetch_depth > 0 else None ) if pending is not None: - gather_event, current_buf_idx = pending - gather_event.wait() + pending.ticket.wait_once() + current_buf_idx = pending.buf_idx _, gathered_buffer = self._ckv_workspace_views( ckv_workspace, current_buf_idx ) @@ -2935,9 +3089,9 @@ def forward_mqa( ) target_event = torch.cuda.Event(blocking=False) target_event.record(prefetch_stream) - prefetch_state.pending_layers[target_idx] = ( + prefetch_state.register_pending_group( + {target_idx: target_buf_idx}, target_event, - target_buf_idx, ) use_decode_kernel = ( diff --git a/vllm/v1/attention/backends/mla/b12x_sparse_ckv_decode.py b/vllm/v1/attention/backends/mla/b12x_sparse_ckv_decode.py index 4234b0a05f73..f144df925242 100644 --- a/vllm/v1/attention/backends/mla/b12x_sparse_ckv_decode.py +++ b/vllm/v1/attention/backends/mla/b12x_sparse_ckv_decode.py @@ -18,7 +18,11 @@ from torch.distributed import ProcessGroup from torch.utils.cpp_extension import load +from vllm import envs from vllm.distributed.parallel_state import register_model_parallel_cleanup_hook +from vllm.logger import init_logger + +logger = init_logger(__name__) @dataclass(frozen=True) @@ -135,6 +139,58 @@ def sparse_decode_prefetch_targets( return targets +def try_exchange_sparse_decode_shared_layers( + *, + enabled: bool, + transport: str, + exchange: object, + records_by_layer: tuple[torch.Tensor, ...], + local_indices_by_destination: torch.Tensor, + outputs_by_layer: tuple[torch.Tensor, ...], + active_records: int, + pool_records: int, +) -> bool: + """Use one copy-engine exchange for a Full layer's S1/S2/S3 payloads.""" + if not enabled: + return False + if transport not in {"auto", "ce"}: + logger.info_once( + "Bulk S1/S2/S3 CKV prefetch requires copy-engine transport; " + "using the per-layer exchange path" + ) + return False + if len(records_by_layer) != 3 or len(outputs_by_layer) != 3: + return False + if not 0 < active_records <= pool_records: + raise ValueError( + f"active sparse records {active_records} exceed pool {pool_records}" + ) + exchange_layers = getattr(exchange, "exchange_layers", None) + if exchange_layers is None: + logger.warning_once( + "Bulk S1/S2/S3 CKV prefetch requires B12X " + "PCIeSelectedRecordCopyExchange.exchange_layers; using the " + "per-layer exchange path" + ) + return False + layered_capacity = getattr(exchange, "layered_max_records", None) + if layered_capacity is not None and active_records > int(layered_capacity): + logger.info_once( + "Bulk S1/S2/S3 CKV prefetch capacity %d is below this batch's " + "%d records; using the per-layer exchange path", + int(layered_capacity), + active_records, + ) + return False + exchange_layers( + records_by_layer=records_by_layer, + local_indices_by_destination=local_indices_by_destination, + outputs_by_layer=outputs_by_layer, + ) + logger.info_once("Using one bulk CKV exchange for S1/S2/S3 prefetch") + return True + + def owner_and_local_ordinal( logical_token: int, *, @@ -310,6 +366,12 @@ def get_selected_record_exchange( lane_key: tuple[str, int], ): """Return one generic B12X exchange per target/draft execution lane.""" + transport = envs.VLLM_B12X_MLA_SPARSE_DECODE_TRANSPORT + if transport not in {"auto", "ce", "direct"}: + raise ValueError( + "VLLM_B12X_MLA_SPARSE_DECODE_TRANSPORT must be auto, ce, or " + f"direct; got {transport!r}" + ) key = ( id(process_group), device.type, @@ -317,17 +379,53 @@ def get_selected_record_exchange( layout.dcp_world_size, layout.pool_records, layout.record_bytes, + transport, lane_key, ) exchange = _SELECTED_RECORD_EXCHANGES.get(key) if exchange is None: - from sparkinfer.comm.pcie import SelectedRecordExchange + from sparkinfer.comm.pcie import ( + SelectedRecordExchange, + SelectedRecordExchangeInitializationError, + ) - exchange = SelectedRecordExchange.from_process_group( - process_group=process_group, - device=device, - max_records=layout.pool_records, - record_bytes=layout.record_bytes, + ce_error: Exception | None = None + if transport in {"auto", "ce"}: + try: + from sparkinfer.comm.pcie import SelectedRecordCopyExchange + + exchange = SelectedRecordCopyExchange.from_process_group( + process_group=process_group, + device=device, + max_records=layout.pool_records, + record_bytes=layout.record_bytes, + ) + except ( + ImportError, + SelectedRecordExchangeInitializationError, + ) as exc: + ce_error = exc + if transport == "ce": + raise RuntimeError( + "strict copy-engine selected-record transport failed " + "to initialize" + ) from exc + if exchange is None: + if ce_error is not None: + logger.warning_once( + "Copy-engine selected-record transport is unavailable; " + "falling back to direct peer writes: %s", + ce_error, + ) + exchange = SelectedRecordExchange.from_process_group( + process_group=process_group, + device=device, + max_records=layout.pool_records, + record_bytes=layout.record_bytes, + ) + logger.info_once( + "Using %s for sparse selected-record CKV exchange", + type(exchange).__name__, ) _SELECTED_RECORD_EXCHANGES[key] = exchange if id(exchange) not in _SELECTED_RECORD_STREAMS: