diff --git a/afd_plugin/compat/npu/feature_validation.py b/afd_plugin/compat/npu/feature_validation.py index ffc54f53d..653f3c301 100644 --- a/afd_plugin/compat/npu/feature_validation.py +++ b/afd_plugin/compat/npu/feature_validation.py @@ -184,9 +184,21 @@ def _fail_if_unsupported_npu_afd_async_features( "with async=true and connector='CAMAsyncAFDConnector'", ) if not bool(vllm_config.model_config.enforce_eager): - raise RuntimeError( - "CAMAsyncAFDConnector supports only eager Attention/FFN execution", - ) + graph_mode = vllm_config.compilation_config.cudagraph_mode.name + if not ( + afd_config.role == "ffn" + and async_cam_layered_gmm_enabled() + and _is_dsv4_target(vllm_config) + and afd_config.compute_gate_on_attention + and extra_info.dynamic_quant == 1 + and not vllm_config.use_v2_model_runner + and graph_mode == "FULL" + ): + raise RuntimeError( + "CAMAsyncAFDConnector graph execution requires FFN FULL, " + "DeepSeek V4, layered W4A8 GMM, Attention-side gate, " + "dynamicQuant=1, and ModelRunnerV1; Attention remains eager" + ) if bool(parallel_config.enable_dbo) or bool(parallel_config.use_ubatching): raise RuntimeError( "CAMAsyncAFDConnector does not support vLLM native ubatching/DBO", diff --git a/afd_plugin/connectors/factory.py b/afd_plugin/connectors/factory.py index 6e94e1ac1..9a47da9ec 100644 --- a/afd_plugin/connectors/factory.py +++ b/afd_plugin/connectors/factory.py @@ -19,6 +19,10 @@ if TYPE_CHECKING: from vllm.config import VllmConfig + from afd_plugin.distributed.afd_process_group import ( + ProcessGroupRendezvousContext, + ) + class AFDConnectorFactory: _registry: dict[str, Callable[[], type[AFDConnectorBase]]] = {} @@ -53,12 +57,29 @@ def create_connector( local_rank: int, vllm_config: VllmConfig, afd_config: AFDConfig | None = None, + *, + rendezvous_context: ProcessGroupRendezvousContext | None = None, ) -> AFDConnectorBase: config = afd_config or parse_afd_config(vllm_config) if config.connector not in cls._registry: raise ValueError(f"unsupported AFD connector type: {config.connector}") connector_cls = cls._registry[config.connector]() role_rank = resolve_role_rank(vllm_config, config) + if rendezvous_context is not None: + from afd_plugin.connectors.npu.async_cam import CAMAsyncAFDConnector + + if not issubclass(connector_cls, CAMAsyncAFDConnector): + raise TypeError( + "rendezvous_context requires a CAMAsyncAFDConnector class" + ) + return connector_cls( + rank, + local_rank, + vllm_config, + config, + role_rank, + rendezvous_context=rendezvous_context, + ) return connector_cls( rank, local_rank, diff --git a/afd_plugin/connectors/npu/async_cam.py b/afd_plugin/connectors/npu/async_cam.py index f79c33070..fb6834b22 100644 --- a/afd_plugin/connectors/npu/async_cam.py +++ b/afd_plugin/connectors/npu/async_cam.py @@ -2,7 +2,7 @@ # SPDX-FileCopyrightText: Copyright contributors to the AFD plugin project """Ascend CAM asynchronous connector for Attention/FFN disaggregation. -``CAMAsyncAFDConnector`` is the eager-only Ascend inference data path. +``CAMAsyncAFDConnector`` is the Ascend inference data path. Attention ranks run MoE routing, submit activations with CAM async dispatch-send, and receive combined expert output with combine-recv. FFN ranks receive routed-expert activations with dispatch-recv, execute their @@ -15,10 +15,10 @@ the CAM operator payload and does not create a separate Gloo DP-metadata control plane. -The supported deployment requires ``async=true``, eager execution, Ascend CAM -operator packages, and matching topology/configuration on every rank. Regular -prefill and autoregressive decode steps use the same connector; vLLM native DBO -and ACL graph execution are not supported. +The supported deployment requires ``async=true``, Ascend CAM operator packages, +and matching topology/configuration on every rank. Attention runs eager; the +layered W4A8 FFN path can use one FULL graph. Regular prefill and autoregressive +decode steps use the same connector; vLLM native DBO is not supported. Optional AFD-managed MoE ubatching is a separate two-stage pipeline using request boundaries or token-balanced stages for DP+TP/SP. See ``docs/npu/CAM_ASYNC_CONNECTOR_USER_GUIDE.md`` for configuration, rank @@ -34,6 +34,7 @@ from typing import TYPE_CHECKING, Any, Final, cast import torch +import torch.distributed as dist from torch import Tensor from vllm.logger import init_logger @@ -58,6 +59,7 @@ AFDTransferState, ) from afd_plugin.distributed import ( + ProcessGroupRendezvousContext, create_hccl_process_group_options, init_afd_process_group, ) @@ -237,6 +239,8 @@ def __init__( vllm_config: VllmConfig, afd_config: AFDConfig, role_rank: int, + *, + rendezvous_context: ProcessGroupRendezvousContext | None = None, ) -> None: """Derive CAM topology, tensor dimensions, and connector state. @@ -262,6 +266,7 @@ def __init__( self.comm_id = CAM_COMM_ID self.tp_size = extra_info.attn_ranks_per_dp self.cam_pg: ProcessGroup | None = None + self._rendezvous_context = rendezvous_context self.topology = build_async_topology( afd_config, role_rank, @@ -304,6 +309,11 @@ def init_afd_connector(self) -> None: pg_options=create_hccl_process_group_options( self.hccl_buffer_size_mb, ), + on_rendezvous=( + self._rendezvous_context.retain_store + if self._rendezvous_context is not None + else None + ), ) backend = self.cam_pg._get_backend(torch.device("npu")) self.group_name = str(backend.get_hccl_comm_name(self.world_rank)) @@ -314,19 +324,21 @@ def init_afd_connector(self) -> None: dtype=self.activation_dtype, device=device, ) + if self._rendezvous_context is not None: + self._rendezvous_context.bind(self.cam_pg) self._initialized = True def close(self) -> None: - """Destroy the HCCL process group and clear pending transfer states.""" + """Destroy the communicator and release its operator buffers.""" if self.cam_pg is not None: - import torch.distributed as dist - dist.destroy_process_group(self.cam_pg) self.cam_pg = None + self._initialized = False self.comm_args = None self._placeholder = None self._pending_attention_payloads.clear() - self._initialized = False + if self._rendezvous_context is not None: + self._rendezvous_context.invalidate() def select_experts(self, **kwargs: Any) -> tuple[Tensor, Tensor]: """Run the pinned vLLM-Ascend expert selector on Attention.""" diff --git a/afd_plugin/distributed/__init__.py b/afd_plugin/distributed/__init__.py index b7b28853b..995ae585c 100644 --- a/afd_plugin/distributed/__init__.py +++ b/afd_plugin/distributed/__init__.py @@ -15,6 +15,7 @@ def __getattr__(name: str): if name in { "DefaultProcessGroupSwitcher", + "ProcessGroupRendezvousContext", "create_hccl_process_group_options", "init_afd_process_group", }: @@ -29,6 +30,7 @@ def __getattr__(name: str): __all__ = [ "AFDRankMapping", "DefaultProcessGroupSwitcher", + "ProcessGroupRendezvousContext", "build_rank_mapping", "create_hccl_process_group_options", "init_afd_process_group", diff --git a/afd_plugin/distributed/afd_process_group.py b/afd_plugin/distributed/afd_process_group.py index 5dbfab507..2ea33db20 100644 --- a/afd_plugin/distributed/afd_process_group.py +++ b/afd_plugin/distributed/afd_process_group.py @@ -4,6 +4,7 @@ from __future__ import annotations +from collections.abc import Callable from datetime import timedelta from typing import Any @@ -12,6 +13,7 @@ from torch.distributed.distributed_c10d import ( PrefixStore, ProcessGroup, + Store, _new_process_group_helper, _update_default_pg, _world, @@ -21,6 +23,35 @@ from vllm.utils.torch_utils import is_torch_equal_or_newer +class ProcessGroupRendezvousContext: + """Borrow the Store and process group created by one CAM rendezvous.""" + + def __init__(self) -> None: + self._store: Store | None = None + self._process_group: ProcessGroup | None = None + self._closed = False + + def retain_store(self, store: Store) -> None: + if self._closed or self._store is not None: + raise RuntimeError("CAM rendezvous context cannot accept another Store") + self._store = store + + def bind(self, process_group: ProcessGroup) -> None: + if self._closed or self._store is None or self._process_group is not None: + raise RuntimeError("CAM rendezvous context is not ready to bind") + self._process_group = process_group + + def borrow(self) -> tuple[Store, ProcessGroup]: + if self._closed or self._store is None or self._process_group is None: + raise RuntimeError("CAM rendezvous context is not bound") + return self._store, self._process_group + + def invalidate(self) -> None: + self._closed = True + self._store = None + self._process_group = None + + class DefaultProcessGroupSwitcher: """Temporarily switch PyTorch's default process group.""" @@ -67,6 +98,7 @@ def init_afd_process_group( group_name: str, timeout: timedelta, pg_options: Any | None = None, + on_rendezvous: Callable[[Store], None] | None = None, ) -> ProcessGroup: """Create a plugin-owned process group without patching vLLM source. @@ -83,6 +115,8 @@ def init_afd_process_group( ) store, rank, world_size = next(rendezvous_iterator) store.set_timeout(timeout) + if on_rendezvous is not None: + on_rendezvous(store) prefixed_store = PrefixStore(group_name, store) backend_value = Backend(backend) if backend else Backend("undefined") pg_options_param_name = ( @@ -118,6 +152,7 @@ def init_afd_process_group( __all__ = [ "DefaultProcessGroupSwitcher", + "ProcessGroupRendezvousContext", "create_hccl_process_group_options", "init_afd_process_group", ] diff --git a/afd_plugin/v1/worker/npu/async_cam_startup.py b/afd_plugin/v1/worker/npu/async_cam_startup.py new file mode 100644 index 000000000..42631cd16 --- /dev/null +++ b/afd_plugin/v1/worker/npu/async_cam_startup.py @@ -0,0 +1,385 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the AFD plugin project +"""Host-side startup coordination for the Async CAM Attention/FFN world.""" + +from __future__ import annotations + +import os +import time +from collections.abc import Callable +from dataclasses import dataclass + +import torch +import torch.distributed as dist +from torch.distributed.distributed_c10d import PrefixStore, Store +from vllm.logger import init_logger + +from afd_plugin.connectors.base import AFDConnectorBase +from afd_plugin.connectors.metadata import AFDTransferContext, AFDTransferMetadata +from afd_plugin.connectors.npu.async_cam import ( + AFD_ASYNC_CAM_GROUP_NAME, + AFDAsyncTopology, +) +from afd_plugin.distributed.afd_process_group import ProcessGroupRendezvousContext + +STARTUP_TIMEOUT_SECONDS = 300 +STARTUP_POLL_SECONDS = 0.25 +STARTUP_NONCE_BYTES = 16 +STARTUP_WARMUP_STAGE = 0 +STARTUP_WARMUP_TOKEN_COUNT = 1 +STARTUP_WARMUP_HIDDEN_VALUE = 0.25 + +logger = init_logger(__name__) + + +@dataclass(frozen=True, slots=True) +class AsyncCamStartupSpec: + """Config-derived, immutable shape and topology of one CAM participant.""" + + topology: AFDAsyncTopology + local_rank: int + tp_size: int + hidden_size: int + topk: int + activation_dtype: torch.dtype + + @property + def attention_dp_size(self) -> int: + if self.topology.attn_size % self.tp_size: + raise ValueError("CAM Attention size must be divisible by TP size") + return self.topology.attn_size // self.tp_size + + +@dataclass(frozen=True, slots=True) +class FFNStartupPlan: + """Selected FFN startup mode and the layer used by graph warmup.""" + + use_graph: bool + first_layer_idx: int | None = None + + def __post_init__(self) -> None: + if self.use_graph and ( + self.first_layer_idx is None or self.first_layer_idx < 0 + ): + raise ValueError("CAM FFN graph startup requires a valid first layer") + if not self.use_graph and self.first_layer_idx is not None: + raise ValueError("CAM FFN eager startup cannot select a graph layer") + + +class AsyncCamStartupCoordinator: + """Run one rank's complete CAM startup without owning its communicator.""" + + def __init__( + self, + connector: AFDConnectorBase, + rendezvous_context: ProcessGroupRendezvousContext, + spec: AsyncCamStartupSpec, + ) -> None: + self._connector = connector + self._rendezvous_context = rendezvous_context + self._spec = spec + self._store: Store | None = None + self._started = False + self._starting = False + self._failure: Exception | None = None + + @property + def failed(self) -> bool: + return self._failure is not None + + @property + def started(self) -> bool: + return self._started + + def _begin(self, role: str) -> bool: + if self._spec.topology.role != role: + raise RuntimeError(f"CAM {role} startup requires {role} rank") + if self._failure is not None: + raise RuntimeError("CAM startup previously failed") from self._failure + if self._started: + self._rendezvous_context.borrow() + return False + if self._starting: + raise RuntimeError("CAM startup is already in progress") + self._starting = True + return True + + def _get_store(self) -> Store: + if self._store is not None: + self._rendezvous_context.borrow() + return self._store + store, process_group = self._rendezvous_context.borrow() + topology = self._spec.topology + nonce = torch.zeros( + STARTUP_NONCE_BYTES, + dtype=torch.uint8, + device=f"npu:{self._spec.local_rank}", + ) + if topology.world_rank == 0: + nonce.copy_( + torch.tensor( + list(os.urandom(STARTUP_NONCE_BYTES)), + dtype=torch.uint8, + device=nonce.device, + ) + ) + logger.info("CAM startup nonce broadcast start rank=%d", topology.world_rank) + dist.broadcast(nonce, src=0, group=process_group) + logger.info("CAM startup nonce broadcast done rank=%d", topology.world_rank) + torch.npu.synchronize() + logger.info("CAM startup nonce synchronize done rank=%d", topology.world_rank) + namespace = bytes(nonce.cpu().tolist()).hex() + self._store = PrefixStore( + f"{AFD_ASYNC_CAM_GROUP_NAME}/startup/{namespace}", store + ) + return self._store + + def _failure_keys(self) -> list[str]: + topology = self._spec.topology + return [f"attn/failure/{rank}" for rank in range(topology.attn_size)] + [ + f"ffn/{rank}" for rank in range(topology.attn_size, topology.world_size) + ] + + def _check_failures(self, store: Store, stage: str) -> None: + for key in self._failure_keys(): + if not store.check([key]): + continue + status = store.get(key).decode() + if key.startswith("attn/failure/") or status.startswith("failed:"): + raise RuntimeError( + f"CAM startup rank={self._spec.topology.world_rank} " + f"stage={stage} {key} {status}" + ) + + def _wait_for_entries(self, keys: list[str], *, stage: str) -> list[str]: + store = self._get_store() + deadline = time.monotonic() + STARTUP_TIMEOUT_SECONDS + while True: + self._check_failures(store, stage) + values = [] + pending = [] + for key in keys: + if store.check([key]): + values.append(store.get(key).decode()) + else: + pending.append(key) + if not pending: + self._check_failures(store, stage) + return values + if time.monotonic() >= deadline: + raise TimeoutError( + f"Timed out waiting for CAM startup rank=" + f"{self._spec.topology.world_rank} stage={stage} " + f"keys={sorted(pending)}" + ) + time.sleep(STARTUP_POLL_SECONDS) + + def _wait_for_ffn_modes(self) -> int | None: + topology = self._spec.topology + modes = self._wait_for_entries( + [ + f"ffn/mode/{rank}" + for rank in range(topology.attn_size, topology.world_size) + ], + stage="mode", + ) + if len(set(modes)) != 1: + raise RuntimeError(f"CAM FFN startup modes disagree: {modes}") + mode = modes[0] + if mode == "eager": + return None + if not mode.startswith("graph:"): + raise RuntimeError(f"Invalid CAM FFN startup mode: {mode}") + layer_idx = int(mode.removeprefix("graph:")) + if layer_idx < 0: + raise RuntimeError(f"Invalid CAM FFN startup layer: {layer_idx}") + return layer_idx + + def _run_attention_warmup(self, layer_idx: int) -> None: + spec = self._spec + topology = spec.topology + device = f"npu:{spec.local_rank}" + hidden = torch.full( + (STARTUP_WARMUP_TOKEN_COUNT, spec.hidden_size), + STARTUP_WARMUP_HIDDEN_VALUE, + dtype=spec.activation_dtype, + device=device, + ) + tp_rank = topology.world_rank % spec.tp_size + expert_ids = [ + ((tp_rank * spec.topk + index) % topology.ffn_size) + * topology.expert_per_rank + + ((tp_rank * spec.topk + index) // topology.ffn_size) + % topology.expert_per_rank + for index in range(spec.topk) + ] + ids = torch.tensor([expert_ids], dtype=torch.int32, device=device) + weights = torch.full( + (STARTUP_WARMUP_TOKEN_COUNT, spec.topk), + 1.0 / spec.topk, + dtype=torch.float32, + device=device, + ) + store = self._get_store() + store.set(f"attn/prepared/{topology.world_rank}", "ready") + self._wait_for_entries( + [ + f"ffn/warmup/start/{rank}" + for rank in range(topology.attn_size, topology.world_size) + ], + stage="warmup-start", + ) + metadata = AFDTransferMetadata.create_attention_metadata( + layer_idx=layer_idx, + stage_idx=STARTUP_WARMUP_STAGE, + seq_len=STARTUP_WARMUP_TOKEN_COUNT, + ) + context = AFDTransferContext(metadata=metadata) + logger.info("CAM Attention startup warmup DS rank=%d", topology.world_rank) + self._connector.send_attn_output( + hidden, context, topk_ids=ids, topk_weights=weights + ) + # Omit context and routing kwargs so the connector consumes its FIFO. + warmup_output = self._connector.recv_ffn_output( + hidden, ubatch_idx=STARTUP_WARMUP_STAGE + ) + torch.npu.synchronize() + del warmup_output + logger.info("CAM Attention startup warmup CR done rank=%d", topology.world_rank) + store.set(f"attn/done/{topology.world_rank}", "ready") + self._wait_for_communication_warmup_done() + + def _wait_for_communication_warmup_done(self) -> None: + topology = self._spec.topology + self._wait_for_entries( + [f"attn/done/{rank}" for rank in range(topology.attn_size)] + + [ + f"ffn/warmup/done/{rank}" + for rank in range(topology.attn_size, topology.world_size) + ], + stage="warmup-done", + ) + + def _wait_for_ffn_ready(self) -> None: + store = self._get_store() + topology = self._spec.topology + keys = [ + f"ffn/{rank}" for rank in range(topology.attn_size, topology.world_size) + ] + deadline = time.monotonic() + STARTUP_TIMEOUT_SECONDS + while True: + self._check_failures(store, "ready") + pending = [] + for key in keys: + if not store.check([key]): + pending.append(key) + continue + status = store.get(key).decode() + if status != "ready": + raise RuntimeError(f"CAM FFN startup {key} {status}") + if not pending: + self._check_failures(store, "ready") + return + if time.monotonic() >= deadline: + raise TimeoutError( + f"Timed out waiting for CAM FFN ranks: {sorted(pending)}" + ) + time.sleep(STARTUP_POLL_SECONDS) + + def report_failure(self, error: Exception) -> None: + """Publish a sticky failure only if the nonce Store already exists.""" + if self._failure is not None: + return + self._failure = error + store = self._store + if store is None: + return + topology = self._spec.topology + key = ( + f"attn/failure/{topology.world_rank}" + if topology.role == "attention" + else f"ffn/{topology.world_rank}" + ) + try: + store.set(key, f"failed:{error}") + except Exception: + logger.warning("CAM startup failure publication failed: %s", key) + + def start_attention(self) -> None: + if not self._begin("attention"): + return + try: + if not self._connector.is_initialized: + self._connector.init_afd_connector() + self._get_store() + layer_idx = self._wait_for_ffn_modes() + if layer_idx is not None: + self._run_attention_warmup(layer_idx) + self._wait_for_ffn_ready() + self._started = True + logger.info( + "CAM Attention startup all FFN ready rank=%d mode=%s", + self._spec.topology.world_rank, + "eager" if layer_idx is None else f"graph:{layer_idx}", + ) + except Exception as exc: + self.report_failure(exc) + raise + finally: + self._starting = False + + def start_ffn( + self, + *, + prepare: Callable[[], FFNStartupPlan], + consume_warmup: Callable[[int], None], + capture: Callable[[], None], + start_receiver: Callable[[], None], + ) -> None: + if not self._begin("ffn"): + return + try: + plan = prepare() + if not self._connector.is_initialized: + self._connector.init_afd_connector() + store = self._get_store() + topology = self._spec.topology + mode = f"graph:{plan.first_layer_idx}" if plan.use_graph else "eager" + store.set(f"ffn/mode/{topology.world_rank}", mode) + observed_layer = self._wait_for_ffn_modes() + if observed_layer != plan.first_layer_idx: + raise RuntimeError( + f"CAM FFN startup rank={topology.world_rank} mode mismatch: " + f"local={mode}, all=graph:{observed_layer}" + ) + if plan.use_graph: + self._wait_for_entries( + [f"attn/prepared/{rank}" for rank in range(topology.attn_size)], + stage="warmup-prepared", + ) + store.set(f"ffn/warmup/start/{topology.world_rank}", "ready") + consume_warmup(self._spec.attention_dp_size) + store.set(f"ffn/warmup/done/{topology.world_rank}", "ready") + self._wait_for_communication_warmup_done() + capture() + start_receiver() + if self._failure is not None: + raise RuntimeError("CAM FFN receiver failed during startup") from ( + self._failure + ) + key = f"ffn/{topology.world_rank}" + status = store.compare_set(key, "", "ready").decode() + if status != "ready": + raise RuntimeError(f"CAM FFN startup {key} {status}") + self._started = True + logger.info( + "CAM FFN startup ready rank=%d mode=%s", topology.world_rank, mode + ) + except Exception as exc: + self.report_failure(exc) + raise + finally: + self._starting = False + + +__all__ = ["AsyncCamStartupCoordinator", "AsyncCamStartupSpec", "FFNStartupPlan"] diff --git a/afd_plugin/v1/worker/npu/attention_model_runner.py b/afd_plugin/v1/worker/npu/attention_model_runner.py index 8f586d04c..7efd8038b 100644 --- a/afd_plugin/v1/worker/npu/attention_model_runner.py +++ b/afd_plugin/v1/worker/npu/attention_model_runner.py @@ -85,7 +85,11 @@ AFDDPMetadata, AFDForwardContextMetadata, ) -from afd_plugin.connectors.npu.async_cam import AFDAsyncExtraInfo +from afd_plugin.connectors.npu.async_cam import ( + AFDAsyncExtraInfo, + CAMAsyncAFDConnector, +) +from afd_plugin.distributed.afd_process_group import ProcessGroupRendezvousContext from afd_plugin.model_executor.models.npu.async_cam_layout import ( ASYNC_MOE_UBATCH_METADATA_KEY, AsyncMoeUbatchMetadata, @@ -105,6 +109,10 @@ _resolve_world_ranks, build_ubatch_dp_metadata_list, ) +from afd_plugin.v1.worker.npu.async_cam_startup import ( + AsyncCamStartupCoordinator, + AsyncCamStartupSpec, +) from afd_plugin.v1.worker.npu.npu_ubatch_wrapper import AscendUBatchWrapper from afd_plugin.v1.worker.npu.ubatch_utils import ( check_enable_ubatch, @@ -139,12 +147,38 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device): ) rank, _ = _resolve_world_ranks() local_rank = int(device.index) - self.connector = AFDConnectorFactory.create_connector( - rank, - local_rank, - vllm_config, - self.afd_config, + rendezvous_context = ( + ProcessGroupRendezvousContext() + if afd_config.connector == AFD_ASYNC_CONNECTOR + else None ) + if rendezvous_context is None: + self.connector = AFDConnectorFactory.create_connector( + rank, local_rank, vllm_config, self.afd_config + ) + else: + self.connector = AFDConnectorFactory.create_connector( + rank, + local_rank, + vllm_config, + self.afd_config, + rendezvous_context=rendezvous_context, + ) + self._async_cam_startup: AsyncCamStartupCoordinator | None = None + if rendezvous_context is not None: + assert isinstance(self.connector, CAMAsyncAFDConnector) + self._async_cam_startup = AsyncCamStartupCoordinator( + self.connector, + rendezvous_context, + AsyncCamStartupSpec( + topology=self.connector.topology, + local_rank=local_rank, + tp_size=self.connector.tp_size, + hidden_size=vllm_config.model_config.hf_config.hidden_size, + topk=vllm_config.model_config.hf_config.num_experts_per_tok, + activation_dtype=vllm_config.model_config.dtype, + ), + ) self.afd_async_extra_info = AFDAsyncExtraInfo() if afd_config.connector == AFD_ASYNC_CONNECTOR: connector_extra_info = self.connector.extra_info @@ -1557,7 +1591,10 @@ def _build_capture_dp_metadata(self, num_tokens: int) -> DPMetadata | AFDDPMetad # Attention and FFN model loading overlap across roles. Connector # initialization is skipped when it already succeeded (the FFN side uses # the same guard), and the rendezvous completes before memory profiling or - # the first model forward. The connector itself is idempotent. + # the first model forward. Async CAM then waits for every FFN service to + # complete a controlled CAM communication warmup, finish capture, and start + # its receive loop before Attention can profile. + # The connector itself is idempotent. # Signature: matches upstream; no added parameters. def load_model(self) -> None: super().load_model() @@ -1567,8 +1604,13 @@ def load_model(self) -> None: # rendezvous is the blocking cross-role collective, so it is # deliberately last: it doubles as the "both roles finished loading # weights" barrier before memory profiling. - if not self.connector.is_initialized: + # ### PATCH START: CAM startup coordination + startup = self._async_cam_startup + if startup is not None: + startup.start_attention() + elif not self.connector.is_initialized: self.connector.init_afd_connector() + # ### PATCH END: CAM startup coordination def _install_ascend_ubatch_wrapper(self) -> None: if isinstance(self.model, AscendUBatchWrapper): @@ -1930,6 +1972,14 @@ def sync_and_slice_intermediate_tensors( def shutdown(self) -> None: stop_afd_npu_profiler(self.prof) + startup = self._async_cam_startup + if startup is not None and self.connector.is_initialized: + if not startup.started or startup.failed: + raise RuntimeError( + "CAM Attention startup did not complete safely; " + "retain communicator resources" + ) + torch.npu.synchronize() self.connector.close() super().shutdown() diff --git a/afd_plugin/v1/worker/npu/ffn_model_runner.py b/afd_plugin/v1/worker/npu/ffn_model_runner.py index 7e6ab1d39..ecbfe4127 100644 --- a/afd_plugin/v1/worker/npu/ffn_model_runner.py +++ b/afd_plugin/v1/worker/npu/ffn_model_runner.py @@ -26,7 +26,7 @@ step_afd_npu_profiler, stop_afd_npu_profiler, ) -from afd_plugin.config import AFDConfig, parse_afd_config +from afd_plugin.config import AFD_ASYNC_CONNECTOR, AFDConfig, parse_afd_config from afd_plugin.connectors import ( AFDConnectorFactory, AFDControlPayload, @@ -39,8 +39,11 @@ AFDAsyncTransferState, CAMAsyncAFDConnector, ) +from afd_plugin.distributed.afd_process_group import ProcessGroupRendezvousContext from afd_plugin.envs import async_cam_layered_gmm_enabled -from afd_plugin.model_executor.npu.async_cam_w4a8 import AsyncCAMW4A8Executor +from afd_plugin.model_executor.npu.async_cam_w4a8 import ( + AsyncCAMW4A8Executor, +) from afd_plugin.v1.worker.attention_metadata import ( _resolve_world_ranks, ) @@ -50,6 +53,11 @@ make_ffn_graph_key, ) from afd_plugin.v1.worker.ffn_model_runner import _set_moe_layer_index +from afd_plugin.v1.worker.npu.async_cam_startup import ( + AsyncCamStartupCoordinator, + AsyncCamStartupSpec, + FFNStartupPlan, +) if TYPE_CHECKING: from vllm.sequence import IntermediateTensors @@ -60,6 +68,7 @@ from afd_plugin.connectors import AFDConnectorBase logger = init_logger(__name__) +FFN_GRAPH_REPLAY_LOG_INTERVAL = 1000 class AFDNPUFFNModelRunner(NPUModelRunner): @@ -78,12 +87,38 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device) -> None: ) rank, _ = _resolve_world_ranks() local_rank = int(device.index) - self.connector = AFDConnectorFactory.create_connector( - rank, - local_rank, - vllm_config, - self.afd_config, + rendezvous_context = ( + ProcessGroupRendezvousContext() + if afd_config.connector == AFD_ASYNC_CONNECTOR + else None ) + if rendezvous_context is None: + self.connector = AFDConnectorFactory.create_connector( + rank, local_rank, vllm_config, self.afd_config + ) + else: + self.connector = AFDConnectorFactory.create_connector( + rank, + local_rank, + vllm_config, + self.afd_config, + rendezvous_context=rendezvous_context, + ) + self._async_cam_startup: AsyncCamStartupCoordinator | None = None + if rendezvous_context is not None: + assert isinstance(self.connector, CAMAsyncAFDConnector) + self._async_cam_startup = AsyncCamStartupCoordinator( + self.connector, + rendezvous_context, + AsyncCamStartupSpec( + topology=self.connector.topology, + local_rank=local_rank, + tp_size=self.connector.tp_size, + hidden_size=vllm_config.model_config.hf_config.hidden_size, + topk=vllm_config.model_config.hf_config.num_experts_per_tok, + activation_dtype=vllm_config.model_config.dtype, + ), + ) if ( isinstance(self.connector, CAMAsyncAFDConnector) and self.afd_config.compute_gate_on_attention @@ -102,13 +137,22 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device) -> None: self._is_shutdown = False self._layered_gmm_requested = async_cam_layered_gmm_enabled() self._layered_executor: AsyncCAMW4A8Executor | None = None + self._async_cam_ffn_graph_enabled = ( + isinstance(self.connector, CAMAsyncAFDConnector) + and not self.model_config.enforce_eager + and self.compilation_config.cudagraph_mode.name == "FULL" + ) + self._async_cam_ffn_graph: torch.npu.NPUGraph | None = None + self._async_cam_ffn_output: torch.Tensor | None = None + self._async_cam_ffn_replays = 0 @staticmethod def parse_config(vllm_config: VllmConfig) -> AFDConfig: return parse_afd_config(vllm_config, expected_role="ffn") def initialize_afd_connector(self) -> None: - self._initialize_layered_executor() + if self._layered_executor is None: + self._initialize_layered_executor() self.connector.init_afd_connector() def _initialize_layered_executor(self) -> None: @@ -125,9 +169,11 @@ def _initialize_layered_executor(self) -> None: raise ValueError("layered W4A8 GMM requires Ascend 910C/A3") if self.connector.dynamic_quant != 1: raise ValueError("layered W4A8 GMM requires dynamicQuant=1") - if self.use_aclgraph or self.vllm_config.use_v2_model_runner: + if self.vllm_config.use_v2_model_runner: + raise ValueError("layered W4A8 GMM requires ModelRunnerV1") + if self.use_aclgraph and not self._async_cam_ffn_graph_enabled: raise ValueError( - "layered W4A8 GMM requires eager ModelRunnerV1" + "layered W4A8 GMM graph requires FFN FULL mode" ) if tuple(sorted(layer.layer_idx for layer in layers)) != tuple( _ffn_layer_indices(self) @@ -149,6 +195,8 @@ def _initialize_layered_executor(self) -> None: layers[0].swiglu_limit != 0.0, ) return + if self._async_cam_ffn_graph_enabled: + raise ValueError(f"CAM async FFN FULL requires layered W4A8 GMM: {reason}") logger.info( "AFD_ASYNC_CAM_LAYERED_GMM requested=%s actual=legacy reason=%s", self._layered_gmm_requested, @@ -203,9 +251,125 @@ def execute_connector_driven_step(self) -> None: "AFD connector", ) step_afd_npu_profiler(self.prof) + if self._async_cam_ffn_graph_enabled: + connector = cast(CAMAsyncAFDConnector, self.connector) + graph = self._async_cam_ffn_graph + if graph is None: + raise RuntimeError( + "CAM async FFN graph was not captured before dispatch" + ) + layer_indices = _ffn_layer_indices(self) + previous_replays = self._async_cam_ffn_replays + for _ in layer_indices: + graph.replay() + self._async_cam_ffn_replays += 1 + if previous_replays == 0: + logger.info( + "CAM async FFN graph replay rank=%d count=%d", + connector.world_rank, + self._async_cam_ffn_replays, + ) + elif ( + self._async_cam_ffn_replays // FFN_GRAPH_REPLAY_LOG_INTERVAL + > previous_replays // FFN_GRAPH_REPLAY_LOG_INTERVAL + ): + logger.debug( + "CAM async FFN graph replay rank=%d count=%d", + connector.world_rank, + self._async_cam_ffn_replays, + ) + return None self._ffn_forward_connector_driven() return None + def prepare_async_cam_ffn_startup(self) -> FFNStartupPlan: + """Prepare the executor and select eager or one FULL graph.""" + if self._layered_executor is None: + self._initialize_layered_executor() + if not self._async_cam_ffn_graph_enabled: + return FFNStartupPlan(use_graph=False) + executor = self._layered_executor + if executor is None or not executor.layer_ids: + raise RuntimeError("CAM async FFN FULL requires layered GMM layers") + return FFNStartupPlan(use_graph=True, first_layer_idx=executor.layer_ids[0]) + + def warmup_async_cam_ffn_communication(self, count: int) -> None: + """Consume one controlled CAM work item from each Attention DP group.""" + if not self._async_cam_ffn_graph_enabled: + return + connector = cast(CAMAsyncAFDConnector, self.connector) + for _ in range(count): + self._execute_layered_work_item() + torch.npu.synchronize() + logger.info( + "CAM async FFN communication warmup done rank=%d items=%d", + connector.world_rank, + count, + ) + + def capture_async_cam_ffn_graph(self) -> None: + """Capture one DR, layered GMM, CS transaction for every MoE layer.""" + if ( + not self._async_cam_ffn_graph_enabled + or self._async_cam_ffn_graph is not None + ): + return + connector = cast(CAMAsyncAFDConnector, self.connector) + if self._layered_executor is None or not connector.is_initialized: + raise RuntimeError("CAM async FFN graph requires weights and HCCL group") + graph = torch.npu.NPUGraph() + logger.info("CAM async FFN graph capture start rank=%d", connector.world_rank) + with torch.npu.graph(graph, pool=self.graph_pool): + logger.info( + "CAM async FFN graph context entered rank=%d", + connector.world_rank, + ) + output = self._execute_layered_work_item(capture_trace=True) + self._async_cam_ffn_output = output + self._async_cam_ffn_graph = graph + logger.info( + "CAM async FFN graph capture complete rank=%d graphs=1", + connector.world_rank, + ) + + def release_async_cam_ffn_graph(self) -> None: + if self._async_cam_ffn_graph is None: + return + connector = cast(CAMAsyncAFDConnector, self.connector) + self._async_cam_ffn_graph = None + self._async_cam_ffn_output = None + logger.info( + "CAM async FFN graph released rank=%d replays=%d", + connector.world_rank, + self._async_cam_ffn_replays, + ) + + def _execute_layered_work_item( + self, *, capture_trace: bool = False + ) -> torch.Tensor: + connector = cast(CAMAsyncAFDConnector, self.connector) + executor = self._layered_executor + if executor is None: + raise RuntimeError("CAM async FFN layered executor is unavailable") + payload = connector.recv_attn_output() + if capture_trace: + logger.info("CAM async FFN graph DR captured rank=%d", connector.world_rank) + states = cast(AFDAsyncTransferState, payload.context.states) + output = executor( + payload.hidden_states, + states.dynamic_scales, + states.group_list, + states.token_nums_rankid_layeridx, + ) + if capture_trace: + logger.info( + "CAM async FFN graph GMM captured rank=%d", connector.world_rank + ) + connector.send_ffn_output(output, payload.context) + if capture_trace: + logger.info("CAM async FFN graph CS captured rank=%d", connector.world_rank) + return output + def execute_model( self, scheduler_output: SchedulerOutput | None = None, @@ -372,15 +536,7 @@ def _ffn_forward_connector_driven( connector = cast(CAMAsyncAFDConnector, self.connector) if self._layered_executor is not None: for _ in _ffn_layer_indices(self): - payload = connector.recv_attn_output() - states = cast(AFDAsyncTransferState, payload.context.states) - rank_ffn_output = self._layered_executor( - payload.hidden_states, - states.dynamic_scales, - states.group_list, - states.token_nums_rankid_layeridx, - ) - connector.send_ffn_output(rank_ffn_output, payload.context) + rank_ffn_output = self._execute_layered_work_item() return rank_ffn_output for _ in _ffn_layer_indices(self): @@ -531,8 +687,14 @@ def shutdown(self) -> None: if self._is_shutdown: return stop_afd_npu_profiler(self.prof) + if self._async_cam_startup is not None and self.connector.is_initialized: + raise RuntimeError( + "CAM FFN receiver must drain and close its connector before " + "model runner shutdown" + ) if self.connector.is_initialized: self.connector.close() + self.release_async_cam_ffn_graph() super().shutdown() self._is_shutdown = True diff --git a/afd_plugin/v1/worker/npu/ffn_worker.py b/afd_plugin/v1/worker/npu/ffn_worker.py index dded65435..69b818750 100644 --- a/afd_plugin/v1/worker/npu/ffn_worker.py +++ b/afd_plugin/v1/worker/npu/ffn_worker.py @@ -22,6 +22,7 @@ fix_all2all_backend_for_afd, npu_afd_num_ubatches, ) +from afd_plugin.connectors.npu.async_cam import CAMAsyncAFDConnector from afd_plugin.model_executor.models.model_utils import get_afd_model_config from afd_plugin.v1.worker.npu.ffn_model_runner import AFDNPUFFNModelRunner from afd_plugin.validation import ( @@ -38,6 +39,7 @@ logger = logging.getLogger(__name__) FFN_SHUTDOWN_TIMEOUT_SECONDS = 5 +FFN_STARTUP_THREAD_TIMEOUT_SECONDS = 30 class AFDNPUFFNWorker(NPUWorker): @@ -56,6 +58,8 @@ def __init__(self, *args: Any, **kwargs: Any) -> None: self._ffn_thread: threading.Thread | None = None self._ffn_shutdown_event: threading.Event | None = None self._ffn_loop_error: BaseException | None = None + self._ffn_loop_started_event: threading.Event | None = None + self._ffn_receiver_drained = False self._cpu_binding_attempted = False def init_device(self) -> None: @@ -91,7 +95,6 @@ def get_kv_cache_spec(self) -> dict[str, KVCacheSpec]: def initialize_from_config(self, kv_cache_config: KVCacheConfig) -> None: self.cache_config.num_gpu_blocks = kv_cache_config.num_blocks self.model_runner.initialize_kv_cache(kv_cache_config) - self.model_runner.initialize_afd_connector() self.start_ffn_server_loop() def compile_or_warm_up_model(self) -> CompilationTimes: @@ -107,39 +110,78 @@ def execute_model( ) def start_ffn_server_loop(self) -> None: + startup = self.model_runner._async_cam_startup if self._ffn_thread is not None and self._ffn_thread.is_alive(): self.raise_ffn_loop_error_if_any() + if startup is not None and startup.failed: + raise RuntimeError("CAM FFN startup previously failed") return - self.raise_ffn_loop_error_if_any() - connector = self.model_runner.connector - if not connector.is_initialized: + if startup is not None and self._ffn_thread is not None: + raise RuntimeError("CAM FFN receiver stopped; startup cannot be reused") + if startup is not None: + startup.start_ffn( + prepare=self.model_runner.prepare_async_cam_ffn_startup, + consume_warmup=self.model_runner.warmup_async_cam_ffn_communication, + capture=self.model_runner.capture_async_cam_ffn_graph, + start_receiver=self._start_ffn_receiver, + ) + return + if not self.model_runner.connector.is_initialized: self.model_runner.initialize_afd_connector() + self._start_ffn_receiver() + def _start_ffn_receiver(self) -> None: + """Start the receive thread and confirm that its NPU device is set.""" self._bind_cpus_once() self._ffn_shutdown_event = threading.Event() + started_event = threading.Event() + self._ffn_loop_started_event = started_event self._ffn_loop_error = None + self._ffn_receiver_drained = False + startup = self.model_runner._async_cam_startup def ffn_worker_loop() -> None: try: self._run_ffn_server_loop() except Exception as exc: shutdown_event = self._ffn_shutdown_event - if shutdown_event is not None and shutdown_event.is_set(): + if ( + startup is None + and shutdown_event is not None + and shutdown_event.is_set() + ): logger.debug( "AFD NPU FFN receive loop stopped during shutdown", exc_info=True, ) return self._ffn_loop_error = exc + started_event.set() logger.exception("AFD NPU FFN worker loop failed") + if startup is not None: + startup.report_failure(exc) self._ffn_thread = threading.Thread( target=ffn_worker_loop, name="afd-npu-ffn-worker-loop", daemon=True, ) - self._ffn_thread.start() + try: + self._ffn_thread.start() + if startup is not None: + if not started_event.wait(timeout=FFN_STARTUP_THREAD_TIMEOUT_SECONDS): + self.raise_ffn_loop_error_if_any() + raise TimeoutError("AFD NPU FFN service thread did not start") + self.raise_ffn_loop_error_if_any() + if not self._ffn_thread.is_alive(): + raise RuntimeError( + "AFD NPU FFN service thread stopped during startup" + ) + except Exception as exc: + if startup is not None: + startup.report_failure(exc) + raise def _bind_cpus_once(self) -> None: if self._cpu_binding_attempted: @@ -167,6 +209,10 @@ def _run_ffn_server_loop(self) -> None: return torch.npu.set_device(self.device) + if isinstance(self.model_runner.connector, CAMAsyncAFDConnector): + started_event = self._ffn_loop_started_event + if started_event is not None: + started_event.set() while not event.is_set(): if self.model_runner.connector.control_plane is None: self.model_runner.execute_connector_driven_step() @@ -188,11 +234,11 @@ def _run_ffn_server_loop(self) -> None: is_graph_replaying=is_graph_replaying, ) torch.npu.synchronize() + self._ffn_receiver_drained = True def raise_ffn_loop_error_if_any(self) -> None: error = self._ffn_loop_error if error is not None: - self._ffn_loop_error = None raise RuntimeError("AFD NPU FFN worker loop failed") from error def stop_ffn_server_loop(self) -> None: @@ -200,12 +246,39 @@ def stop_ffn_server_loop(self) -> None: if event is not None: event.set() - # CAM recv blocks in the connector operator. Release the communicator - # first so the daemon can observe the shutdown event, then wait for it - # before the parent runner releases model tensors. - self.model_runner.connector.close() + # A graph replay may be blocked in DR until Attention sends another + # work item. Do not destroy its communicator or graph-owned tensors + # while the worker thread is still using them. + connector = self.model_runner.connector + cam_connector = ( + connector if isinstance(connector, CAMAsyncAFDConnector) else None + ) thread = self._ffn_thread - if thread is not None: + if cam_connector is not None and thread is not None: + thread.join(timeout=FFN_SHUTDOWN_TIMEOUT_SECONDS) + if thread.is_alive(): + raise RuntimeError( + "AFD NPU FFN receive loop is still active; retain graph and " + "connector resources until the worker process exits" + ) + self.raise_ffn_loop_error_if_any() + if not self._ffn_receiver_drained: + raise RuntimeError( + "AFD NPU FFN receiver did not confirm device completion; " + "retain graph and connector resources" + ) + startup = self.model_runner._async_cam_startup + if cam_connector is not None and startup is not None and startup.failed: + raise RuntimeError( + "CAM FFN startup failed; retain communicator and graph resources" + ) + # Only a normally drained CAM receiver may release captured resources. + if cam_connector is not None: + self.model_runner.release_async_cam_ffn_graph() + cam_connector.close() + else: + connector.close() + if thread is not None and cam_connector is None: thread.join(timeout=FFN_SHUTDOWN_TIMEOUT_SECONDS) if thread.is_alive(): raise RuntimeError( @@ -213,6 +286,7 @@ def stop_ffn_server_loop(self) -> None: ) self._ffn_thread = None self._ffn_shutdown_event = None + self._ffn_loop_started_event = None self.raise_ffn_loop_error_if_any() def shutdown(self) -> None: diff --git a/tests/npu/async_cam_communication_graph_probe.py b/tests/npu/async_cam_communication_graph_probe.py new file mode 100644 index 000000000..89ab282de --- /dev/null +++ b/tests/npu/async_cam_communication_graph_probe.py @@ -0,0 +1,214 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the AFD plugin project +"""Two-rank CAM DR -> simple compute -> CS graph capture probe. + +Run with torchrun --standalone --nproc-per-node=2. Rank 0 is Attention and +rank 1 is FFN. ``cold`` captures immediately; ``warmup`` completes one eager +CAM transaction on both ranks before capture. Both modes send a matched task +for graph replay only after capture completes. +""" + +from __future__ import annotations + +import argparse +import os +from datetime import timedelta + +import torch +import torch.distributed as dist +import torch_npu # noqa: F401 - register torch.npu + +from afd_plugin.compat.npu.ops import ensure_cam_async_ops_available +from afd_plugin.distributed.afd_process_group import init_afd_process_group + +HIDDEN_SIZE = 512 +MAX_SEQUENCE_LENGTH = 262144 +MAX_SEQ_LEN = 1 +TOP_K = 1 +ATTN_RANKS = 1 +FFN_RANKS = 1 +EXPERTS_PER_RANK = 1 +DYNAMIC_QUANT = 1 +COMM_ID = 0 +RENDEZVOUS_TIMEOUT_SECONDS = 120 + + +def log(rank: int, phase: str) -> None: + print(f"rank={rank} {phase}", flush=True) + + +def attention_transaction( + *, rank: int, group_name: str, comm_args: torch.Tensor, anchor: torch.Tensor +) -> None: + ops = torch.ops.afd_ascend + hidden = torch.full( + (MAX_SEQ_LEN, HIDDEN_SIZE), 0.25, device="npu", dtype=torch.bfloat16 + ) + expert_ids = torch.zeros((MAX_SEQ_LEN, TOP_K), device="npu", dtype=torch.int32) + expert_weights = torch.ones((MAX_SEQ_LEN, TOP_K), device="npu", dtype=torch.float32) + log(rank, "DS start") + ops.afd_async_dispatch_send( + hidden, + expert_ids, + comm_args, + COMM_ID, + MAX_SEQ_LEN, + MAX_SEQ_LEN, + HIDDEN_SIZE, + TOP_K, + FFN_RANKS, + ATTN_RANKS, + EXPERTS_PER_RANK, + rank, + ATTN_RANKS + FFN_RANKS, + 0, + ATTN_RANKS, + DYNAMIC_QUANT, + group_name, + ) + log(rank, "DS returned; CR start") + output = ops.afd_async_combine_recv( + anchor, + expert_ids, + expert_weights, + comm_args, + COMM_ID, + MAX_SEQ_LEN, + HIDDEN_SIZE, + TOP_K, + FFN_RANKS, + ATTN_RANKS, + EXPERTS_PER_RANK, + rank, + ATTN_RANKS + FFN_RANKS, + group_name, + ) + torch.npu.synchronize() + torch.testing.assert_close(output, hidden, atol=0.01, rtol=0.01) + log(rank, "CR complete") + + +def ffn_transaction( + *, rank: int, group_name: str, comm_args: torch.Tensor, anchor: torch.Tensor +) -> tuple[torch.Tensor, ...]: + ops = torch.ops.afd_ascend + log(rank, "DR start") + expanded, scales, batch_info, counts = ops.afd_async_dispatch_recv( + anchor, + comm_args, + COMM_ID, + MAX_SEQ_LEN, + HIDDEN_SIZE, + TOP_K, + FFN_RANKS, + ATTN_RANKS, + EXPERTS_PER_RANK, + rank, + ATTN_RANKS + FFN_RANKS, + ATTN_RANKS, + DYNAMIC_QUANT, + group_name, + ) + log(rank, "DR returned; compute start") + output = (expanded.float() * scales.unsqueeze(1)).to(torch.bfloat16) + log(rank, "compute returned; CS start") + ops.afd_async_combine_send( + output, + comm_args, + batch_info, + COMM_ID, + MAX_SEQ_LEN, + HIDDEN_SIZE, + TOP_K, + FFN_RANKS, + ATTN_RANKS, + EXPERTS_PER_RANK, + rank, + ATTN_RANKS + FFN_RANKS, + ATTN_RANKS, + group_name, + ) + log(rank, "CS returned") + # Keep the graph's intermediate output and the receive metadata alive. + return output, expanded, scales, batch_info, counts + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--mode", choices=("cold", "warmup"), required=True) + parser.add_argument("--cam-port", type=int, required=True) + args = parser.parse_args() + rank = int(os.environ["RANK"]) + local_rank = int(os.environ["LOCAL_RANK"]) + if int(os.environ["WORLD_SIZE"]) != ATTN_RANKS + FFN_RANKS: + raise ValueError("Probe requires exactly two ranks") + os.environ["BATCH_SIZE_FACTOR"] = str(1 / MAX_SEQUENCE_LENGTH) + torch.npu.set_device(local_rank) + ensure_cam_async_ops_available() + dist.init_process_group( + "gloo", timeout=timedelta(seconds=RENDEZVOUS_TIMEOUT_SECONDS) + ) + # torchrun sets this for its own agent store. The CAM tcp:// rendezvous + # needs its own rank-0 server on --cam-port. + os.environ["TORCHELASTIC_USE_AGENT_STORE"] = "False" + cam_group = init_afd_process_group( + backend="hccl", + init_method=f"tcp://127.0.0.1:{args.cam_port}", + world_size=ATTN_RANKS + FFN_RANKS, + rank=rank, + group_name="afd_async_cam_communication_graph_probe", + timeout=timedelta(seconds=RENDEZVOUS_TIMEOUT_SECONDS), + ) + group_name = str( + cam_group._get_backend(torch.device("npu")).get_hccl_comm_name(rank) + ) + comm_args = torch.empty((1,), dtype=torch.float16, device="npu") + anchor = torch.empty((1,), dtype=torch.bfloat16, device="npu") + log(rank, "CAM HCCL group initialized") + + with torch.inference_mode(): + if args.mode == "warmup": + log(rank, "eager warmup start") + if rank == 0: + attention_transaction( + rank=rank, group_name=group_name, comm_args=comm_args, anchor=anchor + ) + else: + ffn_transaction( + rank=rank, group_name=group_name, comm_args=comm_args, anchor=anchor + ) + torch.npu.synchronize() + log(rank, "eager warmup complete") + dist.barrier() + + graph = None + references = None + if rank == ATTN_RANKS: + graph = torch.npu.NPUGraph() + log(rank, "graph capture start") + with torch.npu.graph(graph): + log(rank, "graph context entered") + references = ffn_transaction( + rank=rank, group_name=group_name, comm_args=comm_args, anchor=anchor + ) + log(rank, "graph capture complete") + dist.barrier() + + log(rank, "matched replay transaction start") + if rank == 0: + attention_transaction( + rank=rank, group_name=group_name, comm_args=comm_args, anchor=anchor + ) + else: + assert graph is not None and references is not None + graph.replay() + torch.npu.synchronize() + log(rank, "graph replay complete") + dist.barrier() + + dist.destroy_process_group(cam_group) + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/tests/npu/async_cam_ffn_graph_service_smoke.py b/tests/npu/async_cam_ffn_graph_service_smoke.py new file mode 100644 index 000000000..bcf0ea3be --- /dev/null +++ b/tests/npu/async_cam_ffn_graph_service_smoke.py @@ -0,0 +1,132 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the AFD plugin project +"""Reproducible prefill-only API smoke for the AsyncCam FFN graph deployment.""" + +from __future__ import annotations + +import argparse +import json +import time +from concurrent.futures import ThreadPoolExecutor, as_completed +from datetime import datetime, timezone +from pathlib import Path +from threading import Event +from typing import Any +from urllib.error import HTTPError, URLError +from urllib.request import Request, urlopen + +DEFAULT_TIMEOUT_SECONDS = 120 +CHUNK_SIZE = 8192 +CONCURRENT_REQUESTS = 4 + + +def send_completion( + endpoint: str, model: str, name: str, prompt: str, timeout: int +) -> dict[str, Any]: + payload = json.dumps( + {"model": model, "prompt": prompt, "max_tokens": 1, "temperature": 0} + ).encode() + request = Request( + endpoint, + data=payload, + headers={"Content-Type": "application/json"}, + ) + started_at = datetime.now(timezone.utc).isoformat() + start = time.monotonic() + result: dict[str, Any] = {"name": name, "started_at": started_at} + try: + with urlopen(request, timeout=timeout) as response: + body = json.loads(response.read()) + usage = body.get("usage") or {} + result.update( + { + "ok": bool(body.get("choices")) and usage.get("completion_tokens") == 1, + "request_id": body.get("id"), + "usage": usage, + "completion": [ + choice.get("text") for choice in body.get("choices", []) + ], + } + ) + except (HTTPError, URLError, TimeoutError, ValueError) as exc: + result.update({"ok": False, "error": str(exc)}) + result["elapsed_seconds"] = round(time.monotonic() - start, 3) + return result + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--endpoint", default="http://127.0.0.1:17100/v1/completions") + parser.add_argument("--model", default="dsv4") + parser.add_argument("--timeout", type=int, default=DEFAULT_TIMEOUT_SECONDS) + parser.add_argument("--output", type=Path, required=True) + args = parser.parse_args() + + # Unique numbers avoid tokenizer compression of a repeated word and cross + # the deployed 8192-token prefill chunk boundary. Verify from API usage. + long_prompt = "List the next number after: " + " ".join( + f"{number:05d}" for number in range(8500) + ) + medium_prompt = "Continue this sequence: " + " ".join( + f"{number:05d}" for number in range(1000) + ) + cases = [ + ("short", "The capital of France is"), + ("medium", medium_prompt), + ("over_chunk", long_prompt), + ] + args.output.parent.mkdir(parents=True, exist_ok=True) + results = [] + with args.output.open("w") as output_file: + + def record(result: dict[str, Any]) -> None: + results.append(result) + line = json.dumps(result, ensure_ascii=False, sort_keys=True) + print(line, flush=True) + output_file.write(line + "\n") + output_file.flush() + + for name, prompt in cases: + result = send_completion( + args.endpoint, args.model, name, prompt, args.timeout + ) + record(result) + if not result["ok"]: + return 1 + if name == "over_chunk" and ( + (result.get("usage") or {}).get("prompt_tokens", 0) <= CHUNK_SIZE + ): + return 1 + + start_concurrent = Event() + + def concurrent_case(index: int) -> dict[str, Any]: + start_concurrent.wait() + return send_completion( + args.endpoint, + args.model, + f"concurrent_{index}", + f"Continue: {index}. " + medium_prompt, + args.timeout, + ) + + with ThreadPoolExecutor(max_workers=CONCURRENT_REQUESTS) as executor: + futures = [ + executor.submit(concurrent_case, index) + for index in range(CONCURRENT_REQUESTS) + ] + start_concurrent.set() + for future in as_completed(futures): + record(future.result()) + + over_chunk = next(result for result in results if result["name"] == "over_chunk") + prompt_tokens = (over_chunk.get("usage") or {}).get("prompt_tokens", 0) + return ( + 0 + if all(result["ok"] for result in results) and prompt_tokens > CHUNK_SIZE + else 1 + ) + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tests/npu/async_cam_graph_bootstrap_probe.py b/tests/npu/async_cam_graph_bootstrap_probe.py new file mode 100644 index 000000000..89eaaa81b --- /dev/null +++ b/tests/npu/async_cam_graph_bootstrap_probe.py @@ -0,0 +1,97 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the AFD plugin project +"""Two-rank CAM process-group graph bootstrap probe without CAM work items. + +Launch with ``torchrun --nproc-per-node=2`` and pass a free ``--cam-port``. +The default torchrun rendezvous port must be different from ``--cam-port``. +This isolates HCCL startup work from dispatch-recv capture behavior. +""" + +from __future__ import annotations + +import argparse +import os +from datetime import timedelta + +import torch +import torch.distributed as dist +import torch_npu # noqa: F401 - register torch.npu before the probe uses it +from torch.distributed.distributed_c10d import Store + +from afd_plugin.distributed.afd_process_group import init_afd_process_group + +PROBE_TIMEOUT_SECONDS = 120 +NONCE_BYTES = 16 + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--cam-port", type=int, required=True) + parser.add_argument("--skip-nonce", action="store_true") + args = parser.parse_args() + + rank = int(os.environ["RANK"]) + local_rank = int(os.environ["LOCAL_RANK"]) + world_size = int(os.environ["WORLD_SIZE"]) + if world_size != 2: + raise ValueError("This probe requires exactly two ranks") + + torch.npu.set_device(local_rank) + dist.init_process_group( + backend="gloo", timeout=timedelta(seconds=PROBE_TIMEOUT_SECONDS) + ) + # This second rendezvous owns its TCPStore; torchrun's agent serves only + # the default group port, so rank zero must create the CAM store here. + os.environ["TORCHELASTIC_USE_AGENT_STORE"] = "False" + rendezvous_store: Store | None = None + + def retain_store(store: Store) -> None: + nonlocal rendezvous_store + rendezvous_store = store + + cam_group = init_afd_process_group( + backend="hccl", + init_method=f"tcp://127.0.0.1:{args.cam_port}", + world_size=world_size, + rank=rank, + group_name="afd_async_cam_graph_bootstrap_probe", + timeout=timedelta(seconds=PROBE_TIMEOUT_SECONDS), + on_rendezvous=retain_store, + ) + if rendezvous_store is None: + raise RuntimeError("CAM rendezvous store was not retained") + print(f"rank={rank} HCCL group initialized", flush=True) + + if not args.skip_nonce: + nonce = torch.zeros(NONCE_BYTES, dtype=torch.uint8, device=f"npu:{local_rank}") + if rank == 0: + nonce.copy_( + torch.tensor( + list(os.urandom(NONCE_BYTES)), + dtype=torch.uint8, + device=nonce.device, + ) + ) + dist.broadcast(nonce, src=0, group=cam_group) + print(f"rank={rank} nonce broadcast returned", flush=True) + torch.npu.synchronize() + print(f"rank={rank} nonce NPU synchronize returned", flush=True) + + graph = torch.npu.NPUGraph() + value = torch.ones((1,), device=f"npu:{local_rank}") + print(f"rank={rank} graph capture start", flush=True) + with torch.npu.graph(graph): + print(f"rank={rank} graph context entered", flush=True) + result = value + 1 + print(f"rank={rank} graph capture complete", flush=True) + graph.replay() + torch.npu.synchronize() + if result.item() != 2: + raise AssertionError(f"Unexpected graph result on rank {rank}") + print(f"rank={rank} graph replay complete", flush=True) + dist.destroy_process_group(cam_group) + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/tests/npu/test_async_cam_layered_graph_probe.py b/tests/npu/test_async_cam_layered_graph_probe.py new file mode 100644 index 000000000..4d6619941 --- /dev/null +++ b/tests/npu/test_async_cam_layered_graph_probe.py @@ -0,0 +1,125 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the AFD plugin project +"""Opt-in device probe for reusing one layered W4A8 computation graph. + +This isolates the device-controlled layer and row-count inputs. Full AsyncCam +DR/GMM/CS graph capture requires a running Attention/FFN service and is tested +separately on the complete deployment. +""" + +from __future__ import annotations + +import os + +import pytest + +pytestmark = pytest.mark.npu + + +def test_layered_w4a8_graph_reuses_device_metadata(): + if os.environ.get("AFD_RUN_ASCEND_OP_RUNTIME") != "1": + pytest.skip("requires opt-in 910C runtime") + + torch = pytest.importorskip("torch") + torch_npu = pytest.importorskip("torch_npu") + from afd_plugin.compat.npu.ops import ensure_cam_async_ops_available + from afd_plugin.model_executor.npu.async_cam_w4a8 import ( + AsyncCAMW4A8Executor, + W4A8LayerWeights, + ) + + ensure_cam_async_ops_available() + torch.npu.set_device(int(os.environ.get("LOCAL_RANK", "0"))) + torch.npu.config.allow_internal_format = True + torch.manual_seed(19) + + experts = 2 + hidden = 256 + intermediate = 256 + capacity = 32 + + def make_weight(k: int, n: int, factor: float, squeeze: bool): + values = torch.randint(-7, 8, (experts, k, n), dtype=torch.int8) + pairs = values.reshape(-1, 2) + packed = ( + torch.bitwise_or( + torch.bitwise_left_shift(pairs[:, 1], 4), + torch.bitwise_and(pairs[:, 0], 0x0F), + ) + .reshape(experts, k, n // 2) + .clone() + ) + weight = torch_npu.npu_format_cast(packed.npu(), 29).view(torch.int32) + scale = torch.full((experts, 1, n), factor, dtype=torch.float16).float() + compensation = 8 * (values.float() * scale).sum(dim=1) + encoded_scale = scale.view(torch.int32).to(torch.int64) + if squeeze: + encoded_scale = encoded_scale.squeeze(1) + return weight.contiguous(), encoded_scale.npu(), compensation.npu() + + layers = [] + for layer_idx, factor in ((2, 0.01), (5, 0.04)): + w13, s13, b13 = make_weight(hidden, 2 * intermediate, factor, True) + w2, s2, b2 = make_weight(intermediate, hidden, factor, False) + layers.append( + W4A8LayerWeights(layer_idx, w13, w2, s13, s2, b13, b2, True, 0.0, 1.0) + ) + executor = AsyncCAMW4A8Executor(layers) + + static_hidden = torch.zeros((capacity, hidden), dtype=torch.int8, device="npu") + static_scales = torch.full((capacity,), 0.005, device="npu") + static_counts = torch.tensor([1, 0], dtype=torch.int64, device="npu") + static_info = torch.tensor([capacity * 2, 0, 2, 1], dtype=torch.int64, device="npu") + + # Initialize the custom operators before graph capture without changing + # the addresses used by the graph itself. + executor(static_hidden, static_scales, static_counts, static_info) + torch.npu.synchronize() + graph = torch.npu.NPUGraph() + with torch.npu.graph(graph): + captured_output = executor( + static_hidden, static_scales, static_counts, static_info + ) + + shared_hidden = torch.randint(-100, 101, (capacity, hidden), dtype=torch.int8) + layer_outputs = {} + for layer_idx, counts in ( + (2, (1, 0)), + (5, (1, 0)), + (2, (0, 7)), + (5, (16, 16)), + (2, (0, 0)), + ): + valid_rows = sum(counts) + static_hidden.copy_(shared_hidden.npu()) + static_counts.copy_(torch.tensor(counts, dtype=torch.int64).npu()) + static_info.copy_( + torch.tensor( + [capacity * 2, 0, layer_idx, valid_rows], dtype=torch.int64 + ).npu() + ) + graph.replay() + torch.npu.synchronize() + assert captured_output.shape == (capacity, hidden) + graphed = captured_output[:valid_rows].cpu().float() + eager = ( + executor(static_hidden, static_scales, static_counts, static_info)[ + :valid_rows + ] + .cpu() + .float() + ) + if valid_rows: + difference = (graphed - eager).abs() + print( + f"layer={layer_idx} counts={counts} " + f"max_abs={difference.max().item():.6f} " + f"mean_abs={difference.mean().item():.6f}" + ) + torch.testing.assert_close(graphed, eager, rtol=0.04, atol=0.05) + if counts == (1, 0): + layer_outputs[layer_idx] = eager + + assert not torch.allclose( + layer_outputs[2], layer_outputs[5], rtol=0.04, atol=0.05 + ), "synthetic layers must have distinguishable outputs for the same input" diff --git a/tests/unit/connectors/test_async_cam_connector.py b/tests/unit/connectors/test_async_cam_connector.py index fea5658d3..71fbf4317 100644 --- a/tests/unit/connectors/test_async_cam_connector.py +++ b/tests/unit/connectors/test_async_cam_connector.py @@ -33,6 +33,9 @@ CAMAsyncAFDConnector, build_async_topology, ) +from afd_plugin.distributed.afd_process_group import ( # noqa: E402 + ProcessGroupRendezvousContext, +) class _FakeTensorLike: @@ -353,6 +356,7 @@ def fake_init_afd_process_group(**kwargs): "timeout": calls[0]["timeout"], "init_method": "tcp://127.0.0.1:1239", "pg_options": pg_options, + "on_rendezvous": None, }, ] assert connector.cam_pg is not None @@ -511,6 +515,84 @@ def test_async_ffn_side_dispatch_recv_and_combine_send(monkeypatch): assert logs[index][2]["max_seq_len"] == connector.max_num_batched_tokens +def test_startup_warmup_uses_real_connector_fifo_then_formal_task(monkeypatch): + from afd_plugin.v1.worker.npu import async_cam_startup + + fake_torch = _FakeTorch() + routed_ids = [] + synchronizations = [] + fake_torch.full = lambda shape, value, *, dtype, device: _FakeTensor( + shape, dtype=dtype, device=device + ) + + def make_ids(values, *, dtype, device): + routed_ids.append(values) + return _FakeTensor((len(values), len(values[0])), dtype=dtype, device=device) + + fake_torch.tensor = make_ids + fake_torch.npu = SimpleNamespace( + synchronize=lambda: synchronizations.append("done") + ) + monkeypatch.setattr(async_cam_module, "torch", fake_torch) + monkeypatch.setattr(async_cam_startup, "torch", fake_torch) + connector = CAMAsyncAFDConnector( + 0, 0, _vllm_config(), _afd_config(role="attention"), 0 + ) + connector._initialized = True + connector.comm_args = _FakeTensor((1,), dtype="fp16") + + class ReadyStore: + values = { + **{f"ffn/warmup/start/{rank}": b"ready" for rank in (4, 5)}, + **{f"ffn/warmup/done/{rank}": b"ready" for rank in (4, 5)}, + **{f"attn/done/{rank}": b"ready" for rank in (1, 2, 3)}, + } + + def set(self, key, value): + self.values[key] = value.encode() + + def check(self, keys): + return all(key in self.values for key in keys) + + def get(self, key): + return self.values[key] + + context = ProcessGroupRendezvousContext() + store = ReadyStore() + context.retain_store(store) + context.bind(SimpleNamespace()) + startup = async_cam_startup.AsyncCamStartupCoordinator( + connector, + context, + async_cam_startup.AsyncCamStartupSpec( + topology=connector.topology, + local_rank=0, + tp_size=connector.tp_size, + hidden_size=connector.hidden_size, + topk=connector.topk, + activation_dtype=fake_torch.bfloat16, + ), + ) + startup._store = store + startup._run_attention_warmup(layer_idx=3) + + assert routed_ids == [[[0, 4]]] + assert synchronizations == ["done"] + assert connector._pending_attention_payloads == {} + assert store.values["attn/done/0"] == b"ready" + + formal_context = AFDTransferContext( + metadata=AFDTransferMetadata.create_attention_metadata( + layer_idx=5, stage_idx=0, seq_len=1 + ) + ) + formal_hidden = _FakeTensor((1, 16)) + connector.send_attn_output(formal_hidden, formal_context, **_topk_payload(1)) + assert connector._pending_attention_payloads[0][0][0] is formal_context + connector.recv_ffn_output(formal_hidden, ubatch_idx=0) + assert connector._pending_attention_payloads == {} + + def test_async_combine_send_requires_dispatch_recv_token_metadata(monkeypatch): fake_torch = _FakeTorch() monkeypatch.setattr(async_cam_module, "torch", fake_torch) @@ -851,3 +933,71 @@ def test_close_releases_every_stage_routing(monkeypatch): assert set(connector._pending_attention_payloads) == {0, 1} connector.close() assert connector._pending_attention_payloads == {} + + +def test_async_factory_injects_context_only_into_cam_connector(): + context = ProcessGroupRendezvousContext() + connector = AFDConnectorFactory.create_connector( + 0, + 0, + _vllm_config(), + _afd_config(role="attention"), + rendezvous_context=context, + ) + assert isinstance(connector, CAMAsyncAFDConnector) + assert connector._rendezvous_context is context + with pytest.raises(RuntimeError, match="not bound"): + context.borrow() + + with pytest.raises(TypeError, match="requires a CAMAsyncAFDConnector"): + AFDConnectorFactory.create_connector( + 0, + 0, + _vllm_config(), + AFDConfig(connector="CAMP2pAFDConnector", role="attention"), + rendezvous_context=context, + ) + + +def test_async_connector_binds_context_only_after_anchor_init(monkeypatch): + context = ProcessGroupRendezvousContext() + store = object() + backend = SimpleNamespace(get_hccl_comm_name=lambda rank: "cam-test") + process_group = SimpleNamespace(_get_backend=lambda device: backend) + fake_torch = _FakeTorch() + monkeypatch.setattr(async_cam_module, "torch", fake_torch) + monkeypatch.setattr( + async_cam_module, "ensure_cam_async_ops_available", lambda: None + ) + + def fake_init_afd_process_group(**kwargs): + kwargs["on_rendezvous"](store) + with pytest.raises(RuntimeError, match="not bound"): + context.borrow() + return process_group + + monkeypatch.setattr( + async_cam_module, "init_afd_process_group", fake_init_afd_process_group + ) + monkeypatch.setattr( + async_cam_module, "create_hccl_process_group_options", lambda _: None + ) + monkeypatch.setattr( + async_cam_module.dist, + "destroy_process_group", + lambda group: None, + ) + connector = CAMAsyncAFDConnector( + 0, + 0, + _vllm_config(), + _afd_config(role="attention"), + 0, + rendezvous_context=context, + ) + + connector.init_afd_connector() + assert context.borrow() == (store, process_group) + connector.close() + with pytest.raises(RuntimeError, match="not bound"): + context.borrow() diff --git a/tests/unit/v1/worker/test_async_cam_startup.py b/tests/unit/v1/worker/test_async_cam_startup.py new file mode 100644 index 000000000..a2cb84b48 --- /dev/null +++ b/tests/unit/v1/worker/test_async_cam_startup.py @@ -0,0 +1,237 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the AFD plugin project +"""Host-side Async CAM startup ordering and failure tests.""" + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest + +pytest.importorskip("torch_npu") + +from afd_plugin.connectors.npu.async_cam import AFDAsyncTopology # noqa: E402 +from afd_plugin.distributed.afd_process_group import ( # noqa: E402 + ProcessGroupRendezvousContext, +) +from afd_plugin.v1.worker.npu import async_cam_startup # noqa: E402 +from afd_plugin.v1.worker.npu.async_cam_startup import ( # noqa: E402 + AsyncCamStartupCoordinator, + AsyncCamStartupSpec, + FFNStartupPlan, +) + + +class FakeStore: + def __init__(self, values: dict[str, bytes] | None = None): + self.values = dict(values or {}) + + def set(self, key: str, value: str) -> None: + self.values[key] = value.encode() + + def check(self, keys: list[str]) -> bool: + return all(key in self.values for key in keys) + + def get(self, key: str) -> bytes: + return self.values[key] + + def compare_set(self, key: str, expected: str, desired: str) -> bytes: + current = self.values.get(key, b"") + if current == expected.encode(): + current = desired.encode() + self.values[key] = current + return current + + +def make_coordinator( + *, role: str, attn_size: int = 1, ffn_size: int = 1, store: FakeStore +) -> tuple[AsyncCamStartupCoordinator, SimpleNamespace]: + world_rank = 0 if role == "attention" else attn_size + topology = AFDAsyncTopology( + role=role, + role_rank=0, + world_rank=world_rank, + attn_size=attn_size, + ffn_size=ffn_size, + expert_per_rank=1, + ) + connector = SimpleNamespace(is_initialized=True) + context = ProcessGroupRendezvousContext() + context.retain_store(store) + context.bind(SimpleNamespace()) + coordinator = AsyncCamStartupCoordinator( + connector, + context, + AsyncCamStartupSpec( + topology=topology, + local_rank=0, + tp_size=1, + hidden_size=16, + topk=1, + activation_dtype=async_cam_startup.torch.bfloat16, + ), + ) + coordinator._store = store + return coordinator, connector + + +def test_mode_disagreement_prevents_receiver_and_capture(): + store = FakeStore({"ffn/mode/2": b"eager"}) + coordinator, _ = make_coordinator(role="ffn", attn_size=1, ffn_size=2, store=store) + calls = [] + with pytest.raises(RuntimeError, match="modes disagree"): + coordinator.start_ffn( + prepare=lambda: FFNStartupPlan(use_graph=True, first_layer_idx=3), + consume_warmup=lambda count: calls.append("warmup"), + capture=lambda: calls.append("capture"), + start_receiver=lambda: calls.append("receiver"), + ) + assert calls == [] + assert store.values["ffn/1"].startswith(b"failed:") + + +def test_all_attention_and_ffn_warmup_must_finish_before_capture(monkeypatch): + store = FakeStore( + { + "attn/prepared/0": b"ready", + "attn/done/0": b"ready", + "ffn/mode/2": b"graph:3", + } + ) + coordinator, _ = make_coordinator(role="ffn", attn_size=1, ffn_size=2, store=store) + events = [] + + def finish_peer(_: float) -> None: + assert events == ["prepare", "warmup:1"] + store.set("ffn/warmup/done/2", "ready") + + monkeypatch.setattr(async_cam_startup.time, "sleep", finish_peer) + coordinator.start_ffn( + prepare=lambda: ( + events.append("prepare") + or FFNStartupPlan(use_graph=True, first_layer_idx=3) + ), + consume_warmup=lambda count: events.append(f"warmup:{count}"), + capture=lambda: events.append("capture"), + start_receiver=lambda: events.append("receiver"), + ) + assert events == ["prepare", "warmup:1", "capture", "receiver"] + assert store.values["ffn/1"] == b"ready" + assert coordinator.started + + +@pytest.mark.parametrize("failure_phase", ["capture", "receiver"]) +def test_capture_or_receiver_failure_is_sticky(failure_phase): + store = FakeStore({"attn/prepared/0": b"ready", "attn/done/0": b"ready"}) + coordinator, _ = make_coordinator(role="ffn", store=store) + + def fail_if_selected(phase: str) -> None: + if phase == failure_phase: + raise RuntimeError(f"{phase} failed") + + with pytest.raises(RuntimeError, match=f"{failure_phase} failed"): + coordinator.start_ffn( + prepare=lambda: FFNStartupPlan(use_graph=True, first_layer_idx=3), + consume_warmup=lambda count: None, + capture=lambda: fail_if_selected("capture"), + start_receiver=lambda: fail_if_selected("receiver"), + ) + assert store.values["ffn/1"].startswith(b"failed:") + assert not coordinator.started + with pytest.raises(RuntimeError, match="previously failed"): + coordinator.start_ffn( + prepare=lambda: FFNStartupPlan(use_graph=True, first_layer_idx=3), + consume_warmup=lambda count: None, + capture=lambda: None, + start_receiver=lambda: None, + ) + + +def test_ready_cannot_replace_failed_status(): + class RacingStore(FakeStore): + def compare_set(self, key: str, expected: str, desired: str) -> bytes: + self.set(key, "failed:loop stopped") + return self.values[key] + + store = RacingStore() + coordinator, _ = make_coordinator(role="ffn", store=store) + with pytest.raises(RuntimeError, match="loop stopped"): + coordinator.start_ffn( + prepare=lambda: FFNStartupPlan(use_graph=False), + consume_warmup=lambda count: None, + capture=lambda: None, + start_receiver=lambda: None, + ) + assert store.values["ffn/1"].startswith(b"failed:") + + +def test_attention_rereads_ready_ranks_for_later_failure(monkeypatch): + store = FakeStore({"ffn/1": b"ready"}) + coordinator, _ = make_coordinator( + role="attention", attn_size=1, ffn_size=2, store=store + ) + + def fail_after_first_poll(_: float) -> None: + store.set("ffn/1", "failed:receiver stopped") + store.set("ffn/2", "ready") + + monkeypatch.setattr(async_cam_startup.time, "sleep", fail_after_first_poll) + with pytest.raises(RuntimeError, match="receiver stopped"): + coordinator._wait_for_ffn_ready() + + +def test_attention_start_does_not_return_before_ffn_ready(monkeypatch): + store = FakeStore({"ffn/mode/1": b"eager"}) + coordinator, _ = make_coordinator(role="attention", store=store) + polls = [] + + def publish_after_wait(_: float) -> None: + assert not coordinator.started + polls.append(True) + store.set("ffn/1", "ready") + + monkeypatch.setattr(async_cam_startup.time, "sleep", publish_after_wait) + coordinator.start_attention() + assert polls == [True] + assert coordinator.started + + +def test_prepare_failure_before_store_does_not_start_nonce_collective(monkeypatch): + store = FakeStore() + coordinator, connector = make_coordinator(role="ffn", store=store) + connector.is_initialized = False + connector.init_afd_connector = lambda: pytest.fail("group was initialized") + coordinator._store = None + monkeypatch.setattr( + coordinator, + "_get_store", + lambda: pytest.fail("nonce collective was started"), + ) + with pytest.raises(RuntimeError, match="weights unavailable"): + coordinator.start_ffn( + prepare=lambda: (_ for _ in ()).throw(RuntimeError("weights unavailable")), + consume_warmup=lambda count: None, + capture=lambda: None, + start_receiver=lambda: None, + ) + assert coordinator.failed + assert store.values == {} + + +def test_repeated_eager_start_is_idempotent_and_closed_context_is_rejected(): + store = FakeStore() + coordinator, connector = make_coordinator(role="ffn", store=store) + events = [] + connector.init_afd_connector = lambda: events.append("init") + callbacks = dict( + prepare=lambda: events.append("prepare") or FFNStartupPlan(False), + consume_warmup=lambda count: events.append("warmup"), + capture=lambda: events.append("capture"), + start_receiver=lambda: events.append("receiver"), + ) + coordinator.start_ffn(**callbacks) + coordinator.start_ffn(**callbacks) + assert events == ["prepare", "receiver"] + coordinator._rendezvous_context.invalidate() + with pytest.raises(RuntimeError, match="not bound"): + coordinator.start_ffn(**callbacks) diff --git a/tests/unit/v1/worker/test_npu_runtime.py b/tests/unit/v1/worker/test_npu_runtime.py index d63cfefe8..e74897c6f 100644 --- a/tests/unit/v1/worker/test_npu_runtime.py +++ b/tests/unit/v1/worker/test_npu_runtime.py @@ -446,6 +446,11 @@ def _new_ffn_runner(): runner = object.__new__(AFDNPUFFNModelRunner) runner._layered_executor = None runner._layered_gmm_requested = False + runner._async_cam_ffn_graph_enabled = False + runner._async_cam_ffn_graph = None + runner._async_cam_ffn_output = None + runner._async_cam_ffn_replays = 0 + runner._async_cam_startup = None runner.prof = None runner.device = SimpleNamespace(type="npu") runner._is_shutdown = False @@ -468,6 +473,7 @@ def _new_ffn_worker(): worker = object.__new__(AFDNPUFFNWorker) worker._ffn_loop_error = None + worker._ffn_receiver_drained = False # Most tests exercise the daemon loop rather than CPU placement. Tests for # the startup binding path explicitly reset this guard. worker._cpu_binding_attempted = True @@ -1961,6 +1967,23 @@ def test_npu_ffn_runner_shutdown_is_idempotent(monkeypatch): assert parent_calls == [runner] +def test_npu_cam_eager_runner_rejects_direct_close_with_live_connector(monkeypatch): + _require_npu_runtime() + from afd_plugin.v1.worker.npu import ffn_model_runner + + runner = _new_ffn_runner() + closes = [] + runner._async_cam_startup = SimpleNamespace(started=True, failed=False) + runner.connector = SimpleNamespace( + is_initialized=True, close=lambda: closes.append(True) + ) + monkeypatch.setattr(ffn_model_runner, "stop_afd_npu_profiler", lambda _: None) + + with pytest.raises(RuntimeError, match="receiver must drain"): + runner.shutdown() + assert closes == [] + + def test_npu_ffn_worker_scheduler_execute_model_fails_fast(): worker = _new_ffn_worker() @@ -1991,6 +2014,7 @@ def test_npu_ffn_worker_start_binds_physical_npu_once_before_daemon( worker.local_rank = 3 worker.model_runner = SimpleNamespace( connector=SimpleNamespace(is_initialized=True), + _async_cam_startup=None, ) events: list[tuple[str, object]] = [] monkeypatch.setattr( @@ -2056,6 +2080,7 @@ def test_npu_ffn_worker_cpu_binding_failure_does_not_abort_daemon_start( worker.local_rank = 5 worker.model_runner = SimpleNamespace( connector=SimpleNamespace(is_initialized=True), + _async_cam_startup=None, ) thread_starts: list[bool] = [] monkeypatch.setattr( @@ -2101,6 +2126,7 @@ def test_npu_ffn_worker_loop_error_is_propagated(caplog): worker._ffn_loop_error = None worker.model_runner = SimpleNamespace( connector=SimpleNamespace(is_initialized=True), + _async_cam_startup=None, ) expected_error = RuntimeError("boom") @@ -2132,6 +2158,7 @@ def test_npu_ffn_worker_ignores_receive_error_during_shutdown(caplog): worker._ffn_loop_error = None worker.model_runner = SimpleNamespace( connector=SimpleNamespace(is_initialized=True), + _async_cam_startup=None, ) def stop_while_receiving(): @@ -2169,7 +2196,7 @@ def is_alive(self): connector = SimpleNamespace(close=lambda: calls.append(("close", None))) worker._ffn_thread = _StoppingThread() - worker.model_runner = SimpleNamespace(connector=connector) + worker.model_runner = SimpleNamespace(connector=connector, _async_cam_startup=None) monkeypatch.setattr( ffn_worker.NPUWorker, "shutdown", @@ -2207,6 +2234,7 @@ def is_alive(self): worker._ffn_thread = thread worker.model_runner = SimpleNamespace( connector=SimpleNamespace(close=lambda: calls.append(("close", None))), + _async_cam_startup=None, ) monkeypatch.setattr( ffn_worker.NPUWorker, @@ -2440,10 +2468,46 @@ def test_npu_async_feature_validation_requires_async_config_and_eager(): config = _vllm_config(connector="CAMAsyncAFDConnector", async_dp=True) config.model_config.enforce_eager = False - with pytest.raises(RuntimeError, match="only eager"): + with pytest.raises(RuntimeError, match="requires FFN FULL"): fail_if_unsupported_npu_afd_features(config) +def test_npu_async_feature_validation_allows_only_layered_dsv4_ffn_full( + monkeypatch, +): + _require_npu_runtime() + from afd_plugin.compat.npu import feature_validation + + config = _vllm_config( + role="ffn", + connector="CAMAsyncAFDConnector", + async_dp=True, + compute_gate_on_attention=True, + extra_config={"dynamicQuant": "1"}, + ) + config.model_config.enforce_eager = False + config.model_config.hf_config = SimpleNamespace(model_type="deepseek_v4") + config.use_v2_model_runner = False + monkeypatch.setattr( + feature_validation, "async_cam_layered_gmm_enabled", lambda: True + ) + fail_if_unsupported_npu_afd_features(config) + + for graph_mode in ( + "NONE", + "PIECEWISE", + "FULL_AND_PIECEWISE", + "FULL_DECODE_ONLY", + ): + config.compilation_config.cudagraph_mode.name = graph_mode + with pytest.raises(RuntimeError, match="requires FFN FULL"): + fail_if_unsupported_npu_afd_features(config) + + config.model_config.enforce_eager = True + config.compilation_config.cudagraph_mode.name = "NONE" + fail_if_unsupported_npu_afd_features(config) + + @pytest.mark.parametrize( ("parallel_override", "error"), [ @@ -2695,7 +2759,8 @@ def test_npu_ffn_runner_disables_mask_only_for_cam_mrv1( def fake_native_init(self, vllm_config, device): self.model_config = SimpleNamespace( - hf_config=SimpleNamespace(num_hidden_layers=1) + hf_config=SimpleNamespace(num_hidden_layers=1), + enforce_eager=True, ) module.ascend_context._reserved_mc2_mask = mask @@ -2714,7 +2779,10 @@ def fake_native_init(self, vllm_config, device): module.AFDNPUFFNModelRunner, "parse_config", staticmethod( - lambda config: SimpleNamespace(compute_gate_on_attention=compute_gate) + lambda config: SimpleNamespace( + compute_gate_on_attention=compute_gate, + connector="P2pNcclAFDConnector", + ) ), ) @@ -2792,6 +2860,7 @@ def test_npu_attention_runner_load_model_initializes_connector_after_weights( connector = _LifecycleConnector(events) runner = object.__new__(attention_model_runner.AFDNPUAttentionModelRunner) runner.connector = connector + runner._async_cam_startup = None runner.vllm_config = SimpleNamespace( parallel_config=SimpleNamespace(use_ubatching=use_ubatching), ) @@ -2829,6 +2898,7 @@ def test_npu_attention_runner_afd_ubatching_does_not_install_native_wrapper( connector = _LifecycleConnector(events) runner = object.__new__(attention_model_runner.AFDNPUAttentionModelRunner) runner.connector = connector + runner._async_cam_startup = None runner.vllm_config = SimpleNamespace( parallel_config=SimpleNamespace( use_ubatching=False, @@ -2849,3 +2919,179 @@ def test_npu_attention_runner_afd_ubatching_does_not_install_native_wrapper( runner.load_model() assert events == ["model_load", "connector_init"] + + +def test_npu_async_cam_ffn_graph_replays_one_item_per_moe_layer(monkeypatch): + _require_npu_runtime() + from afd_plugin.v1.worker.npu import ffn_model_runner + + class FakeGraph: + def __init__(self): + self.replays = 0 + + def replay(self): + self.replays += 1 + + graph = FakeGraph() + runner = object.__new__(ffn_model_runner.AFDNPUFFNModelRunner) + runner.connector = SimpleNamespace(control_plane=None, world_rank=8) + runner.prof = None + runner._async_cam_ffn_graph_enabled = True + runner._async_cam_ffn_graph = graph + runner._async_cam_ffn_replays = 0 + monkeypatch.setattr(ffn_model_runner, "step_afd_npu_profiler", lambda _: None) + monkeypatch.setattr(ffn_model_runner, "_ffn_layer_indices", lambda _: [1, 4, 7]) + + runner.execute_connector_driven_step() + + assert graph.replays == 3 + assert runner._async_cam_ffn_replays == 3 + + +def test_npu_async_cam_ffn_graph_capture_is_idempotent_and_failure_is_atomic( + monkeypatch, +): + _require_npu_runtime() + from afd_plugin.v1.worker.npu import ffn_model_runner + + @contextmanager + def fake_graph_context(graph, *, pool): + yield graph + + graph_creations = [] + + def make_graph(): + graph = SimpleNamespace() + graph_creations.append(graph) + return graph + + monkeypatch.setattr(ffn_model_runner.torch.npu, "NPUGraph", make_graph) + monkeypatch.setattr(ffn_model_runner.torch.npu, "graph", fake_graph_context) + runner = object.__new__(ffn_model_runner.AFDNPUFFNModelRunner) + runner.connector = SimpleNamespace(is_initialized=True, world_rank=8) + runner._layered_executor = object() + runner._async_cam_ffn_graph_enabled = True + runner._async_cam_ffn_graph = None + runner._async_cam_ffn_output = None + runner.graph_pool = None + runner._execute_layered_work_item = lambda **_: "captured-output" + + runner.capture_async_cam_ffn_graph() + runner.capture_async_cam_ffn_graph() + assert len(graph_creations) == 1 + assert runner._async_cam_ffn_output == "captured-output" + + runner._async_cam_ffn_graph = None + runner._async_cam_ffn_output = None + runner._execute_layered_work_item = lambda **_: (_ for _ in ()).throw( + RuntimeError("capture failed") + ) + with pytest.raises(RuntimeError, match="capture failed"): + runner.capture_async_cam_ffn_graph() + assert runner._async_cam_ffn_graph is None + assert runner._async_cam_ffn_output is None + + +def test_npu_async_cam_ffn_worker_delegates_one_complete_startup(): + _require_npu_runtime() + + events = [] + worker = _new_ffn_worker() + worker._ffn_thread = None + worker._start_ffn_receiver = lambda: events.append("receiver") + + class FakeStartup: + failed = False + + def start_ffn(self, *, prepare, consume_warmup, capture, start_receiver): + events.append("coordinator") + plan = prepare() + assert plan.use_graph and plan.first_layer_idx == 3 + consume_warmup(2) + capture() + start_receiver() + + worker.model_runner = SimpleNamespace( + _async_cam_startup=FakeStartup(), + prepare_async_cam_ffn_startup=lambda: SimpleNamespace( + use_graph=True, first_layer_idx=3 + ), + warmup_async_cam_ffn_communication=lambda count: events.append( + ("warmup", count) + ), + capture_async_cam_ffn_graph=lambda: events.append("capture"), + ) + worker.start_ffn_server_loop() + assert events == ["coordinator", ("warmup", 2), "capture", "receiver"] + + +def test_npu_async_cam_graph_shutdown_releases_after_thread_stops(): + _require_npu_runtime() + from afd_plugin.connectors.npu.async_cam import CAMAsyncAFDConnector + + events = [] + connector = object.__new__(CAMAsyncAFDConnector) + connector.world_rank = 8 + connector.close = lambda: events.append("close group and anchors") + worker = _new_ffn_worker() + worker._ffn_shutdown_event = threading.Event() + worker._ffn_loop_started_event = threading.Event() + worker._ffn_receiver_drained = True + + class FakeThread: + active = True + + def is_alive(self): + return self.active + + def join(self, *, timeout): + events.append("join") + self.active = False + + worker._ffn_thread = FakeThread() + worker.model_runner = SimpleNamespace( + connector=connector, + _async_cam_ffn_graph_enabled=True, + _async_cam_startup=SimpleNamespace(failed=False, started=True), + release_async_cam_ffn_graph=lambda: events.append("release graph"), + ) + + worker.stop_ffn_server_loop() + + assert events == [ + "join", + "release graph", + "close group and anchors", + ] + + +def test_npu_async_cam_graph_shutdown_retains_anchors_if_join_times_out(): + _require_npu_runtime() + from afd_plugin.connectors.npu.async_cam import CAMAsyncAFDConnector + + events = [] + connector = object.__new__(CAMAsyncAFDConnector) + connector.world_rank = 8 + connector.close = lambda: events.append("close group and anchors") + worker = _new_ffn_worker() + worker._ffn_shutdown_event = threading.Event() + + class BlockedThread: + def is_alive(self): + return True + + def join(self, *, timeout): + events.append("join timeout") + + worker._ffn_thread = BlockedThread() + worker.model_runner = SimpleNamespace( + connector=connector, + _async_cam_ffn_graph_enabled=True, + _async_cam_startup=SimpleNamespace(failed=False, started=True), + release_async_cam_ffn_graph=lambda: events.append("release graph"), + ) + + with pytest.raises(RuntimeError, match="still active"): + worker.stop_ffn_server_loop() + + assert events == ["join timeout"]