Skip to content
18 changes: 15 additions & 3 deletions afd_plugin/compat/npu/feature_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
21 changes: 21 additions & 0 deletions afd_plugin/connectors/factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]]] = {}
Expand Down Expand Up @@ -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,
Expand Down
30 changes: 21 additions & 9 deletions afd_plugin/connectors/npu/async_cam.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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

Expand All @@ -58,6 +59,7 @@
AFDTransferState,
)
from afd_plugin.distributed import (
ProcessGroupRendezvousContext,
create_hccl_process_group_options,
init_afd_process_group,
)
Expand Down Expand Up @@ -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.

Expand All @@ -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,
Expand Down Expand Up @@ -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))
Expand All @@ -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."""
Expand Down
2 changes: 2 additions & 0 deletions afd_plugin/distributed/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
def __getattr__(name: str):
if name in {
"DefaultProcessGroupSwitcher",
"ProcessGroupRendezvousContext",
"create_hccl_process_group_options",
"init_afd_process_group",
}:
Expand All @@ -29,6 +30,7 @@ def __getattr__(name: str):
__all__ = [
"AFDRankMapping",
"DefaultProcessGroupSwitcher",
"ProcessGroupRendezvousContext",
"build_rank_mapping",
"create_hccl_process_group_options",
"init_afd_process_group",
Expand Down
35 changes: 35 additions & 0 deletions afd_plugin/distributed/afd_process_group.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

from __future__ import annotations

from collections.abc import Callable
from datetime import timedelta
from typing import Any

Expand All @@ -12,6 +13,7 @@
from torch.distributed.distributed_c10d import (
PrefixStore,
ProcessGroup,
Store,
_new_process_group_helper,
_update_default_pg,
_world,
Expand All @@ -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."""

Expand Down Expand Up @@ -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.

Expand All @@ -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 = (
Expand Down Expand Up @@ -118,6 +152,7 @@ def init_afd_process_group(

__all__ = [
"DefaultProcessGroupSwitcher",
"ProcessGroupRendezvousContext",
"create_hccl_process_group_options",
"init_afd_process_group",
]
Loading
Loading