diff --git a/.github/workflows/scripts/estimated_times.yaml b/.github/workflows/scripts/estimated_times.yaml index 33aaf8fc9c11..2ed85bec1bfa 100644 --- a/.github/workflows/scripts/estimated_times.yaml +++ b/.github/workflows/scripts/estimated_times.yaml @@ -144,3 +144,4 @@ estimated_times: tests/e2e/pull_request/one_card/test_thinking_budget.py: 80 tests/e2e/pull_request/one_card/test_model_runner_v1_with_device.py: 70 tests/e2e/pull_request/one_card/test_msa_index_score.py: 20 + tests/e2e/pull_request/four_card/fault_tolerance/test_fault_tolerance_e2e.py: 900 diff --git a/tests/e2e/pull_request/four_card/fault_tolerance/__init__.py b/tests/e2e/pull_request/four_card/fault_tolerance/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/tests/e2e/pull_request/four_card/fault_tolerance/test_fault_tolerance_e2e.py b/tests/e2e/pull_request/four_card/fault_tolerance/test_fault_tolerance_e2e.py new file mode 100644 index 000000000000..c9e954d2d7f5 --- /dev/null +++ b/tests/e2e/pull_request/four_card/fault_tolerance/test_fault_tolerance_e2e.py @@ -0,0 +1,615 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""End-to-end tests for the fault-tolerance framework on Ascend NPU. + +Requires 4 NPUs (DP=4). Retry is gated behind ``has_npu_ft_capability()``; +scale-down additionally behind ``has_npu_scale_down_capability()`` (CANN V3+). +""" + +import contextlib +import json +import os +import threading +import time +import uuid +from concurrent.futures import ThreadPoolExecutor +from typing import Any + +import psutil +import pytest +import regex as re +import requests +import torch + +from tests.e2e.conftest import RemoteOpenAIServer + +MODEL_NAME = os.getenv("MODEL_NAME", "vllm-ascend/Qwen3-30B-A3B-W8A8") +DP_SIZE = 4 + +# Fault-detection timeout budget: +# - CPU: Gloo DP allreduce timeout detects the dead peer. +# - NPU: HCSP operator timeout detects the dead peer. +# - Deadline: slowest fallback + margin. +CPU_DISTRIBUTED_TIMEOUT_S = 15 +FT_COMMUNICATION_ABORT_TIMEOUT_S = 10 +FAULT_DETECTION_DEADLINE_S = 45 + +# scale_down re-hosts the dead rank's experts on the survivors and reloads +# the reassigned weights from disk; give the redistribution + dummy-batch +# check more headroom than fault detection (still under +# engine_recovery_timeout_sec=120, the busy-loop's own give-up point). +SCALE_DOWN_DEADLINE_S = 90 + +# Qwen3-30B-A3B has 128 routed experts; on EP=4, 48 redundant experts give +# 44 physical slots per rank, so the 3 surviving ranks still have +# 3 * 44 = 132 >= 128 slots after one rank is removed (the strict minimum +# ``check_redundancy_sufficient`` accepts for a 4 -> 3 shrink is 44). +NUM_REDUNDANT_EXPERTS = 48 + +# Post-recovery accuracy check: after retry recovery, every DP rank must +# answer these factual prompts correctly. Each expected answer is a single +# word, matched case-insensitively on word boundaries anywhere in the +# completion, so harmless phrasing variations do not fail the check. +_ANSWER_CASES = [ + ("The capital of France is", ("Paris",)), + ("The largest planet in our solar system is", ("Jupiter",)), + ("The first month of the year is", ("January",)), +] +_ANSWER_MAX_TOKENS = 16 + + +# --------------------------------------------------------------------------- +# Fault-injection via sitecustomize.py +# --------------------------------------------------------------------------- +# Patches ``dp_utils.sync_cudagraph_and_dp_padding`` to raise on ``rank`` after +# the DP all_reduce at a chosen step. Gated on VLLM_FT_TEST_INJECT_FAULT. +# +# The import hook waits for ``vllm.v1.worker.gpu.dp_utils`` to land in +# sys.modules and for ``sync_cudagraph_and_dp_padding`` to be defined, then +# wraps the function with a step-counting wrapper. The wrapper calls the +# original (so the all_reduce completes) and raises after it returns. +_FAULT_INJECT_SITECUSTOMIZE = """\ +import builtins +import os +import sys + +_SPEC = os.environ.get("VLLM_FT_TEST_INJECT_FAULT") +_MODULE = "vllm.v1.worker.gpu.dp_utils" +_FUNC = "sync_cudagraph_and_dp_padding" + +if _SPEC: + _f = dict(kv.split("=", 1) for kv in _SPEC.split(",")) + _RANK, _STEP = int(_f["rank"]), int(_f["step"]) + _steps = [0] + + def _patch(m): + import inspect + + _orig = getattr(m, _FUNC) + _sig = inspect.signature(_orig) + + def _wrapped(*args, **kwargs): + result = _orig(*args, **kwargs) + bound = _sig.bind(*args, **kwargs) + bound.apply_defaults() + dp_rank = bound.arguments.get("dp_rank") + if dp_rank == _RANK: + _steps[0] += 1 + if _steps[0] == _STEP: + raise RuntimeError( + "FT test fault injection (rank=%d step=%d)" + % (_RANK, _STEP) + ) + return result + + setattr(m, _FUNC, _wrapped) + + _real_import = builtins.__import__ + + def _hook(name, *a, **k): + module = _real_import(name, *a, **k) + m = sys.modules.get(_MODULE) + if ( + m is not None + and hasattr(m, _FUNC) + and not getattr(m, "_ft_patched", False) + ): + m._ft_patched = True + _patch(m) + return module + + builtins.__import__ = _hook +""" + + +def _install_fault_injection(monkeypatch, tmp_path, rank: int, step: int) -> None: + """Arrange for the DP all_reduce sync to raise on ``rank`` at serving ``step``. + + Writes a ``sitecustomize.py`` and prepends its dir to PYTHONPATH so every + vLLM subprocess picks it up; the fault spec is read from the environment. + """ + site_dir = tmp_path / "ft_inject" + site_dir.mkdir() + (site_dir / "sitecustomize.py").write_text(_FAULT_INJECT_SITECUSTOMIZE) + existing = os.environ.get("PYTHONPATH", "") + monkeypatch.setenv( + "PYTHONPATH", + str(site_dir) + (os.pathsep + existing if existing else ""), + ) + monkeypatch.setenv("VLLM_FT_TEST_INJECT_FAULT", f"rank={rank},step={step}") + + +# --------------------------------------------------------------------------- +# Server management +# --------------------------------------------------------------------------- + + +def _ft_server_args(extra_args: list[str] | None = None) -> list[str]: + # Quantized path end to end: --quantization ascend (W8A8); no --dtype, + # the checkpoint's own config decides. MODEL_NAME must point at a W8A8 + # checkpoint (e.g. vllm-ascend/Qwen3-30B-A3B-W8A8). + return [ + "--quantization", + "ascend", + "--max-model-len", + "37364", + "--max-num-seqs", + "128", + "--enable-expert-parallel", + "--enable-fault-tolerance", + "--cpu-distributed-timeout-seconds", + str(CPU_DISTRIBUTED_TIMEOUT_S), + "--fault-tolerance-config", + '{"engine_recovery_timeout_sec": 120}', + "--additional-config", + f'{{"ft_communication_abort_timeout": {FT_COMMUNICATION_ABORT_TIMEOUT_S}}}', + *(extra_args or []), + ] + + +class FTServerManager: + """Manages DP=4 vLLM server instances for fault-tolerance testing. + + Starts one process per DP rank with fixed ports (8000 + rank). + """ + + def __init__( + self, + model_name: str, + dp_size: int, + base_server_args: list[str], + tp_size: int = 1, + ): + self.model_name = model_name + self.dp_size = dp_size + self.tp_size = tp_size + self.base_server_args = base_server_args + self.servers: list[tuple[RemoteOpenAIServer, list[str]]] = [] + self.server_threads: list[threading.Thread] = [] + + def __enter__(self) -> list[tuple[RemoteOpenAIServer, list[str]]]: + for rank in range(self.dp_size): + server_args = self.base_server_args.copy() + server_args.extend( + [ + "--data-parallel-size", + str(self.dp_size), + "--data-parallel-rank", + str(rank), + "--data-parallel-size-local", + "1", + "--tensor-parallel-size", + str(self.tp_size), + "--port", + str(8000 + rank), + "--api-server-count", + "1", + ] + ) + + def start_server(r: int, sargs: list[str]) -> None: + try: + server = RemoteOpenAIServer( + self.model_name, + sargs, + server_host="localhost", + server_port=8000 + r, + auto_port=False, + env_dict={ + "ASCEND_RT_VISIBLE_DEVICES": str(r), + "VLLM_USE_V2_MODEL_RUNNER": "1", + }, + ) + self.servers.append((server, sargs)) + except Exception: + print(f"Failed to start server rank {r}") + raise + + thread = threading.Thread(target=start_server, args=(rank, server_args)) + thread.start() + self.server_threads.append(thread) + + for thread in self.server_threads: + thread.join() + + if len(self.servers) != self.dp_size: + raise RuntimeError(f"Only {len(self.servers)}/{self.dp_size} servers started") + + return self.servers + + def __exit__(self, exc_type, exc_val, exc_tb) -> None: + for server, _ in reversed(self.servers): + with contextlib.suppress(Exception): + server.__exit__(None, None, None) + self.servers.clear() + + +def _ft_manager(extra_args: list[str] | None = None) -> FTServerManager: + return FTServerManager( + MODEL_NAME, + DP_SIZE, + base_server_args=_ft_server_args(extra_args), + tp_size=1, + ) + + +def _server_for_rank(servers: list[tuple[RemoteOpenAIServer, list[str]]], rank: int) -> RemoteOpenAIServer: + """Locate the server for a DP rank.""" + for server, sargs in servers: + if "--data-parallel-rank" in sargs: + idx = sargs.index("--data-parallel-rank") + if int(sargs[idx + 1]) == rank: + return server + raise AssertionError(f"no server found for DP rank {rank}") + + +# --------------------------------------------------------------------------- +# Test primitives +# --------------------------------------------------------------------------- + + +def _complete(client) -> Any: + """Issue one completion request; used to drive the serving loop.""" + return client.completions.create( + model=MODEL_NAME, + prompt="Hello, my name is", + max_tokens=5, + temperature=0.0, + ) + + +def _in_parallel(fn, servers) -> list[Any]: + """Run ``fn(server)`` for all servers concurrently; return in order.""" + with ThreadPoolExecutor(max_workers=len(servers)) as ex: + return list(ex.map(fn, servers)) + + +def _get_ft_status(server: RemoteOpenAIServer) -> dict: + resp = requests.get(server.url_for("v1/fault_tolerance/status"), timeout=10) + resp.raise_for_status() + return resp.json() + + +def _apply_ft( + server: RemoteOpenAIServer, + instruction: str, + params: dict | None = None, + request_id: str | None = None, +) -> dict: + """POST an FT instruction; assert it is accepted (202) and return body.""" + resp = requests.post( + server.url_for("v1/fault_tolerance/apply"), + json={ + "instruction": instruction, + "params": params or {}, + "request_id": request_id or str(uuid.uuid4()), + }, + timeout=10, + ) + assert resp.status_code == 202, resp.text + return resp.json() + + +def _assert_serving_and_healthy( + servers: tuple[RemoteOpenAIServer, ...], +) -> None: + """Wait until every engine is healthy, then serve one request per server.""" + healthy = _wait_for_engines(list(servers), match_key="status", match_values={"healthy"}) + assert all(healthy), healthy + _in_parallel(lambda s: _complete(s.get_client()), servers) + + +def _assert_correct_answers(servers: tuple[RemoteOpenAIServer, ...]) -> None: + """Send factual prompts to every DP rank and assert the answer appears. + + Must only be called after FT recovery (retry or scale_down) has fully + completed (see ``_assert_serving_and_healthy``), so the requests are not + sent into a still-faulting cluster. + """ + for server in servers: + client = server.get_client() + for prompt, answers in _ANSWER_CASES: + resp = client.completions.create( + model=MODEL_NAME, + prompt=prompt, + max_tokens=_ANSWER_MAX_TOKENS, + temperature=0.0, + ) + completion = resp.choices[0].text + matched = any(re.search(rf"\b{re.escape(answer)}\b", completion, re.IGNORECASE) for answer in answers) + print(f"[accuracy] rank {server.port} prompt {prompt!r}: {completion!r}") + assert matched, ( + f"[post-recovery] rank {server.port} answered {prompt!r} incorrectly: " + f"expected one of {answers!r}, got {completion!r}" + ) + + +def _kill_worker_process(server: RemoteOpenAIServer) -> None: + """SIGKILL only the worker proc, leaving EngineCore and API server alive.""" + workers = [p for p in psutil.Process(server.proc.pid).children(recursive=True) if "Worker" in " ".join(p.cmdline())] + assert len(workers) == 1, f"expected 1 worker proc, found: {workers}" + workers[0].kill() + + +def _wait_for_engines( + servers: list[RemoteOpenAIServer], + match_key: str, + match_values: set[str], + deadline_s: int = FAULT_DETECTION_DEADLINE_S, +) -> list[dict[str, Any] | None]: + """Poll ``/v1/fault_tolerance/status`` until each server's engine status matches. + + A server matches when its engine-status dict has ``match_key`` equal to + one of ``match_values``. Returns one engine-status dict per server. + Servers still unmatched after ``deadline_s`` get None. + """ + results: dict[int, dict[str, Any]] = {} + pending = dict(enumerate(servers)) + start = time.time() + while pending and time.time() - start < deadline_s: + for i, server in list(pending.items()): + with contextlib.suppress(Exception): + for engine_status in _get_ft_status(server)["engines"]: + if engine_status.get(match_key) in match_values: + results[i] = engine_status + del pending[i] + break + if pending: + time.sleep(1.0) + return [results.get(i) for i in range(len(servers))] + + +@contextlib.contextmanager +def _driving(*servers: RemoteOpenAIServer): + """Pump completions at each server in the background for the block's duration. + + Keeps every engine stepping into its failed component so a fault surfaces. + Errors are expected once faulted and are ignored. + """ + stop = threading.Event() + + def _drive(server): + client = server.get_client() + while not stop.is_set(): + with contextlib.suppress(Exception): + _complete(client) + time.sleep(0.2) + + threads = [threading.Thread(target=_drive, args=(s,), daemon=True) for s in servers] + for t in threads: + t.start() + try: + yield + finally: + stop.set() + for t in threads: + t.join(timeout=2) + + +def _wait_for_ft_apply_outcome(server: RemoteOpenAIServer, request_id: str, deadline_s: int) -> str | None: + """Wait until ``/v1/fault_tolerance/status`` records the FT apply outcome.""" + engine_status = _wait_for_engines( + [server], + match_key="last_ft_request_id", + match_values={request_id}, + deadline_s=deadline_s, + )[0] + return engine_status.get("ft_error") if engine_status else None + + +# --------------------------------------------------------------------------- +# Feature guard +# --------------------------------------------------------------------------- + + +def has_npu_ft_capability() -> bool: + """Require at least 4 visible NPUs for DP=4 fault-tolerance tests.""" + if not torch.npu.is_available(): + return False + try: + return torch.npu.device_count() >= DP_SIZE + except Exception: + return False + + +def has_npu_scale_down_capability() -> bool: + """scale_down additionally requires the V3 dispatch op (CANN V3+). + + Mirrors the engine-side precondition in ``WorkerSentinel + ._validate_scale_down_preconditions`` so the test skips cleanly on + older CANN/torch_npu stacks instead of timing out mid-recovery. + """ + if not has_npu_ft_capability(): + return False + try: + import torch_npu + + return hasattr(torch_npu, "npu_moe_distribute_dispatch_v2") + except Exception: + return False + + +# --------------------------------------------------------------------------- +# Tests +# --------------------------------------------------------------------------- + + +@pytest.mark.skipif( + not has_npu_ft_capability(), + reason="Requires at least 4 NPUs for DP=4 fault-tolerance testing", +) +def test_injected_fault_retry_recovers_all_ranks(monkeypatch, tmp_path): + """An exception injected after the DP allreduce drives full + retry recovery on all 4 DP ranks. + + Inject a fault at a chosen step on rank 3: + + - Rank 3 raises after allreduce and goes UNHEALTHY. + - Ranks 0, 1, 2 detect the now-absent peer via the Gloo DP allreduce + timeout and also go UNHEALTHY. + + All 4 being UNHEALTHY is the precondition for ``retry``. The fault + is patched via a generated ``sitecustomize.py``. + + After retry recovery completes, every rank must answer factual prompts + correctly (greedy decoding; case-insensitive whole-word match, not an + exact-match golden, so the check is independent of environment numerics). + """ + fault_step = int(os.getenv("FT_FAULT_STEP", "100")) + _install_fault_injection(monkeypatch, tmp_path, rank=3, step=fault_step) + + with _ft_manager() as servers: + assert len(servers) == DP_SIZE + all_ranks = tuple(_server_for_rank(servers, r) for r in range(DP_SIZE)) + + # 1. All engines healthy and serving. + _assert_serving_and_healthy(all_ranks) + + # 2. Drive all ranks so rank 3 accumulates steps and trips the + # injected fault; ranks 0,1,2 then time out on DP allreduce. + with _driving(*all_ranks): + faulted = _wait_for_engines( + list(all_ranks), + match_key="status", + match_values={"unhealthy"}, + ) + + for rank, engine_status in enumerate(faulted): + assert engine_status is not None, ( + f"rank {rank} did not report UNHEALTHY within {FAULT_DETECTION_DEADLINE_S}s -- it likely hung" + ) + + # The rank that raised carries the fault info from its own exception. + assert faulted[3] is not None + assert faulted[3].get("fault_info"), faulted[3] + + # 3. retry all engines. + round_id = str(uuid.uuid4()) + for server in all_ranks: + _apply_ft(server, "retry", request_id=round_id) + + # 4. Recovery completes: all engines return to healthy and serve again. + _assert_serving_and_healthy(all_ranks) + + # 5. Only now (after retry fully completed) send the accuracy prompts + # to every DP rank directly. The injected fault is a one-shot step + # guard, so these requests do not re-trigger it. + _assert_correct_answers(all_ranks) + + +@pytest.mark.skipif( + not has_npu_scale_down_capability(), + reason=("Requires at least 4 NPUs and npu_moe_distribute_dispatch_v2 (CANN V3+) for DP=4 scale-down testing"), +) +def test_scale_down_removes_dead_rank_and_recovers(): + """scale_down removes the dead DP rank; survivors keep serving. + + SIGKILL rank 3's worker: survivors go UNHEALTHY, the victim goes DEAD + and rejects ``retry`` (recovery requires UNHEALTHY). ``scale_down`` + with ``removed_dp_ranks=[3]`` masks the dead rank, redistributes its + EPLB experts onto the survivors and reloads the reassigned weights. + Rank 3 (not rank 0, the DP store master) is removed so no + ``dp_master_ip`` / ``dp_store_port`` params are needed. After recovery + every survivor must answer factual prompts correctly. + """ + victim_rank = DP_SIZE - 1 + extra_args = [ + "--compilation-config", + json.dumps( + { + "cudagraph_mode": "FULL_AND_PIECEWISE", + "cudagraph_capture_sizes": [4, 8, 12, 16, 20, 24, 28, 32], + } + ), + "--enable-eplb", + "--eplb-config.num_redundant_experts", + str(NUM_REDUNDANT_EXPERTS), + ] + with _ft_manager(extra_args) as servers: + assert len(servers) == DP_SIZE + servers_by_rank = {r: _server_for_rank(servers, r) for r in range(DP_SIZE)} + victim = servers_by_rank[victim_rank] + survivor_ranks = [r for r in range(DP_SIZE) if r != victim_rank] + survivors = tuple(servers_by_rank[r] for r in survivor_ranks) + all_ranks = tuple(servers_by_rank[r] for r in range(DP_SIZE)) + + # 1. All engines healthy and serving. + _assert_serving_and_healthy(all_ranks) + + # 2. Kill the victim's worker; survivors detect the peer fault. + _kill_worker_process(victim) + with _driving(*all_ranks): + faulted_results = _wait_for_engines( + list(all_ranks), + match_key="status", + match_values={"dead", "unhealthy"}, + ) + + for rank, engine_status in enumerate(faulted_results): + assert engine_status is not None, ( + f"rank {rank} did not report the fault within {FAULT_DETECTION_DEADLINE_S}s -- it likely hung" + ) + if rank == victim_rank: + # The victim is DEAD: its own worker is gone. + assert engine_status["status"] == "dead", engine_status + else: + # Survivors report the peer fault as UNHEALTHY with fault info + # (UNHEALTHY is the precondition for accepting scale_down). + assert engine_status["status"] == "unhealthy", engine_status + assert engine_status.get("fault_info"), engine_status + + # 3. The DEAD victim rejects retry: recovery requires UNHEALTHY. + # retry is accepted at the HTTP layer (202 = background dispatch)... + request_id = _apply_ft(victim, "retry")["request_id"] + # 4. ...but the rejection is recorded in /v1/fault_tolerance/status. + ft_error = _wait_for_ft_apply_outcome(victim, request_id, FAULT_DETECTION_DEADLINE_S) + assert ft_error is not None, "rejection was never recorded in /v1/fault_tolerance/status" + assert "status is DEAD" in ft_error, ft_error + + # 5. scale_down to every survivor: remove the dead rank. + round_id = str(uuid.uuid4()) + for server in survivors: + _apply_ft( + server, + "scale_down", + {"removed_dp_ranks": [victim_rank]}, + request_id=round_id, + ) + + # 6. Recovery completes: survivors return to healthy and serve again. + recovered = _wait_for_engines( + list(survivors), + match_key="status", + match_values={"healthy"}, + deadline_s=SCALE_DOWN_DEADLINE_S, + ) + for rank, engine_status in zip(survivor_ranks, recovered): + assert engine_status is not None, ( + f"survivor {rank} did not recover within {SCALE_DOWN_DEADLINE_S}s -- " + "expert redistribution or weight reload likely failed" + ) + _in_parallel(lambda s: _complete(s.get_client()), survivors) + + # 7. Factual answers on the survivors: the re-hosted experts must + # produce correct completions on the shrunken DP group. + _assert_correct_answers(survivors) diff --git a/tests/ut/worker/test_worker_v1.py b/tests/ut/worker/test_worker_v1.py index 96651acc1f85..60ec17d1f9aa 100644 --- a/tests/ut/worker/test_worker_v1.py +++ b/tests/ut/worker/test_worker_v1.py @@ -676,6 +676,7 @@ def test_init_device( worker.parallel_config.data_parallel_size = 1 worker.parallel_config.assigned_physical_gpu_ids = None worker.parallel_config.distributed_executor_backend = "ray" + worker.parallel_config.enable_fault_tolerance = False worker.vllm_config = MagicMock() worker.vllm_config.kv_transfer_config = None worker.cache_config = MagicMock() diff --git a/vllm_ascend/ascend_config.py b/vllm_ascend/ascend_config.py index 827935ef3245..cf293d9c1b49 100644 --- a/vllm_ascend/ascend_config.py +++ b/vllm_ascend/ascend_config.py @@ -498,6 +498,10 @@ class AscendConfig: msmonitor_use_daemon: bool = False enable_transpose_kv_cache_by_block: bool = True weight_nz_mode: int = 1 + # ---- fault-tolerance: comm op abort timeout (s); 0 = disable ---- + # Drives HCCL_EVENT_TIMEOUT / HCCL_EXEC_TIMEOUT (= timeout - 1) and + # set_op_timeout_ms(timeout * 1000). Validated to be 0 or >= 2. + ft_communication_abort_timeout: int = 0 # ---- sub-configs (no vllm_config dep): pydantic dict→dataclass coercion ---- ascend_compilation_config: AscendCompilationConfig = dataclasses.field(default_factory=AscendCompilationConfig) @@ -821,6 +825,17 @@ def derive_and_validate(self, vllm_config: VllmConfig) -> AscendConfig: # sparse KV offload vs sparse SFA C8 main cache mutex self._validate_sparse_c8_kv_offload_compatibility() + + # ft_communication_abort_timeout only takes effect when fault tolerance + # is enabled, so validate it only then: a stray value in additional_config + # must not fail startup for non-FT runs. + if vc.parallel_config.enable_fault_tolerance and ( + self.ft_communication_abort_timeout != 0 and self.ft_communication_abort_timeout < 2 + ): + raise ValueError( + f"ft_communication_abort_timeout must be 0 (disabled) or an integer of at least " + f"2 seconds, got {self.ft_communication_abort_timeout}" + ) return self def _validate_mc2_comm_alg(self, vllm_config: VllmConfig) -> None: diff --git a/vllm_ascend/distributed/device_communicators/npu_communicator.py b/vllm_ascend/distributed/device_communicators/npu_communicator.py index cd38345ed229..f45ae1d8518f 100644 --- a/vllm_ascend/distributed/device_communicators/npu_communicator.py +++ b/vllm_ascend/distributed/device_communicators/npu_communicator.py @@ -21,20 +21,83 @@ class _NpuAll2AllManager: - """No-op all2all_manager for NPU. Used by vLLM main's fault-tolerance - check (data_parallel_size > 1 and is_moe); NPU does not register a real - one because it uses mc2 / all_gather for MoE communication. + """All2All-manager adapter for MC2 fault tolerance. + + Owns the dead-rank mask, encoded into the ``elastic_info`` tensor consumed + by the MC2 dispatch/combine operators. The public interface mirrors the + upstream All2AllManagerBase mask API. """ - @property - def support_fault_tolerance(self) -> bool: - return False + # MC2 kernels do not detect faults themselves; the mask is written + # host-side by FT recovery, so a per-step query can never observe one. + support_fault_tolerance = False - def query_fault(self) -> torch.Tensor: - return torch.zeros(1, dtype=torch.bool, device="cpu") + def __init__(self, ep_world_size: int, device: torch.device | None = None) -> None: + self._ep_world_size = ep_world_size + self._device = device + self._dead: set[int] = set() + self._num_local_experts: int = 0 + + # elastic_info layout: [is_scaling_down, dense ep world size, + # shared_expert_rank_num, num_physical_experts] + table1(orig->dense) + # + table2(dense->orig). num_physical_experts is derived from the + # dead set and num_local_experts on every rebuild. + size = 4 + 2 * ep_world_size + self._elastic_info_host = torch.zeros(size, dtype=torch.int32) + if device is None: + device = torch.device("npu", torch.npu.current_device()) + self._elastic_info = torch.zeros(size, dtype=torch.int32, device=device) + + def update_mask(self, rank: int, masked: bool = True) -> None: + """Mark an EP rank dead/alive and rebuild elastic_info in place.""" + if masked: + self._dead.add(rank) + else: + self._dead.discard(rank) + self._rebuild_elastic_info() def query_active_mask(self) -> torch.Tensor: - return torch.zeros(1, dtype=torch.bool, device="cpu") + """Per-EP-rank mask (1=dead, 0=live) as a CPU tensor, matching the + upstream mask-buffer convention. + + Built on CPU on purpose: this is called while a fault is being + probed, when the NPU may be hung — any device op would fail. + """ + mask = torch.zeros(self._ep_world_size, dtype=torch.int32) + for rank in self._dead: + mask[rank] = 1 + return mask + + def query_fault(self) -> torch.Tensor: + # MC2 has no in-kernel fault detection; faults surface as aborted ops. + return torch.tensor(False) + + def clean_buffers(self) -> None: + """No-op, kept for the upstream retry flow which calls it unconditionally.""" + + def get_elastic_info(self) -> torch.Tensor: + """The device elastic_info tensor for the next MC2 dispatch/combine.""" + return self._elastic_info + + def set_num_local_physical_experts(self, num_local_experts: int) -> None: + """Record the physical expert slots per EP rank.""" + self._num_local_experts = num_local_experts + + def _rebuild_elastic_info(self) -> None: + """Rebuild elastic_info from the dead set into the existing device + tensor (never reallocates, so captured graphs stay valid).""" + + world_size = self._ep_world_size + alive = sorted(set(range(world_size)) - self._dead) + num_physical_experts = len(alive) * self._num_local_experts + table1 = torch.full((world_size,), -1, dtype=torch.int32) + table1[alive] = torch.arange(len(alive), dtype=torch.int32) + table2 = torch.full((world_size,), -1, dtype=torch.int32) + table2[: len(alive)] = torch.tensor(alive, dtype=torch.int32) + self._elastic_info_host.copy_( + torch.cat([torch.tensor([1, len(alive), 0, num_physical_experts], dtype=torch.int32), table1, table2]) + ) + self._elastic_info.copy_(self._elastic_info_host, non_blocking=True) class NPUCommunicator(DeviceCommunicatorBase): @@ -59,4 +122,6 @@ def __init__( # Keep the shared coordinator protocol available without enabling the # FlashInfer PCIe IPC backend on NPU. self.fi_pcie_ipc_ar_comm = None - self.all2all_manager = _NpuAll2AllManager() + # Only the EP group's instance is ever looked up (via the upstream + # get_ep_all2all_manager()); the rest stay dormant. + self.all2all_manager = _NpuAll2AllManager(dist.get_world_size(cpu_group), device) diff --git a/vllm_ascend/models/deepseek_v4/model.py b/vllm_ascend/models/deepseek_v4/model.py index 9e51dc85eac6..3f527b3ee945 100644 --- a/vllm_ascend/models/deepseek_v4/model.py +++ b/vllm_ascend/models/deepseek_v4/model.py @@ -28,6 +28,7 @@ from collections.abc import Callable, Iterable from itertools import islice +import regex as re import torch import torch.nn.functional as F import vllm.envs as envs @@ -65,6 +66,7 @@ ) from vllm.model_executor.models.utils import ( PPMissingLayer, + WeightsMapper, is_pp_missing_parameter, make_layers, maybe_prefix, @@ -1045,6 +1047,25 @@ class AscendDeepseekV4ForCausalLM(nn.Module, SupportsPP, DeepseekV2MixtureOfExpe } model_cls = DeepseekV4Model + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_regex={re.compile(r"^(?!model)"): "model."}, + orig_to_new_substr={ + ".w1.": ".gate_proj.", + ".w2.": ".down_proj.", + ".w3.": ".up_proj.", + "embed.": "embed_tokens.", + ".attn.": ".self_attn.", + ".ffn.": ".mlp.", + ".ffn_norm.": ".post_attention_layernorm.", + ".attn_norm.": ".input_layernorm.", + }, + orig_to_new_prefix={ + "model.head.": "lm_head.", + "model.lm_head.": "lm_head.", + }, + orig_to_new_suffix={".scale": ".weight_scale"}, + ) + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): super().__init__() config = vllm_config.model_config.hf_config @@ -1166,33 +1187,10 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: if spec_layer is not None: continue # skip spec decode layers for main model - # TODO: - if not name.startswith("model"): - name = f"model.{name}" - - if ".w1." in name: - name = name.replace(".w1.", ".gate_proj.") - if ".w2." in name: - name = name.replace(".w2.", ".down_proj.") - if ".w3." in name: - name = name.replace(".w3.", ".up_proj.") - - if "model.head." in name and "model.lm_head." not in name: - name = name.replace("model.head.", "lm_head.") - if "model.lm_head." in name: - name = name.replace("model.lm_head.", "lm_head.") - if "embed." in name and "embed_token." not in name: - name = name.replace("embed.", "embed_tokens.") - if "attn" in name and "self_attn" not in name: - name = name.replace(".attn.", ".self_attn.") - if ".ffn." in name: - name = name.replace(".ffn.", ".mlp.") - if ".ffn_norm." in name: - name = name.replace(".ffn_norm.", ".post_attention_layernorm.") - if ".attn_norm." in name: - name = name.replace(".attn_norm.", ".input_layernorm.") - if name.endswith(".scale"): - name = name.replace(".scale", ".weight_scale") + mapped_name = self.hf_to_vllm_mapper._map_name(name) + if mapped_name is None: + continue # mapper-declared drop (none today; parity with AutoWeightsLoader) + name = mapped_name if "rotary_emb.inv_freq" in name: continue diff --git a/vllm_ascend/ops/fused_moe/token_dispatcher.py b/vllm_ascend/ops/fused_moe/token_dispatcher.py index cb9e65787ea5..af87ad77c851 100644 --- a/vllm_ascend/ops/fused_moe/token_dispatcher.py +++ b/vllm_ascend/ops/fused_moe/token_dispatcher.py @@ -27,6 +27,7 @@ import torch_npu from vllm.config import get_current_vllm_config from vllm.distributed.parallel_state import get_ep_group +from vllm.model_executor.layers.fused_moe.all2all_utils import get_ep_all2all_manager from vllm_ascend.ascend_config import get_ascend_config from vllm_ascend.ascend_forward_context import get_mc2_tokens_capacity @@ -143,6 +144,11 @@ def __init__(self, **kwargs): # use the real global_bs and do NOT pass mc2_mask. self.global_bs = _max_global_bs if should_skip_allreduce_across_dp_group(vllm_config) else 0 + # Fault tolerance: the dead-rank mask is passed to the MC2 operators + # as an explicit elastic_info tensor on every call; dead ranks are + # excluded purely via elastic_info. + self._ft_enabled = vllm_config.parallel_config.enable_fault_tolerance + def refresh_hccl_group(self) -> None: """Refresh MC2 communicator metadata after HCCL groups are recreated.""" device_group = get_mc2_group().device_group @@ -182,6 +188,8 @@ def get_dispatch_mc2_kwargs( "global_bs": self.global_bs, "expert_token_nums_type": expert_token_nums_type, } + if self._ft_enabled: + kwargs_mc2["elastic_info"] = get_ep_all2all_manager().get_elastic_info() if self.global_bs == 0: kwargs_mc2["x_active_mask"] = token_dispatch_input.routing.mc2_mask @@ -294,6 +302,10 @@ def get_combine_mc_kwargs(self, hidden_states: torch.Tensor, combine_metadata: M "moe_expert_num": self.moe_expert_num, "global_bs": self.global_bs, } + if self._ft_enabled: + # The combine's alltoallv spans the same EP group and must + # exclude dead ranks with the identical elastic_info tensor. + kwargs_mc2["elastic_info"] = get_ep_all2all_manager().get_elastic_info() if self.global_bs == 0: kwargs_mc2["x_active_mask"] = combine_metadata.mc2_mask diff --git a/vllm_ascend/patch/worker/patch_distributed.py b/vllm_ascend/patch/worker/patch_distributed.py index f5dd4d40378c..0ca5ad1f3e8f 100644 --- a/vllm_ascend/patch/worker/patch_distributed.py +++ b/vllm_ascend/patch/worker/patch_distributed.py @@ -130,6 +130,7 @@ def __init__( # type: ignore[misc] self.use_cpu_custom_send_recv = False self.group_name = group_name self.group_ranks = group_ranks + self.dead_dp_ranks: set[int] = set() try: self._init_device_groups(create_cpu_group=True) diff --git a/vllm_ascend/patch/worker/patch_v2/patch_async_output.py b/vllm_ascend/patch/worker/patch_v2/patch_async_output.py new file mode 100644 index 000000000000..251e7421ed32 --- /dev/null +++ b/vllm_ascend/patch/worker/patch_v2/patch_async_output.py @@ -0,0 +1,20 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from vllm.v1.outputs import ModelRunnerOutput +from vllm.v1.worker.gpu.async_utils import AsyncOutput + +from vllm_ascend.worker.sentinel.npu_worker_sentinel import fault_barrier_wrapper + + +class AscendAsyncOutput(AsyncOutput): + """AsyncOutput whose ``get_output`` runs behind the FT fault barrier. + + HCCL faults (e.g. after a peer faults) escaping ``get_output`` can make + the destructor/teardown that follows abort at the C level and kill the + process; with FT the worker must survive to recover, so the barrier + quarantines it (stop_device) before any teardown executes. + """ + + @fault_barrier_wrapper + def get_output(self) -> ModelRunnerOutput: + return super().get_output() diff --git a/vllm_ascend/worker/sentinel/__init__.py b/vllm_ascend/worker/sentinel/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/vllm_ascend/worker/sentinel/eplb_redistribute.py b/vllm_ascend/worker/sentinel/eplb_redistribute.py new file mode 100644 index 000000000000..900f6d68d681 --- /dev/null +++ b/vllm_ascend/worker/sentinel/eplb_redistribute.py @@ -0,0 +1,497 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM Ascend project +"""NPU expert redistribution and weight reload for fault-tolerance scale-down. + +The reload deliberately does not degrade: a checkpoint, weight layout or quant +scheme it cannot handle raises. Every fallback that used to sit here read the +whole checkpoint, or wrote tensors that did not match the runtime layout, on +paths that could not succeed anyway. +""" + +import glob +import os +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from functools import partial + +import torch +import torch_npu +from safetensors import safe_open +from transformers.utils import SAFE_WEIGHTS_INDEX_NAME +from vllm.config import VllmConfig +from vllm.distributed import get_ep_group +from vllm.logger import logger +from vllm.model_executor.model_loader.weight_utils import ( + filter_duplicate_safetensors_files, +) + +from vllm_ascend.ops.fused_moe.routed_experts import AscendUnquantizedFusedMoEMethod +from vllm_ascend.quantization.methods.w8a8.w8a8_dynamic import ( + AscendW8A8DynamicFusedMoEMethod, + scale_from_float_to_int64, +) +from vllm_ascend.utils import ACL_FORMAT_FRACTAL_NZ, maybe_trans_nz + +__all__ = [ + "build_orig_to_dense_rank_table", + "densify_routing_table_physical_ids", + "reload_experts_from_disk", +] + +# Checkpoint name suffixes (relative to "..") of the +# tensors a reload may consume. weight_offset is loaded at startup but never +# consumed by the MoE apply path, so it is not reloaded here. +_W13_WEIGHT_SUFFIXES = ("gate_up_proj.weight", "gate_proj.weight", "up_proj.weight") +_W2_WEIGHT_SUFFIX = "down_proj.weight" +_W13_SCALE_SUFFIXES = ("gate_up_proj.weight_scale", "gate_proj.weight_scale", "up_proj.weight_scale") +_W2_SCALE_SUFFIX = "down_proj.weight_scale" + +# Quant schemes the reload supports (gated in _reload_local_slots). +_SUPPORTED_QUANT_TYPES = { + AscendUnquantizedFusedMoEMethod, + AscendW8A8DynamicFusedMoEMethod, +} + +# (layer, expert) pairs per chunk: a larger chunk amortises the device rounds +# further but grows the stack's transient device memory. +_RELOAD_CHUNK_SIZE = 32 + +# Threads assembling one chunk; bounded because worker processes share host +# memory bandwidth. +_GATHER_MAX_WORKERS = 8 + + +def _get_ckpt_name_mapper(model: torch.nn.Module) -> Callable[[str], str]: + """Map raw checkpoint weight names into the model's runtime namespace.""" + mapper = getattr(model, "hf_to_vllm_mapper", None) + if mapper is None: + # No declared mapping means the checkpoint stores weights under + # runtime names already (e.g. Qwen3-MoE), so identity is correct. + return lambda name: name + return lambda name: mapper._map_name(name) or name + + +def build_orig_to_dense_rank_table(ep_world_size: int, dead_ranks: set[int]) -> torch.Tensor: + """Build the orig-rank -> densified-rank mapping table. + + Returns a ``[ep_world_size]`` int32 tensor where ``table[orig_rank]`` is + the densified rank (-1 for dead ranks), i.e. the kernel's table1 view of + the elastic_info layout. Densified ranks are assigned in ascending + original-rank order over the survivors. + """ + table = torch.full((ep_world_size,), -1, dtype=torch.int32) + alive = sorted(set(range(ep_world_size)) - dead_ranks) + for dense_rank, orig_rank in enumerate(alive): + table[orig_rank] = dense_rank + return table + + +def densify_routing_table_physical_ids( + routing_table: torch.Tensor, + orig_to_dense_rank: torch.Tensor, + num_local_experts: int, +) -> None: + """Renumber a routing table's physical ids into the densified id space. + + In scale-down mode the MC2 dispatch kernel computes a token's destination + as ``table2[expert_id // num_local]`` and the combine kernel drops any id + >= the shrunk physical expert count, so the ids produced by the EPLB + mapping must be dense-rank-major: ``dense_rank * num_local + slot``. + Keeping original ids only works when the dead ranks happen to be a suffix + (then table2 is the identity on the alive prefix); a dead rank in the + middle misroutes tokens or crashes the kernel on a -1 rank lookup. + + The update is an in-place ``copy_`` with an unchanged shape, so captured + graphs keep pointing at valid storage. + + Args: + routing_table: ``expert_replica_routing_table`` of one MoE layer + (device, int32), holding original global physical ids. + orig_to_dense_rank: ``[ep_world_size]`` original EP rank -> densified + rank (-1 for dead ranks), i.e. elastic_info's table1. + num_local_experts: physical slots per EP rank (unchanged by + scale-down). + """ + ids = routing_table.to(torch.int64) + if bool((ids < 0).any()): + raise RuntimeError( + "[FT] expert replica routing table references empty slots after " + "redistribution; every logical expert must have a live replica." + ) + orig_rank = torch.div(ids, num_local_experts, rounding_mode="floor") + dense_rank = orig_to_dense_rank.to(device=ids.device, dtype=torch.int64)[orig_rank] + if bool((dense_rank < 0).any()): + raise RuntimeError( + "[FT] expert replica routing table references dead EP ranks after " + "redistribution; the placement did not vacate the dead ranks." + ) + dense_ids = dense_rank * num_local_experts + ids % num_local_experts + routing_table.copy_(dense_ids.to(routing_table.dtype)) + + +def _tp_shard_info(layer) -> tuple[int, int]: + """(tp_rank, tp_size) of the MoE weights of a routed-experts module.""" + parallel_config = layer.moe_config.moe_parallel_config + return parallel_config.tp_rank, parallel_config.tp_size + + +def _shard_row(t: torch.Tensor, tp_rank: int, tp_size: int) -> torch.Tensor: + """Take this TP rank's row shard (dim 0) of a full checkpoint tensor.""" + if tp_size == 1: + return t + shard = t.shape[0] // tp_size + return t.narrow(0, tp_rank * shard, shard) + + +def _shard_col(t: torch.Tensor, tp_rank: int, tp_size: int) -> torch.Tensor: + """Take this TP rank's column shard (dim 1) of a full checkpoint tensor.""" + if tp_size == 1: + return t + shard = t.shape[1] // tp_size + return t.narrow(1, tp_rank * shard, shard) + + +def _gather_w13( + tensors: dict[str, torch.Tensor], + tp_rank: int, + tp_size: int, +) -> torch.Tensor: + """Assemble one expert's w13 ([2I_local, H]) from checkpoint tensors.""" + fused = tensors.get("gate_up_proj.weight") + if fused is not None: + half = fused.shape[0] // 2 + gate, up = fused[:half], fused[half:] + else: + gate, up = tensors["gate_proj.weight"], tensors["up_proj.weight"] + return torch.cat([_shard_row(gate, tp_rank, tp_size), _shard_row(up, tp_rank, tp_size)], dim=0) + + +def _gather_w13_scale( + tensors: dict[str, torch.Tensor], + tp_rank: int, + tp_size: int, +) -> torch.Tensor: + """1-D w13 scale ([2I_local]) from checkpoint tensors.""" + fused = tensors.get("gate_up_proj.weight_scale") + if fused is not None: + half = fused.shape[0] // 2 + gate, up = fused[:half], fused[half:] + else: + gate = tensors["gate_proj.weight_scale"] + up = tensors["up_proj.weight_scale"] + return torch.cat([_shard_row(gate, tp_rank, tp_size), _shard_row(up, tp_rank, tp_size)], dim=0).view(-1) + + +def _assemble_expert( + pair: tuple[int, int], + routed_layers: list, + buckets: dict[tuple[int, int], dict[str, torch.Tensor]], + w8a8: bool, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None, torch.Tensor | None]: + """Assemble one expert's (w13, w2, w13_scale, w2_scale) on the CPU; thread-safe, + as it does not modify shared state. Both scales are None for the unquantized + scheme; the w2 scale is a view into its bucket, so callers must not write it. + """ + layer_idx, logical_id = pair + routed = routed_layers[layer_idx] + if getattr(routed, "w13_bias", None) is not None or getattr(routed, "w2_bias", None) is not None: + raise NotImplementedError("[FT] scale_down weight reload does not support MoE expert bias yet.") + tp_rank, tp_size = _tp_shard_info(routed) + tensors = buckets[(layer_idx, logical_id)] + w13 = _gather_w13(tensors, tp_rank, tp_size).transpose(0, 1).contiguous() + w2 = _shard_col(tensors[_W2_WEIGHT_SUFFIX], tp_rank, tp_size).transpose(0, 1).contiguous() + if not w8a8: + return w13, w2, None, None + return w13, w2, _gather_w13_scale(tensors, tp_rank, tp_size), tensors[_W2_SCALE_SUFFIX].view(-1) + + +def _reload_local_slots( + routed_layers: list, + local_slots: dict[tuple[int, int], int], + buckets: dict[tuple[int, int], dict[str, torch.Tensor]], +) -> int: + """Reload all local (layer, expert) pairs in uniformly-laid-out chunks. + + Every pair shares one quant scheme, weight shape and dtype, so a chunk + stacks into one tensor, e.g. DeepSeek-V4 43-layers, 8 experts to reload, + 344 pairs become ceil(344 / 32) = 11 chunks instead of 344 per-expert reloads. + """ + pairs: list[tuple[int, int]] = [] + layouts: set[tuple] = set() + for (layer_idx, logical_id), slot in sorted(local_slots.items()): + routed = routed_layers[layer_idx] + # Quantized layers wrap the scheme in AscendFusedMoEMethod; the + # scheme class lives in its .quant_method attribute. + quant_method = getattr(routed.quant_method, "quant_method", routed.quant_method) + if type(quant_method) not in _SUPPORTED_QUANT_TYPES: + raise NotImplementedError( + f"[FT] scale_down weight reload is not implemented for quant method {type(quant_method).__name__}." + ) + w13_weight_list = getattr(routed, "w13_weight_list", None) + if w13_weight_list is not None: + w13_shape, w13_dtype = w13_weight_list[slot].shape, w13_weight_list[slot].dtype + w2_shape, w2_dtype = routed.w2_weight_list[slot].shape, routed.w2_weight_list[slot].dtype + else: + w13_shape, w13_dtype = routed.w13_weight.shape[1:], routed.w13_weight.dtype + w2_shape, w2_dtype = routed.w2_weight.shape[1:], routed.w2_weight.dtype + layouts.add((type(quant_method), w13_shape, w13_dtype, w2_shape, w2_dtype)) + pairs.append((layer_idx, logical_id)) + + if len(layouts) != 1: + raise RuntimeError( + f"[FT] scale_down weight reload requires exactly one expert layout, " + f"found {len(layouts)}: {sorted(layouts, key=str)}." + ) + quant_type = next(iter(layouts))[0] + + reloaded = 0 + for start in range(0, len(pairs), _RELOAD_CHUNK_SIZE): + chunk = pairs[start : start + _RELOAD_CHUNK_SIZE] + reloaded += _reload_chunk(quant_type, routed_layers, chunk, local_slots, buckets) + return reloaded + + +def _reload_chunk( + quant_type: type, + routed_layers: list, + pairs: list[tuple[int, int]], + local_slots: dict[tuple[int, int], int], + buckets: dict[tuple[int, int], dict[str, torch.Tensor]], +) -> int: + """Apply one chunk of pairs sharing a quant scheme and layout. + + Assemble the chunk's (w13, w2) pairs on the CPU in parallel, stack them, + then pay one H2D + one format cast for the whole stack and scatter it with + one non-blocking copy per slot, drained by a single sync. Per pair that + replaces a CPU transpose-copy plus its own H2D, cast and scatter. + + Shapes must match (see ``_reload_local_slots``); stacking and the device-side + cast assume it. + """ + w13s: list[torch.Tensor] = [] + w2s: list[torch.Tensor] = [] + w13_scales: list[torch.Tensor] = [] + w2_scales: list[torch.Tensor] = [] + w8a8 = quant_type is AscendW8A8DynamicFusedMoEMethod + + # CPU phase: assemble pairs in parallel (read-only shared state). + assemble = partial(_assemble_expert, routed_layers=routed_layers, buckets=buckets, w8a8=w8a8) + with ThreadPoolExecutor(max_workers=min(len(pairs), _GATHER_MAX_WORKERS)) as pool: + for w13, w2, w13_scale, w2_scale in pool.map(assemble, pairs): + w13s.append(w13) + w2s.append(w2) + if w8a8: + w13_scales.append(w13_scale) + w2_scales.append(w2_scale) + + w13_batch = torch.stack(w13s) + w2_batch = torch.stack(w2s) + if w8a8: + w13_scale_batch = torch.stack(w13_scales) + w2_scale_batch = torch.stack(w2_scales) + + # Device phase: one H2D + one format cast per weight class. + first_layer = routed_layers[pairs[0][0]] + first_slot = local_slots[pairs[0]] + w13_weight_list = getattr(first_layer, "w13_weight_list", None) + if w13_weight_list is not None: + device = w13_weight_list[first_slot].device + w13_dtype, w2_dtype = w13_weight_list[first_slot].dtype, first_layer.w2_weight_list[first_slot].dtype + else: + device = first_layer.w13_weight.device + w13_dtype, w2_dtype = first_layer.w13_weight.dtype, first_layer.w2_weight.dtype + + if w8a8: + w13_batch = w13_batch.to(device=device) + w2_batch = w2_batch.to(device=device) + w13_scale_batch = w13_scale_batch.to(device=device, dtype=torch.float32) + w2_scale_batch = w2_scale_batch.to(device=device) + else: + # Whole-tensor layout policy; no-op when it does not force NZ. + w13_batch = w13_batch.to(device=device, dtype=w13_dtype) + w2_batch = w2_batch.to(device=device, dtype=w2_dtype) + if w8a8: + w13_batch = torch_npu.npu_format_cast(w13_batch, ACL_FORMAT_FRACTAL_NZ) + w2_batch = torch_npu.npu_format_cast(w2_batch, ACL_FORMAT_FRACTAL_NZ) + else: + w13_batch = maybe_trans_nz(w13_batch) + w2_batch = maybe_trans_nz(w2_batch) + + # Scatter phase: non-blocking copies drained by one sync. + for i, (layer_idx, logical_id) in enumerate(pairs): + slot = local_slots[(layer_idx, logical_id)] + routed = routed_layers[layer_idx] + w13_weight_list = getattr(routed, "w13_weight_list", None) + if w13_weight_list is not None: + w13_weight_list[slot].copy_(w13_batch[i], non_blocking=True) + routed.w2_weight_list[slot].copy_(w2_batch[i], non_blocking=True) + else: + # Whole-tensor layout: the slot slice is one expert matrix. + routed.w13_weight.data[slot].copy_(w13_batch[i], non_blocking=True) + routed.w2_weight.data[slot].copy_(w2_batch[i], non_blocking=True) + if w8a8: + routed.w13_weight_scale_fp32_list[slot].copy_(w13_scale_batch[i], non_blocking=True) + w2_scale_target = routed.w2_weight_scale_list[slot] + routed.w2_weight_scale_list[slot].copy_(w2_scale_batch[i].to(w2_scale_target.dtype), non_blocking=True) + # Optional fused scales (enable_fused_mc2, currently rejected + # for scale_down). + fused_w1_scale_list = getattr(routed, "fused_w1_scale_list", None) + fused_w2_scale_list = getattr(routed, "fused_w2_scale_list", None) + if fused_w1_scale_list is not None and fused_w2_scale_list is not None: + fused_w1_scale_list[slot].copy_(scale_from_float_to_int64(w13_scales[i])) + fused_w2_scale_list[slot].copy_(scale_from_float_to_int64(w2_scales[i])) + torch.npu.synchronize() + return len(pairs) + + +def _resolve_expert_tensor( + name: str, + wanted: dict[str, dict[int, tuple[int, int]]], +) -> tuple[tuple[int, int], str] | None: + """Split a normalized name into (key, suffix) if it names a wanted expert tensor. + + E.g. ``model.layers.3.ffn.experts.12.gate_proj.weight`` -> ``((3, 12), "gate_proj.weight")``. + """ + head, _, suffix = name.rpartition(".") + layer_name, _, expert_str = head.rpartition(".") + if not expert_str.isdigit(): + # 2-segment suffix ("down_proj.weight"): peel one more segment. + layer_name, _, last = layer_name.rpartition(".") + suffix = f"{expert_str}.{suffix}" + expert_str = last + by_expert = wanted.get(layer_name) + if by_expert is None or not expert_str.isdigit(): + return None + key = by_expert.get(int(expert_str)) + return (key, suffix) if key is not None else None + + +def _collect_safetensors_shards(vllm_config: VllmConfig, model: torch.nn.Module) -> list[str]: + """Collect the safetensors shards the model loads from the local checkpoint directory.""" + model_path = vllm_config.model_config.model + if not os.path.isdir(model_path): + raise RuntimeError( + f"[FT] scale_down expert reload reads a local checkpoint only, but '{model_path}' is not a directory." + ) + # A model may narrow the patterns it loads; honour that over our default. + allow_patterns = getattr(model, "allow_patterns_overrides", None) or ["*.safetensors"] + shards: list[str] = [] + for pattern in allow_patterns: + shards += glob.glob(os.path.join(model_path, pattern)) + if shards: + break + # An override may select non-safetensors files; the header walk has + # nothing to read there. + shards = [shard for shard in shards if shard.endswith(".safetensors")] + if len(shards) > 1: + # Sharded and consolidated safetensors can coexist; the index file + # records which set the model actually loads. It ships with the + # checkpoint, and the filter no-ops without it. + shards = filter_duplicate_safetensors_files(shards, model_path, SAFE_WEIGHTS_INDEX_NAME) + return sorted(shards) + + +def _collect_matching_weights( + shards: list[str], + normalize: Callable[[str], str], + wanted: dict[str, dict[int, tuple[int, int]]], + wanted_suffixes: set[str], + buckets: dict[tuple[int, int], dict[str, torch.Tensor]], + matched: set[tuple[int, int]], +) -> None: + """Read only the wanted experts' tensors by walking safetensors shard headers. + + ``safe_open(...).keys()`` reads a shard's JSON header alone, so a name + ``_resolve_expert_tensor`` rejects never reaches ``get_tensor``: reads are + limited to the wanted experts plus one header per shard. + """ + for st_file in shards: + with safe_open(st_file, framework="pt") as f: + for raw_name in f.keys(): # noqa: SIM118 + found = _resolve_expert_tensor(normalize(raw_name), wanted) + if found is None: + continue + key, suffix = found + matched.add(key) + if suffix in wanted_suffixes: + buckets.setdefault(key, {})[suffix] = f.get_tensor(raw_name) + + +def reload_experts_from_disk( + model: torch.nn.Module, + vllm_config: VllmConfig, + reassignments: set[tuple[int, int]], +) -> int: + """Reload reassigned (moe_layer_idx, logical_expert_id) weights from disk. + + Ascend keeps expert weights in runtime layout (transposed, NZ-cast, + per-slot lists, derived quant scales), so ``model.load_weights`` cannot + write them back. Only reassigned experts whose replica lands in this + rank's physical block are reloaded, into the slot given by + ``eplb_state.logical_to_physical_map`` -- hence the call must follow the + upstream ``rebuild_model_expert_maps``. + """ + if not reassignments: + return 0 + + moe_layers = list(model.moe_layers) + routed_layers = [getattr(layer, "routed_experts", layer) for layer in moe_layers] + ep_rank = get_ep_group().rank_in_group + + local_slots: dict[tuple[int, int], int] = {} + for layer_idx, logical_id in reassignments: + layer_state = getattr(moe_layers[layer_idx], "eplb_state", None) + l2p = getattr(layer_state, "logical_to_physical_map", None) + if l2p is None: + raise RuntimeError( + f"[FT] MoE layer {layer_idx} has no EPLB placement map; the reassigned expert has nowhere to land." + ) + num_local = routed_layers[layer_idx].moe_config.num_local_experts + start = ep_rank * num_local + for physical_id in l2p[logical_id].tolist(): + if start <= physical_id < start + num_local: + local_slots[(layer_idx, logical_id)] = physical_id - start + break + + if not local_slots: + return 0 + + normalize = _get_ckpt_name_mapper(model) + + # "." -> key for O(1) matching. + wanted: dict[str, dict[int, tuple[int, int]]] = {} + for layer_idx, logical_id in local_slots: + wanted.setdefault(routed_layers[layer_idx].layer_name, {})[logical_id] = (layer_idx, logical_id) + + wanted_suffixes = set(_W13_WEIGHT_SUFFIXES + _W13_SCALE_SUFFIXES) + wanted_suffixes.add(_W2_WEIGHT_SUFFIX) + wanted_suffixes.add(_W2_SCALE_SUFFIX) + + buckets: dict[tuple[int, int], dict[str, torch.Tensor]] = {} + matched: set[tuple[int, int]] = set() + + logger.info("[FT] Reloading %d reassigned (layer, expert) pair(s) on this rank from disk.", len(local_slots)) + shards = _collect_safetensors_shards(vllm_config, model) + if not shards: + raise RuntimeError( + "[FT] scale_down expert reload requires a safetensors checkpoint; " + f"{vllm_config.model_config.model} has none." + ) + _collect_matching_weights(shards, normalize, wanted, wanted_suffixes, buckets, matched) + # Same reasoning: a pair with no checkpoint weight would leave its slot + # holding the previous occupant's weights. + unmatched = [pair for pair in local_slots if pair not in matched] + if unmatched: + raise RuntimeError( + f"[FT] {len(unmatched)} (layer, expert) pair(s) had no matching " + f"checkpoint weight, e.g. {unmatched[:5]}. The model's expert " + "weights likely use a layout that does not follow " + "'..' (e.g. fused experts), or the " + "checkpoint's naming differs from the runtime namespace without " + "an hf_to_vllm_mapper declared on the model class." + ) + + reloaded = _reload_local_slots(routed_layers, local_slots, buckets) + + logger.info("[FT] Expert weight reload complete: %d (layer, expert) pair(s).", reloaded) + return reloaded diff --git a/vllm_ascend/worker/sentinel/npu_worker_sentinel.py b/vllm_ascend/worker/sentinel/npu_worker_sentinel.py new file mode 100644 index 000000000000..84705c0ee349 --- /dev/null +++ b/vllm_ascend/worker/sentinel/npu_worker_sentinel.py @@ -0,0 +1,234 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from collections.abc import Callable +from datetime import timedelta +from typing import TYPE_CHECKING + +import torch +import torch_npu +from vllm.distributed.eplb.eplb_state import _commit_eplb_maps +from vllm.distributed.parallel_state import ( + get_dp_group, + get_ep_group, + get_tp_group, +) +from vllm.distributed.utils import set_gloo_backend_timeout +from vllm.logger import logger +from vllm.model_executor.layers.fused_moe.all2all_utils import get_ep_all2all_manager +from vllm.v1.fault_tolerance.utils import FaultToleranceRequest +from vllm.v1.outputs import EMPTY_MODEL_RUNNER_OUTPUT +from vllm.v1.worker.sentinel.eplb_redistribute import ( + mark_dead_expert_slots_inplace, + rebuild_model_expert_maps, + redistribute_expert_placement, +) +from vllm.v1.worker.sentinel.gpu_worker_sentinel import ( + WorkerSentinel as GPUWorkerSentinel, +) + +from vllm_ascend.ascend_config import get_ascend_config +from vllm_ascend.distributed.eplb.state import refresh_model_routing_tables +from vllm_ascend.platform import NPUPlatform +from vllm_ascend.worker.sentinel.eplb_redistribute import ( + build_orig_to_dense_rank_table, + densify_routing_table_physical_ids, + reload_experts_from_disk, +) + +if TYPE_CHECKING: + from vllm.v1.worker.gpu_worker import Worker + + +_sentinel: "WorkerSentinel | None" = None + + +def fault_barrier_wrapper(func: Callable): + """Barrier between device faults and the async step loop. + + On the first device-touching fault (e.g. an EP-group allreduce failing on a + dead peer that poisons the model stream) it quarantines the worker and + resets the device immediately, before any tensor teardown on the broken + stream can std::terminate the process. While quarantined, wrapped methods + short-circuit with an empty output so in-flight async steps drain safely + without re-hitting the device; retry lifts the quarantine only after the + groups are rebuilt. + """ + + def wrapped(self, *args, **kwargs): + sentinel = _sentinel + if sentinel is None: + return func(self, *args, **kwargs) + if sentinel.worker_faulted: + return EMPTY_MODEL_RUNNER_OUTPUT + try: + return func(self, *args, **kwargs) + except SystemExit: + raise + except Exception as exc: + sentinel.worker_faulted = True + logger.warning( + "[FT] Quarantining %s after fault: %s", + getattr(self, "rank", "worker"), + exc, + ) + try: + sentinel.reset_device() + except Exception: + logger.exception( + "[FT] self device reset failed on %s.", + getattr(self, "rank", "worker"), + ) + return EMPTY_MODEL_RUNNER_OUTPUT + + return wrapped + + +class WorkerSentinel(GPUWorkerSentinel): + """Per-worker sentinel for fault tolerance on Ascend NPU. + + Handles commands dispatched from EngineCoreSentinel via collective_rpc, + including device restart and DP group re-initialization on retry, and + MC2 elastic_info masking + expert redistribution on scale_down. + """ + + def __init__(self, worker: "Worker", device: torch.device): + global _sentinel + self.device = device + self.worker = worker + # Set once a device-touching method faults, to keep this worker off the + # device until FT recovery rebuilds the groups. + self.worker_faulted = False + _sentinel = self + + def query_mask(self, ft_request: FaultToleranceRequest) -> dict: + """Report the dead-rank mask (upstream convention: 0=live, 1=dead).""" + return {"mask": get_ep_all2all_manager().query_active_mask().tolist()} + + def reset_device(self) -> None: + NPUPlatform.set_device(self.device) + torch_npu.npu.stop_device(self.device.index) + torch_npu.npu.restart_device(self.device.index) + torch_npu.distributed.reinit_process_group(None, False) + torch.npu.synchronize() + + def retry(self, ft_request: FaultToleranceRequest): + # Reset first so hung device collectives are aborted, then run the + # base flow and lift the quarantine after the groups are rebuilt. + self.reset_device() + super().retry(ft_request) + self.worker_faulted = False + + def init_num_local_experts(self) -> None: + """Record the per-rank physical expert slot count after model load.""" + if self.worker.model_runner.eplb_state is None: + return + eplb_model_state = self._eplb_model_state() + num_local_experts = eplb_model_state.physical_to_logical_map.shape[1] // get_ep_group().world_size + get_ep_all2all_manager().set_num_local_physical_experts(num_local_experts) + + def scale_down(self, ft_request: FaultToleranceRequest): + """Scale down over the surviving DP ranks, reusing the upstream flow. + + ``super().scale_down`` runs the deterministic dead-rank masking and + expert redistribution, dispatching ``retry`` / ``_redistribute_experts`` + to the Ascend overrides below. Ascend adds its platform preconditions + and a dummy-batch runnability check on top. + """ + self._validate_scale_down_preconditions() + super().scale_down(ft_request) + + # Verify the redistributed model is runnable before reporting healthy. + parallel_config = self.worker.parallel_config + timeout = timedelta(seconds=self.worker.parallel_config.fault_tolerance_config.engine_recovery_timeout_sec) + if parallel_config.data_parallel_size > 1: + set_gloo_backend_timeout(get_dp_group().cpu_group, timeout) + if parallel_config.tensor_parallel_size > 1: + set_gloo_backend_timeout(get_tp_group().cpu_group, timeout) + self.worker.execute_dummy_batch() + self.activate_cpu_group_timeouts(ft_request) + torch.npu.synchronize() + + def _validate_scale_down_preconditions(self) -> None: + if not self.worker.use_v2_model_runner: + raise ValueError("[FT] scale_down on Ascend NPU requires the v2 model runner.") + model_runner = self.worker.model_runner + eplb_config = self.worker.parallel_config.eplb_config + if model_runner.eplb_state is None or eplb_config.num_redundant_experts <= 0: + raise ValueError( + "[FT] scale_down requires EPLB with num_redundant_experts > 0 to re-host the dead rank's experts." + ) + ascend_config = get_ascend_config() + if ascend_config.enable_fused_mc2: + raise ValueError( + "[FT] scale_down is not supported with enable_fused_mc2: the " + "fused dispatch_ffn_combine operators take no elastic_info." + ) + if ascend_config.enable_mc2_hierarchy_comm: + raise ValueError( + "[FT] scale_down (elastic_info) is mutually exclusive with mc2 hierarchy comm (comm_alg='hierarchy')." + ) + if not hasattr(torch_npu, "npu_moe_distribute_dispatch_v2"): + raise ValueError( + "[FT] scale_down requires npu_moe_distribute_dispatch_v2 " + "(aclnn V3+); please upgrade the CANN/torch_npu version." + ) + + def _redistribute_experts(self, dead_ep_ranks: set[int]) -> None: + """Redistribute experts onto the surviving slots after scale-down. + + Overrides the upstream flow so reassigned weights reload through the + Ascend reloader: the upstream reloader writes via ``model.load_weights``, + which cannot produce Ascend's runtime expert layout (transpose / NZ / + per-slot lists / quant scales). The shared redistribution steps (mark + dead slots, steal spare slots, rebuild the logical maps) are kept; on + top of that, refreshes the Ascend kernel-facing routing tables into + the densified id space and shrinks the MC2 physical-expert width. + """ + model_runner = self.worker.model_runner + eplb_model_state = self._eplb_model_state() + + p2l = eplb_model_state.physical_to_logical_map + num_logical = eplb_model_state.logical_replica_count.shape[1] + ep_world_size = get_ep_group().world_size + num_local_experts = p2l.shape[1] // ep_world_size + + mark_dead_expert_slots_inplace(p2l, dead_ep_ranks, num_local_experts) + reassignments = redistribute_expert_placement(p2l, num_logical, num_local_experts) + # p2l was updated in place; the commit derives l2p/lrc from it. + _commit_eplb_maps(eplb_model_state, p2l.cpu()) + rebuild_model_expert_maps(model_runner.model, p2l, num_local_experts) + + if reassignments: + reload_experts_from_disk(model_runner.model, self.worker.vllm_config, reassignments) + + logger.info( + "[FT] Expert redistribution: num_logical=%d, ep_world_size=%d, reassignments=%d", + num_logical, + ep_world_size, + len(reassignments), + ) + + # Propagate the new placement into the Ascend routing tables (in-place, + # so captured graphs keep pointing at valid storage), then renumber + # their ids into the densified space for the MC2 kernels. + refresh_model_routing_tables(eplb_model_state) + self._densify_routing_tables(eplb_model_state) + + def _densify_routing_tables(self, eplb_model_state) -> None: + """Renumber the kernel-facing routing tables into the densified id space. + + The MC2 kernels consume the routing tables in the densified id space, + so the kernel-facing values are renumbered in place after the refresh. + The dead set is cumulative (accumulated across recovery rounds), so it + is derived from the manager's mask rather than this round's ranks. + """ + p2l = eplb_model_state.physical_to_logical_map + ep_world_size = get_ep_group().world_size + num_local_experts = p2l.shape[1] // ep_world_size + active_mask = get_ep_all2all_manager().query_active_mask() + dead_ranks = {rank for rank, is_dead in enumerate(active_mask.tolist()) if is_dead} + orig_to_dense_rank = build_orig_to_dense_rank_table(ep_world_size, dead_ranks) + for layer in eplb_model_state.model.moe_layers: + routing_table = getattr(getattr(layer, "eplb_state", None), "expert_replica_routing_table", None) + if routing_table is not None: + densify_routing_table_physical_ids(routing_table, orig_to_dense_rank, num_local_experts) diff --git a/vllm_ascend/worker/worker.py b/vllm_ascend/worker/worker.py index f0bd5b66658e..5a48541d8021 100644 --- a/vllm_ascend/worker/worker.py +++ b/vllm_ascend/worker/worker.py @@ -21,6 +21,7 @@ import gc import inspect import logging +import os from contextlib import AbstractContextManager, nullcontext from types import NoneType from typing import Any @@ -60,6 +61,7 @@ ) from vllm.v1.outputs import EMPTY_MODEL_RUNNER_OUTPUT, AsyncModelRunnerOutput, DraftTokenIds, ModelRunnerOutput from vllm.v1.utils import report_usage_stats +from vllm.v1.worker.gpu.async_utils import AsyncOutput from vllm.v1.worker.gpu_worker import AsyncIntermediateTensors from vllm.v1.worker.startup_plan import ( maybe_apply_startup_plan, @@ -93,6 +95,7 @@ ) from vllm_ascend.distributed.parallel_state import init_ascend_model_parallel from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton +from vllm_ascend.patch.worker.patch_v2.patch_async_output import AscendAsyncOutput from vllm_ascend.profiler.torch_npu_profiler import TorchNPUProfilerWrapper from vllm_ascend.utils import ( check_ascend_device_type, @@ -101,6 +104,7 @@ setup_ascend_local_comm_res, ) from vllm_ascend.worker.model_runner_v1 import NPUModelRunner +from vllm_ascend.worker.sentinel.npu_worker_sentinel import WorkerSentinel, fault_barrier_wrapper torch._dynamo.trace_rules.clear_lru_cache() # noqa: E402 from torch._dynamo.variables import TorchInGraphFunctionVariable # noqa: E402 @@ -184,6 +188,8 @@ def __init__( if "UnquantizedLinearMethod" in WEIGHT_LOADER_V2_SUPPORTED: WEIGHT_LOADER_V2_SUPPORTED.remove("UnquantizedLinearMethod") + self.worker_sentinel: WorkerSentinel | None = None + self.use_v2_model_runner = self.vllm_config.use_v2_model_runner self._kvpp_cache_allocation_plan: KVPPPhysicalCachePlan | None = None self._pp_send_work: list[Handle] = [] @@ -206,6 +212,10 @@ def signal_handler(signum, frame): signal.signal(signal.SIGTERM, signal_handler) signal.signal(signal.SIGINT, signal_handler) + def handle_ft_command(self, ft_request): + assert self.worker_sentinel is not None + return self.worker_sentinel.handle_command(ft_request) + def uninstall_static_kernel(self): import fcntl import os @@ -382,6 +392,50 @@ def _init_device(self): # shift self.local_rank by dp_local_rank * tp_pp_world_size so # that each DP group binds to a distinct set of NPUs. parallel_config = self.parallel_config + if self.parallel_config.enable_fault_tolerance: + if self.use_v2_model_runner: + # Model Runner V2 + fault tolerance task queue + # (TASK_QUEUE_ENABLE) hangs abnormally; force it off. + os.environ["TASK_QUEUE_ENABLE"] = "0" + logger.warning( + "Fault tolerance with Model Runner V2 does not support the " + "task queue (TASK_QUEUE_ENABLE); forcing TASK_QUEUE_ENABLE=0." + ) + + if parallel_config.tensor_parallel_size > 1: + # TP>1 relies on collective HCCL comms; disable HCCL's async + # error handling so it cannot abort the process out-of-band + # during fault-tolerance recovery. + os.environ["HCCL_ASYNC_ERROR_HANDLING"] = "0" + + abort_timeout = get_ascend_config().ft_communication_abort_timeout + if abort_timeout > 0: + # User-provided HCCL timeouts win; otherwise derive them + # from the config value. HCCL_EVENT_TIMEOUT must be + # greater than HCCL_EXEC_TIMEOUT, hence EXEC defaults to + # abort_timeout - 1. + os.environ.setdefault("HCCL_EVENT_TIMEOUT", str(abort_timeout)) + os.environ.setdefault("HCCL_EXEC_TIMEOUT", str(abort_timeout - 1)) + if int(os.environ["HCCL_EVENT_TIMEOUT"]) <= int(os.environ["HCCL_EXEC_TIMEOUT"]): + raise ValueError( + f"HCCL_EVENT_TIMEOUT ({os.environ['HCCL_EVENT_TIMEOUT']}) " + "must be greater than HCCL_EXEC_TIMEOUT " + f"({os.environ['HCCL_EXEC_TIMEOUT']})" + ) + if ( + int(os.environ["HCCL_EVENT_TIMEOUT"]) != abort_timeout + or int(os.environ["HCCL_EXEC_TIMEOUT"]) != abort_timeout - 1 + ): + logger.warning( + "Fault tolerance: HCCL communication timeouts are taken from the " + "HCCL_EVENT_TIMEOUT (%s s) / HCCL_EXEC_TIMEOUT (%s s) environment " + "variables instead of the NPU operator timeout " + "(ft_communication_abort_timeout=%s s).", + os.environ["HCCL_EVENT_TIMEOUT"], + os.environ["HCCL_EXEC_TIMEOUT"], + abort_timeout, + ) + torch.npu.set_op_timeout_ms(abort_timeout * 1000) if ( parallel_config.distributed_executor_backend not in ("ray", "external_launcher") and parallel_config.data_parallel_backend != "ray" @@ -473,6 +527,8 @@ def _init_device(self): # Initialize the distributed environment. self._init_worker_distributed_environment() + if self.parallel_config.enable_fault_tolerance: + self.worker_sentinel = WorkerSentinel(worker=self, device=device) # Set random seed. set_random_seed(self.model_config.seed) # Initialize device properties used by triton kernels. @@ -743,6 +799,7 @@ def log_memory_stats(self) -> None: self.torch_allocated / GiB_bytes, ) + @fault_barrier_wrapper def execute_model( self, scheduler_output: "SchedulerOutput", @@ -812,6 +869,7 @@ def execute_model( return output @torch.inference_mode() + @fault_barrier_wrapper def sample_tokens(self, grammar_output: "GrammarOutput") -> ModelRunnerOutput | AsyncModelRunnerOutput: output = self.model_runner.sample_tokens(grammar_output) _attach_profiling_chunk_execution_time( @@ -819,6 +877,11 @@ def sample_tokens(self, grammar_output: "GrammarOutput") -> ModelRunnerOutput | self.model_runner, output, ) + # Only the async (MRV2) path returns an AsyncOutput; re-class it so + # get_output is guarded by the fault barrier. A quarantined step drains + # as a plain ModelRunnerOutput and passes through untouched. + if self.parallel_config.enable_fault_tolerance and isinstance(output, AsyncOutput): + output.__class__ = AscendAsyncOutput return output def load_model(self) -> None: @@ -834,6 +897,9 @@ def load_model(self) -> None: with context, set_current_vllm_config(self.vllm_config): self.model_runner.load_model() + if self.worker_sentinel is not None and self.use_v2_model_runner: + self.worker_sentinel.init_num_local_experts() + if self.vllm_config.weight_transfer_config is not None: from vllm.distributed.weight_transfer.factory import ( WeightTransferEngineFactory,