diff --git a/.buildkite/test_areas/fault_tolerance.yaml b/.buildkite/test_areas/fault_tolerance.yaml new file mode 100644 index 000000000000..e2f700a8bd86 --- /dev/null +++ b/.buildkite/test_areas/fault_tolerance.yaml @@ -0,0 +1,26 @@ +group: Fault Tolerance +depends_on: + - image-build +steps: +- label: Fault Tolerance E2E (2xH100) + key: fault-tolerance-e2e-2xh100 + timeout_in_minutes: 35 + device: h100 + num_devices: 2 + working_dir: "/vllm-workspace/tests" + source_file_dependencies: + - vllm/v1/fault_tolerance/ + - vllm/v1/worker/sentinel/ + - vllm/entrypoints/serve/fault_tolerance/ + - vllm/distributed/elastic_ep/ + - vllm/distributed/device_communicators/ + - vllm/v1/engine/ + - vllm/v1/worker/ + - tests/v1/fault_tolerance/ + - tests/v1/distributed/test_external_lb_dp.py + commands: + # Base image has no nixl; install it or has_nixl_ep() skips the tests. + - bash /vllm-workspace/.buildkite/scripts/install-kv-connectors.sh + # https://github.com/NVIDIA/nccl/issues/1838 + - export NCCL_CUMEM_HOST_ENABLE=0 + - pytest -v -s v1/fault_tolerance/test_fault_tolerance_e2e.py diff --git a/tests/test_config.py b/tests/test_config.py index 71e078ef3a2a..1fc00a8f8a94 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -1547,6 +1547,16 @@ def test_needs_dp_coordination( assert vllm_config.needs_dp_coordinator == expected_needs_coordinator +def test_fault_tolerance_requires_single_api_server(): + """Fault tolerance assumes one AsyncMPClient manages all engines, so it + is incompatible with API server scale-out (_api_process_count > 1).""" + with pytest.raises(ValueError, match="single API server"): + ParallelConfig(enable_fault_tolerance=True, _api_process_count=2) + + # Single API server (the FT-supported topology) is accepted. + ParallelConfig(enable_fault_tolerance=True, _api_process_count=1) + + def test_renderer_num_workers_with_mm_cache(): """Disallow renderer_num_workers > 1 when mm processor cache is enabled, since neither cache type is thread-safe.""" diff --git a/tests/v1/fault_tolerance/__init__.py b/tests/v1/fault_tolerance/__init__.py new file mode 100644 index 000000000000..208f01a7cb5e --- /dev/null +++ b/tests/v1/fault_tolerance/__init__.py @@ -0,0 +1,2 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project diff --git a/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py b/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py new file mode 100644 index 000000000000..f6d15343b56f --- /dev/null +++ b/tests/v1/fault_tolerance/test_fault_tolerance_e2e.py @@ -0,0 +1,378 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""End-to-end tests for the elastic fault-tolerance framework. + +Requires nixl_ep FT hardware; gated behind ``has_nixl_ep()``. +""" + +import contextlib +import os +import threading +import time +from concurrent.futures import ThreadPoolExecutor +from typing import Any + +import psutil +import pytest +import requests + +from tests.utils import RemoteOpenAIServer, multi_gpu_test +from vllm.utils.import_utils import has_nixl_ep + +MODEL_NAME = os.getenv("MODEL_NAME", "ibm-research/PowerMoE-3b") +DP_SIZE = 2 + +# Fault-detection timeout budget: +# - CPU: Gloo DP allreduce timeout (30s) detects the dead peer. +# - nixl_ep: kernel masks the dead rank after Buffer's default timeout_ms=30000 (30s). +# - Deadline (45s): slowest fallback (30s) + margin. +CPU_DISTRIBUTED_TIMEOUT_S = 30 +FAULT_DETECTION_DEADLINE_S = 45 + + +# Patches ``gpu.dp_utils.sync_cudagraph_and_dp_padding`` to raise on ``rank`` at +# a chosen step. Gated on VLLM_FT_TEST_INJECT_FAULT. +_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" +_ATTR = "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, _ATTR) + _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, _ATTR, _wrapped) + + _real_import = builtins.__import__ + + def _hook(name, *a, **k): + module = _real_import(name, *a, **k) + m = sys.modules.get(_MODULE) + # During vLLM's circular import the module lands in sys.modules before + # its functions are defined; hasattr guards against patching too early. + if ( + m is not None + and hasattr(m, _ATTR) + 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-sync fn 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}") + + +def _ft_server_args() -> list[str]: + return [ + "--enforce-eager", + "--dtype", + "bfloat16", + "--max-model-len", + "2048", + "--max-num-seqs", + "128", + "--enable-expert-parallel", + "--all2all-backend", + "nixl_ep", + "--enable-fault-tolerance", + "--cpu-distributed-timeout-seconds", + str(CPU_DISTRIBUTED_TIMEOUT_S), + "--fault-tolerance-config", + '{"engine_recovery_timeout_sec": 120}', + ] + + +def _ft_manager(): + """Build the shared DP+EP fault-tolerant server topology (one engine/server).""" + from tests.v1.distributed.test_external_lb_dp import ExternalLBServerManager + + return ExternalLBServerManager( + MODEL_NAME, + DP_SIZE, + api_server_count=1, # FT requires a single API server per engine + base_server_args=_ft_server_args(), + tp_size=1, + ) + + +def _server_for_rank(servers, rank: int): + """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}") + + +def _complete(client): + """Issue the one standard completion the tests use everywhere.""" + return client.completions.create( + model=MODEL_NAME, + prompt="Hello, my name is", + max_tokens=5, + temperature=0.0, + timeout=10.0, + ) + + +def _in_parallel(fn, servers) -> list: + """Run ``fn(server)`` for all servers concurrently; return results in order.""" + with ThreadPoolExecutor(max_workers=len(servers)) as ex: + return list(ex.map(fn, servers)) + + +def _get_ft_status(server) -> dict: + resp = requests.get(server.url_for("fault_tolerance/status"), timeout=10) + resp.raise_for_status() + return resp.json() + + +def _assert_serving_and_healthy(servers) -> 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 _apply_ft(server, instruction: str, params: dict | None = None) -> dict: + """POST an FT instruction; assert it is accepted (202) and return the body.""" + resp = requests.post( + server.url_for("fault_tolerance/apply"), + json={"instruction": instruction, "params": params or {}}, + timeout=10, + ) + assert resp.status_code == 202, resp.text + return resp.json() + + +def _kill_worker_process(server) -> 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 ``/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): + """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, request_id: str, deadline_s: int) -> str | None: + """Wait until ``/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 + + +@pytest.mark.skipif(not has_nixl_ep(), reason="Requires nixl_ep all2all backend") +@multi_gpu_test(num_gpus=2) +def test_injected_fault_retry_recovers_all_ranks(monkeypatch, tmp_path): + """An exception injected into the inference path drives full retry recovery. + + Injecting an exception into ``sync_cudagraph_and_dp_padding`` at a chosen + step on rank 1. + + - Rank 1 raises inside the busy loop and goes UNHEALTHY. + - Rank 0 detects the now-absent peer via the communication timeout and also + goes UNHEALTHY. + + Both being UNHEALTHY is the precondition for ``retry``. The fault is patched + into the DP-sync fn from the test (via a generated ``sitecustomize``). + """ + fault_step = int(os.getenv("FT_FAULT_STEP", "50")) + _install_fault_injection(monkeypatch, tmp_path, rank=1, step=fault_step) + + with _ft_manager() as servers: + assert len(servers) == DP_SIZE + rank0 = _server_for_rank(servers, 0) + rank1 = _server_for_rank(servers, 1) + + # 1. Both engines healthy and serving. + _assert_serving_and_healthy((rank0, rank1)) + + # 2. Drive both ranks so rank 1 accumulates execute_model steps and trips + # the injected fault; rank 0 then times out on the DP allreduce. + with _driving(rank0, rank1): + faulted = _wait_for_engines( + [rank0, rank1], 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 " + f"{FAULT_DETECTION_DEADLINE_S}s -- it likely hung" + ) + # The rank that raised carries the fault info from its own exception. + assert faulted[1] is not None + assert faulted[1].get("fault_info"), faulted[1] + + # 3. retry both engines. + for server in (rank0, rank1): + _apply_ft(server, "retry") + + # 4. Recovery completes: both engines return to healthy and serve again. + _assert_serving_and_healthy((rank0, rank1)) + + +@pytest.mark.skipif(not has_nixl_ep(), reason="Requires nixl_ep all2all backend") +@multi_gpu_test(num_gpus=2) +def test_worker_kill_survivor_unhealthy_and_dead_rejects_retry(): + """One worker kill surfaces two status transitions at once. + + SIGKILLing only rank 1's worker leaves both EngineCores alive, so the same + fault is seen two ways: + + - Survivor (rank 0): detects the dead peer via Gloo allreduce / nixl_ep + kernel timeout. Its own executor is fine, so ``on_fault`` marks it + UNHEALTHY with a ``fault_info``. + - Victim (rank 1): detects its own executor failure and marks itself DEAD. + + Recovery is gated on UNHEALTHY: the DEAD engine accepts ``retry`` at the + HTTP layer (202 = background dispatch) but rejects it in the engine, + recording the reason as ``ft_error``. + """ + with _ft_manager() as servers: + assert len(servers) == DP_SIZE + survivor = _server_for_rank(servers, 0) + victim = _server_for_rank(servers, 1) + + # 1. Confirm both engines are healthy and serving. + _assert_serving_and_healthy((survivor, victim)) + + # 2. Kill only the victim's worker; both EngineCores stay alive. + _kill_worker_process(victim) + + # 3. Drive both engines so each keeps stepping into the failed component. + with _driving(survivor, victim): + survivor_faulted, victim_faulted = _wait_for_engines( + [survivor, victim], + match_key="status", + match_values={"dead", "unhealthy"}, + ) + + assert survivor_faulted is not None, ( + "survivor did not report the peer fault within " + f"{FAULT_DETECTION_DEADLINE_S}s -- it likely hung" + ) + # The survivor's own executor is fine, so it must be UNHEALTHY, not DEAD. + assert survivor_faulted["status"] == "unhealthy", survivor_faulted + assert survivor_faulted.get("fault_info"), survivor_faulted + + assert victim_faulted is not None, ( + "victim did not report its worker's death within " + f"{FAULT_DETECTION_DEADLINE_S}s" + ) + assert victim_faulted["status"] == "dead", victim_faulted + + # 4. retry is accepted at the HTTP layer (202 = background dispatch)... + request_id = _apply_ft(victim, "retry")["request_id"] + + # 5. ...but the DEAD engine must reject it: recovery requires UNHEALTHY. + 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 /fault_tolerance/status" + ) + assert "status is DEAD" in ft_error, ft_error diff --git a/vllm/config/__init__.py b/vllm/config/__init__.py index 6070a3f82382..24a8b5a31713 100644 --- a/vllm/config/__init__.py +++ b/vllm/config/__init__.py @@ -13,6 +13,7 @@ from vllm.config.diffusion import DiffusionConfig from vllm.config.ec_manager_config import EncoderCacheManagerConfig from vllm.config.ec_transfer import ECTransferConfig +from vllm.config.fault_tolerance import FaultToleranceConfig from vllm.config.kernel import KernelConfig from vllm.config.kv_events import KVEventsConfig from vllm.config.kv_transfer import KVTransferConfig @@ -124,6 +125,8 @@ "StructuredOutputsConfig", # From vllm.config.profiler "ProfilerConfig", + # From vllm.config.fault_tolerance + "FaultToleranceConfig", # From vllm.config.utils "ConfigType", "SupportsMetricsInfo", diff --git a/vllm/config/fault_tolerance.py b/vllm/config/fault_tolerance.py new file mode 100644 index 000000000000..7ed095d204f6 --- /dev/null +++ b/vllm/config/fault_tolerance.py @@ -0,0 +1,18 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + + +from vllm.config.utils import config + + +@config +class FaultToleranceConfig: + """Configuration for fault tolerance.""" + + engine_recovery_timeout_sec: int = 120 + """Timeout (in seconds) to wait for error handling instructions + before raising an exception. If the EngineCore encounters an + error, it waits up to this many seconds for vLLM to receive + instructions on how to handle the error and then recover from the fault. + If vLLM does not recover during this time, the original error is raised. + """ diff --git a/vllm/config/parallel.py b/vllm/config/parallel.py index ce038a99ea9a..949eb298a170 100644 --- a/vllm/config/parallel.py +++ b/vllm/config/parallel.py @@ -13,6 +13,7 @@ from typing_extensions import Self import vllm.envs as envs +from vllm.config.fault_tolerance import FaultToleranceConfig from vllm.config.utils import config from vllm.logger import init_logger from vllm.platforms import current_platform @@ -22,6 +23,7 @@ from ray.runtime_env import RuntimeEnv from ray.util.placement_group import PlacementGroup + from vllm.config.fault_tolerance import FaultToleranceConfig from vllm.v1.executor import Executor else: RuntimeEnv = Any @@ -393,6 +395,16 @@ class is dynamically inherited by the worker class. This is used to inject should only be set by API server scale-out. """ + enable_fault_tolerance: bool = False + """Enable fault tolerance for detailed error recovery, + such as scaling down fault DPEngineCore. + """ + + fault_tolerance_config: FaultToleranceConfig = Field( + default_factory=FaultToleranceConfig + ) + """The configurations for fault tolerance.""" + @field_validator("disable_nccl_for_dp_synchronization", mode="wrap") @classmethod def _skip_none_validation(cls, value: Any, handler: Callable) -> Any: @@ -445,6 +457,13 @@ def _validate_parallel_config(self) -> Self: f"but found: {self._api_process_rank}" ) + if self.enable_fault_tolerance and self._api_process_count > 1: + raise ValueError( + "Fault tolerance requires a single API server process " + f"(--api-server-count=1), but got {self._api_process_count}. " + "The FT system assumes one AsyncMPClient manages all engines." + ) + if self.all2all_backend in ["pplx", "naive"]: logger.warning( "The '%s' all2all backend has been removed. " diff --git a/vllm/distributed/device_communicators/all2all.py b/vllm/distributed/device_communicators/all2all.py index 679764f6a82e..ee404f688a19 100644 --- a/vllm/distributed/device_communicators/all2all.py +++ b/vllm/distributed/device_communicators/all2all.py @@ -8,6 +8,7 @@ import torch.distributed as dist import vllm.envs as envs +from vllm.config import get_current_vllm_config from vllm.distributed import get_dp_group, get_ep_group, get_pcp_group from vllm.distributed.utils import StatelessProcessGroup from vllm.forward_context import get_forward_context @@ -278,7 +279,9 @@ class DeepEPLLAll2AllManager(DeepEPAll2AllManagerBase): def __init__(self, cpu_group, tcp_store_group=None): super().__init__(cpu_group, tcp_store_group) - self.support_fault_tolerance = False # TODO: set to True when FT is supported. + self.support_fault_tolerance = ( + get_current_vllm_config().parallel_config.enable_fault_tolerance + ) def _make_all2all_kwargs( self, @@ -360,6 +363,16 @@ def query_fault(self) -> torch.Tensor: has_fault = (current != DeepEPLLAll2AllManager._last_mask).any() return has_fault + def clean_buffers(self) -> None: + buf = DeepEPLLAll2AllManager._buffer + if buf is None: + return + buf.get_local_buffer_tensor(dtype=torch.int8, use_rdma_buffer=True).zero_() + torch.accelerator.synchronize() + buf.low_latency_clean_mask_buffer() + torch.accelerator.synchronize() + DeepEPLLAll2AllManager._last_mask = None + @dataclass class _NixlEPBufferState: @@ -565,6 +578,18 @@ def query_fault(self) -> torch.Tensor: has_fault = (current != last).any() return has_fault + def clean_buffers(self) -> None: + if NixlEPAll2AllManager._buffer is None: + return + state = NixlEPAll2AllManager._buffer + state.buffer.get_local_buffer_tensor( + dtype=torch.int8, use_rdma_buffer=True + ).zero_() + torch.accelerator.synchronize() + state.buffer.clean_mask_buffer() + torch.accelerator.synchronize() + NixlEPAll2AllManager._last_mask = None + class FlashInferNVLinkTwoSidedManager(All2AllManagerBase): """ diff --git a/vllm/distributed/device_communicators/base_device_communicator.py b/vllm/distributed/device_communicators/base_device_communicator.py index 70f1fb5d62c5..dc2671433891 100644 --- a/vllm/distributed/device_communicators/base_device_communicator.py +++ b/vllm/distributed/device_communicators/base_device_communicator.py @@ -105,10 +105,30 @@ def dispatch( raise NotImplementedError def query_active_mask(self) -> torch.Tensor: + """Return the all2all liveness mask for the EP ranks. + + Returns: + An int32 device tensor where 0 marks a live rank and 1 marks a + masked (dead/unreachable) rank. + """ raise NotImplementedError def query_fault(self) -> torch.Tensor: - """Returns has_fault scalar.""" + """Return a scalar bool tensor, True if a new fault appeared. + + Compares the current mask against the baseline recorded at the last + recovery point. + """ + raise NotImplementedError + + def clean_buffers(self) -> None: + """Reset this rank's RDMA buffers and all2all mask state (rank-local). + + Post-fault cleanup: a dispatch/combine that hit a dead peer or timed + out can leave partially-written or stale tokens in the RDMA receive + buffer, so it is zeroed to stop the next forward from reading that + contaminated data. + """ raise NotImplementedError def set_num_sms(self, num_sms: int): diff --git a/vllm/engine/arg_utils.py b/vllm/engine/arg_utils.py index a14de27190ec..0244ff192cd6 100644 --- a/vllm/engine/arg_utils.py +++ b/vllm/engine/arg_utils.py @@ -42,6 +42,7 @@ DiffusionConfig, ECTransferConfig, EPLBConfig, + FaultToleranceConfig, KernelConfig, KVEventsConfig, KVTransferConfig, @@ -722,6 +723,11 @@ class EngineArgs: optimization_level: OptimizationLevel = VllmConfig.optimization_level performance_mode: PerformanceMode = VllmConfig.performance_mode + fault_tolerance_config: FaultToleranceConfig = get_field( + ParallelConfig, "fault_tolerance_config" + ) + enable_fault_tolerance: bool = ParallelConfig.enable_fault_tolerance + kv_offloading_size: float | None = CacheConfig.kv_offloading_size kv_offloading_backend: KVOffloadingBackend = CacheConfig.kv_offloading_backend tokens_only: bool = False @@ -754,6 +760,16 @@ def __post_init__(self): self.weight_transfer_config = WeightTransferConfig( **self.weight_transfer_config ) + if isinstance(self.fault_tolerance_config, dict): + if not self.enable_fault_tolerance: + logger.warning( + "--fault-tolerance-config was passed. Fault tolerance is being " + "automatically enabled." + ) + self.enable_fault_tolerance = True + self.fault_tolerance_config = FaultToleranceConfig( + **self.fault_tolerance_config + ) if isinstance(self.ir_op_priority, dict): self.ir_op_priority = IrOpPriorityConfig(**self.ir_op_priority) @@ -1151,6 +1167,12 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: parallel_group.add_argument( "--worker-extension-cls", **parallel_kwargs["worker_extension_cls"] ) + parallel_group.add_argument( + "--enable-fault-tolerance", **parallel_kwargs["enable_fault_tolerance"] + ) + parallel_group.add_argument( + "--fault-tolerance-config", **parallel_kwargs["fault_tolerance_config"] + ) # KV cache arguments cache_kwargs = get_kwargs(CacheConfig) @@ -1617,6 +1639,7 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: def from_cli_args(cls, args: argparse.Namespace): # Get the list of attributes of this dataclass. attrs = [attr.name for attr in dataclasses.fields(cls)] + # Set the attributes from the parsed arguments. engine_args = cls( **{attr: getattr(args, attr) for attr in attrs if hasattr(args, attr)} @@ -2022,6 +2045,12 @@ def create_engine_config( data_parallel_external_lb = ( self.data_parallel_external_lb or self.data_parallel_rank is not None ) + if self.enable_fault_tolerance and not data_parallel_external_lb: + raise ValueError( + "Fault tolerance requires external load balancer mode " + "(--data-parallel-external-lb or --data-parallel-rank). " + "Internal LB mode is not supported." + ) if ( self.data_parallel_size > 1 and data_parallel_external_lb @@ -2179,6 +2208,8 @@ def create_engine_config( _api_process_count=self._api_process_count, _api_process_rank=self._api_process_rank, assigned_physical_gpu_ids=self._resolve_device_ids(), + enable_fault_tolerance=self.enable_fault_tolerance, + fault_tolerance_config=self.fault_tolerance_config, numa_bind=self.numa_bind, numa_bind_nodes=self.numa_bind_nodes, numa_bind_cpus=self.numa_bind_cpus, diff --git a/vllm/engine/protocol.py b/vllm/engine/protocol.py index c54123bea9e5..ef3be178ac8d 100644 --- a/vllm/engine/protocol.py +++ b/vllm/engine/protocol.py @@ -20,6 +20,7 @@ from vllm.tasks import SupportedTask from vllm.v1.engine import EngineCoreRequest from vllm.v1.engine.input_processor import InputProcessor +from vllm.v1.fault_tolerance.utils import FaultToleranceRequest, FaultToleranceResult if TYPE_CHECKING: from vllm.v1.engine import PauseMode @@ -234,6 +235,16 @@ async def collective_rpc( """Perform a collective RPC call to the given path.""" raise NotImplementedError + async def handle_fault( + self, fault_tolerance_request: FaultToleranceRequest + ) -> FaultToleranceResult: + """send fault tolerance instruction to the engine""" + raise NotImplementedError + + async def get_status(self): + """Get fault tolerance status of all engines.""" + raise NotImplementedError + async def get_supported_tasks(self) -> tuple[SupportedTask, ...]: """Get supported tasks""" raise NotImplementedError diff --git a/vllm/entrypoints/openai/api_server.py b/vllm/entrypoints/openai/api_server.py index 59c7ee84caef..9103dd7fae95 100644 --- a/vllm/entrypoints/openai/api_server.py +++ b/vllm/entrypoints/openai/api_server.py @@ -269,6 +269,13 @@ def build_app( register_pooling_api_routers(app, supported_tasks, model_config) + if args.enable_fault_tolerance: + from vllm.entrypoints.serve.fault_tolerance.api_router import ( + register_fault_tolerance_api_router, + ) + + register_fault_tolerance_api_router(app) + # Endpoint plugins are attached last so their routes are registered after all core # routers. This runs even for the CPU only render server. A plugin eligible for # the `render` task still gets its routes registered. It receives diff --git a/vllm/entrypoints/serve/fault_tolerance/__init__.py b/vllm/entrypoints/serve/fault_tolerance/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/vllm/entrypoints/serve/fault_tolerance/api_router.py b/vllm/entrypoints/serve/fault_tolerance/api_router.py new file mode 100644 index 000000000000..960833af2155 --- /dev/null +++ b/vllm/entrypoints/serve/fault_tolerance/api_router.py @@ -0,0 +1,100 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import json +import uuid +from http import HTTPStatus + +from fastapi import APIRouter, BackgroundTasks, Depends, FastAPI, HTTPException, Request +from fastapi.responses import JSONResponse + +from vllm.engine.protocol import EngineClient +from vllm.entrypoints.openai.engine.protocol import ErrorResponse +from vllm.entrypoints.serve.utils.api_utils import validate_json_request +from vllm.logger import init_logger +from vllm.v1.fault_tolerance.utils import FaultToleranceRequest + +logger = init_logger(__name__) + +router = APIRouter() + +_ALLOWED_INSTRUCTIONS = {"retry"} + + +def _validate_payload(body: dict) -> tuple[str, dict]: + if not isinstance(body, dict): + raise HTTPException(400, "Request body must be a JSON object.") + instruction = body.get("instruction") + if not instruction: + raise HTTPException(400, "'instruction' is required.") + if instruction not in _ALLOWED_INSTRUCTIONS: + raise HTTPException(400, f"Invalid instruction: '{instruction}'.") + params = body.get("params", {}) + if not isinstance(params, dict): + raise HTTPException(400, "'params' must be an object.") + return instruction, params + + +@router.post( + "/fault_tolerance/apply", + dependencies=[Depends(validate_json_request)], + responses={ + HTTPStatus.ACCEPTED.value: {"model": dict}, + HTTPStatus.BAD_REQUEST.value: {"model": ErrorResponse}, + }, +) +async def process_fault_tolerance_instruction( + raw_request: Request, background_tasks: BackgroundTasks +): + try: + body = await raw_request.json() + except json.JSONDecodeError as e: + raise HTTPException(400, "Invalid JSON format") from e + + instruction, params = _validate_payload(body) + ft_request = FaultToleranceRequest( + instruction=instruction, + params=params, + request_id=str(uuid.uuid4()), + ) + + client: EngineClient = raw_request.app.state.engine_client + # Recovery runs cross-rank collective ops that only complete once every rank + # has been dispatched. Run it in the background and return immediately so the + # orchestrator can dispatch to all ranks without blocking; completion is + # observed by polling GET /fault_tolerance/status. + background_tasks.add_task(_run_fault_recovery, client, ft_request) + return JSONResponse( + status_code=HTTPStatus.ACCEPTED.value, + content={ + "message": "Request accepted; poll /fault_tolerance/status for updates.", + "request_id": ft_request.request_id, + }, + background=background_tasks, + ) + + +async def _run_fault_recovery( + client: EngineClient, ft_request: FaultToleranceRequest +) -> None: + """Drive recovery to completion after the 202 response is sent.""" + try: + result = await client.handle_fault(ft_request) + except Exception: + logger.exception("[FT] Recovery dispatch failed.") + return + if not result.success: + logger.error( + "[FT] Recovery failed for request %s: %s", + ft_request.request_id, + result.reason, + ) + + +@router.get("/fault_tolerance/status") +async def get_status(raw_request: Request): + client: EngineClient = raw_request.app.state.engine_client + return JSONResponse(content=await client.get_status()) + + +def register_fault_tolerance_api_router(app: FastAPI): + app.include_router(router) diff --git a/vllm/v1/engine/__init__.py b/vllm/v1/engine/__init__.py index 919402a16ab0..4ac27be5068f 100644 --- a/vllm/v1/engine/__init__.py +++ b/vllm/v1/engine/__init__.py @@ -31,6 +31,8 @@ EEP_NOTIFICATION_CALL_ID = -1 +FT_STATUS_CALL_ID = -2 + class EEPNotificationType(enum.Enum): NEW_CORE_ENGINES_INIT_READY = "NEW_CORE_ENGINES_INIT_READY" @@ -282,3 +284,9 @@ class ReconfigureRankType(enum.IntEnum): KEEP_CURRENT_RANK = -1 SHUTDOWN_CURRENT_RANK = -2 + + +class EngineStatusType(enum.IntEnum): + HEALTHY = 0 + DEAD = 1 + UNHEALTHY = 2 diff --git a/vllm/v1/engine/async_llm.py b/vllm/v1/engine/async_llm.py index 93e02abf7479..f1e7132339c3 100644 --- a/vllm/v1/engine/async_llm.py +++ b/vllm/v1/engine/async_llm.py @@ -44,6 +44,7 @@ from vllm.v1.engine.output_processor import OutputProcessor, RequestOutputCollector from vllm.v1.engine.parallel_sampling import ParentRequest from vllm.v1.executor import Executor +from vllm.v1.fault_tolerance.utils import FaultToleranceRequest, FaultToleranceResult from vllm.v1.metrics.loggers import ( StatLoggerFactory, StatLoggerManager, @@ -1041,6 +1042,15 @@ async def scale_elastic_ep( finally: set_scaling_elastic_ep(False) + async def handle_fault( + self, fault_tolerance_request: FaultToleranceRequest + ) -> FaultToleranceResult: + """send fault tolerance instruction to the engine""" + return await self.engine_core.handle_fault(fault_tolerance_request) + + async def get_status(self): + return await self.engine_core.get_status() + @property def is_running(self) -> bool: # Is None before the loop is started. diff --git a/vllm/v1/engine/core.py b/vllm/v1/engine/core.py index 476f53d46115..8a62c8b2b1be 100644 --- a/vllm/v1/engine/core.py +++ b/vllm/v1/engine/core.py @@ -78,6 +78,11 @@ get_physical_gpu_ids_for_local_dp_rank, ) from vllm.v1.executor import Executor +from vllm.v1.fault_tolerance.engine_core_sentinel import ( + FT_UTILITY_METHOD, + EngineCoreSentinel, + fault_tolerant_wrapper, +) from vllm.v1.kv_cache_interface import KVCacheConfig, get_kv_cache_spec_kind from vllm.v1.metrics.stats import SchedulerIterationDetails, SchedulerStats from vllm.v1.outputs import ModelRunnerOutput @@ -1070,6 +1075,16 @@ def __init__( internal_dp_balancing, ) + # Initialize fault tolerance settings. + self.enable_fault_tolerance = ( + vllm_config.parallel_config.enable_fault_tolerance + ) + if self.enable_fault_tolerance: + self.ft_sentinel = EngineCoreSentinel( + engine=self, + parallel_config=vllm_config.parallel_config, + ) + # Background Threads and Queues for IO. These enable us to # overlap ZMQ socket IO with GPU since they release the GIL, # and to overlap some serialization/deserialization with the @@ -1355,6 +1370,7 @@ def is_running(self) -> bool: """Returns true if shutdown has not been requested.""" return self.shutdown_state == EngineShutdownState.RUNNING + @fault_tolerant_wrapper def run_busy_loop(self): """Core busy loop of the EngineCore.""" while self._handle_shutdown(): @@ -1672,6 +1688,14 @@ def process_input_sockets( except Exception: self._handle_request_preproc_error(req) continue + elif request_type == EngineCoreRequestType.UTILITY: + request = generic_decoder.decode(data_frames) + client_idx, call_id, method, args = request + if method == FT_UTILITY_METHOD: + self.ft_sentinel.handle_command( + client_idx, call_id, args[0] + ) + continue else: request = generic_decoder.decode(data_frames) @@ -2021,6 +2045,7 @@ def _should_throttle_prefills(self) -> bool: and self.step_counter % self.prefill_schedule_interval != 0 ) + @fault_tolerant_wrapper def run_busy_loop(self): """Core busy loop of the EngineCore for data parallel case.""" diff --git a/vllm/v1/engine/core_client.py b/vllm/v1/engine/core_client.py index bcb441e7564a..f83e32096a90 100644 --- a/vllm/v1/engine/core_client.py +++ b/vllm/v1/engine/core_client.py @@ -16,6 +16,7 @@ from threading import Thread from typing import Any, TypeAlias, TypeVar +import msgspec import msgspec.msgpack import zmq import zmq.asyncio @@ -35,6 +36,7 @@ ) from vllm.v1.engine import ( EEP_NOTIFICATION_CALL_ID, + FT_STATUS_CALL_ID, EEPNotificationType, EngineCoreOutputs, EngineCoreReadyResponse, @@ -56,6 +58,11 @@ launch_core_engines, ) from vllm.v1.executor import Executor +from vllm.v1.fault_tolerance.engine_core_sentinel import FT_UTILITY_METHOD +from vllm.v1.fault_tolerance.utils import ( + FaultToleranceRequest, + FaultToleranceResult, +) from vllm.v1.pool.late_interaction import get_late_interaction_engine_index from vllm.v1.serial_utils import MsgpackDecoder, MsgpackEncoder, bytestr @@ -272,6 +279,14 @@ async def collective_rpc_async( ) -> list[_R]: raise NotImplementedError + async def handle_fault( + self, fault_tolerance_request: FaultToleranceRequest + ) -> FaultToleranceResult: + raise NotImplementedError + + async def get_status(self): + raise NotImplementedError + class InprocClient(EngineCoreClient): """ @@ -971,6 +986,14 @@ def __init__( self.client_count = client_count self.client_index = client_index self.outputs_queue = asyncio.Queue[EngineCoreOutputs | Exception]() + + # locally-cached engine status + self._engine_status: dict[int, dict] = {} + if self.vllm_config.parallel_config.enable_fault_tolerance: + self._engine_status = { + rank: {"id": rank, "status": "healthy"} + for rank in self.engine_ranks_managed + } try: # If we are running in an asyncio event loop, start the queue task. # Otherwise, it will be started lazily. If it is not started here, @@ -994,7 +1017,7 @@ def _ensure_output_queue_task(self): output_handler: ( Callable[[AsyncMPClient, EngineCoreOutputs], Awaitable[None]] | None ) = getattr(self.__class__, "process_engine_outputs", None) - _self_ref = weakref.ref(self) if output_handler else None + _self_ref = weakref.ref(self) output_socket = resources.output_socket assert output_socket is not None @@ -1025,6 +1048,14 @@ async def process_outputs_socket(): asyncio.create_task( notification_callback_handler(_self, notification_data) ) + elif outputs.utility_output.call_id == FT_STATUS_CALL_ID: + _self = _self_ref() + if not _self: + return + if outputs.utility_output.result is not None: + _self._engine_status[outputs.engine_index] = ( + outputs.utility_output.result.result + ) else: _process_utility_output( outputs.utility_output, utility_results @@ -1196,6 +1227,25 @@ async def collective_rpc_async( "collective_rpc", method, timeout, args, kwargs ) + async def handle_fault( + self, ft_request: FaultToleranceRequest + ) -> FaultToleranceResult: + res = await self.call_utility_async(FT_UTILITY_METHOD, ft_request) + result = msgspec.convert(res, FaultToleranceResult) + if not result.success: + status = self._engine_status.get(self.engine_ranks_managed[0]) + if status is not None: + status["last_ft_request_id"] = result.request_id + status["ft_error"] = result.reason + return result + + async def get_status(self): + return { + "schema_version": 1, + "total_engines": len(self.engine_ranks_managed), + "engines": list(self._engine_status.values()), + } + class DPAsyncMPClient(AsyncMPClient): """Asyncio-compatible client for multi-proc, multi-engine (data parallel) diff --git a/vllm/v1/engine/utils.py b/vllm/v1/engine/utils.py index 093f065475ab..db1896b09468 100644 --- a/vllm/v1/engine/utils.py +++ b/vllm/v1/engine/utils.py @@ -234,9 +234,6 @@ def monitor_engine_liveness(self) -> None: if exitcode != 0 and not self.manager_stopped.is_set(): self.failed_proc_name = proc.name if died_sentinels: - # Any engine exit currently triggers a shutdown. Future - # work (e.g., Elastic and fault-tolerant EP) will add finer-grained - # handling for different exit scenarios. break self.shutdown() diff --git a/vllm/v1/fault_tolerance/__init__.py b/vllm/v1/fault_tolerance/__init__.py new file mode 100644 index 000000000000..ee54b4eb2e9c --- /dev/null +++ b/vllm/v1/fault_tolerance/__init__.py @@ -0,0 +1,8 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from .engine_core_sentinel import EngineCoreSentinel, fault_tolerant_wrapper + +__all__ = [ + "EngineCoreSentinel", + "fault_tolerant_wrapper", +] diff --git a/vllm/v1/fault_tolerance/engine_core_sentinel.py b/vllm/v1/fault_tolerance/engine_core_sentinel.py new file mode 100644 index 000000000000..1d82dd9f7435 --- /dev/null +++ b/vllm/v1/fault_tolerance/engine_core_sentinel.py @@ -0,0 +1,197 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""EngineCoreSentinel and fault_tolerant_wrapper for the engine core.""" + +import json +import threading +from collections.abc import Callable +from typing import TYPE_CHECKING + +import msgspec + +from vllm.config import set_current_vllm_config +from vllm.distributed import stateless_destroy_torch_distributed_process_group +from vllm.distributed.utils import stateless_init_torch_distributed_process_group +from vllm.logger import init_logger +from vllm.utils.network_utils import get_open_port +from vllm.v1.engine import ( + FT_STATUS_CALL_ID, + EngineCoreOutputs, + EngineStatusType, + UtilityOutput, +) +from vllm.v1.fault_tolerance.utils import FaultToleranceRequest, FaultToleranceResult +from vllm.v1.request import RequestStatus +from vllm.v1.serial_utils import UtilityResult, run_method + +if TYPE_CHECKING: + from vllm.v1.engine.core import EngineCoreProc + +logger = init_logger(__name__) + +FT_UTILITY_METHOD = "handle_fault_tolerance" + + +class EngineCoreSentinel: + """Manages fault tolerance state for a single engine core.""" + + def __init__(self, engine: "EngineCoreProc", parallel_config): + self.engine = engine + self.engine_index = engine.engine_index + self.parallel_config = parallel_config + ft_config = parallel_config.fault_tolerance_config + self.engine_recovery_timeout_sec = ft_config.engine_recovery_timeout_sec + + self.resumed = threading.Event() + self.resumed.set() + self.status_type = EngineStatusType.HEALTHY + self.fault_info: str | None = None + self._dp_reinit_epoch = 0 + + def handle_command(self, client_idx: int, call_id: int, ft_args: dict): + """Dispatch an FT command by instruction name.""" + ft_request = FaultToleranceRequest(**ft_args) + if self.status_type != EngineStatusType.UNHEALTHY: + reason = ( + f"[FT] Rejecting {ft_request.instruction} on engine " + f"{self.engine_index}: status is {self.status_type.name}" + ) + logger.warning(reason) + result = FaultToleranceResult( + request_id=ft_request.request_id, + success=False, + reason=reason, + ) + else: + try: + result = run_method(self, ft_request.instruction, (ft_request,), {}) + except Exception as e: + logger.exception("[FT] Instruction '%s' failed", ft_request.instruction) + result = FaultToleranceResult( + request_id=ft_request.request_id, success=False, reason=str(e) + ) + + uo = UtilityOutput(call_id) + uo.result = UtilityResult(msgspec.structs.asdict(result)) + self.engine.output_queue.put_nowait( + (client_idx, EngineCoreOutputs(utility_output=uo)) + ) + + def on_fault(self, exc: Exception): + """Called by the wrapper when the busy loop raises an exception.""" + self.resumed.clear() + logger.warning( + "[FT] Busy loop raised %s. Waiting for recovery.", type(exc).__name__ + ) + + engine = self.engine + aborted = engine.scheduler.finish_requests(None, RequestStatus.FINISHED_ABORTED) + engine._send_abort_outputs(aborted) + if engine.batch_queue is not None: + engine.batch_queue.clear() + if ( + hasattr(engine.model_executor, "is_failed") + and engine.model_executor.is_failed + ): + self.status_type = EngineStatusType.DEAD + else: + self.status_type = EngineStatusType.UNHEALTHY + self.fault_info = f"{type(exc).__name__}" + logger.info( + "[FT] Engine %d status -> %s:", + self.engine_index, + self.status_type.name, + exc_info=exc, + ) + self._push_status() + + def _push_status(self): + """Push current health to the client so it can refresh its cache.""" + payload = {"id": self.engine_index, "status": self.status_type.name.lower()} + if self.status_type == EngineStatusType.UNHEALTHY: + payload["fault_info"] = self.fault_info + outputs = EngineCoreOutputs( + utility_output=UtilityOutput( + call_id=FT_STATUS_CALL_ID, + result=UtilityResult(payload), + ) + ) + outputs.engine_index = self.engine_index + self.engine.output_queue.put_nowait((0, outputs)) + + def retry(self, ft_request: FaultToleranceRequest) -> FaultToleranceResult: + engine = self.engine + executor = engine.model_executor + + with set_current_vllm_config(engine.vllm_config): + ft_request.params.update(self._reinit_dp_group()) + if hasattr(engine, "step_counter"): + engine.step_counter = 0 + + executor.collective_rpc("handle_ft_command", args=(ft_request,)) + + self.status_type = EngineStatusType.HEALTHY + logger.info("[FT] Engine %d status -> HEALTHY", self.engine_index) + self.resumed.set() + self._push_status() + return FaultToleranceResult(request_id=ft_request.request_id, success=True) + + def _reinit_dp_group(self) -> dict: + """Reinit DP process group if in DP mode. Returns worker params.""" + engine = self.engine + if not hasattr(engine, "dp_group") or not hasattr(engine, "dp_store"): + return {} + + parallel_config = engine.vllm_config.parallel_config + worker_key = f"ft_worker_dp_ports_{self._dp_reinit_epoch}" + engine_key = f"ft_engine_dp_port_{self._dp_reinit_epoch}" + self._dp_reinit_epoch += 1 + + if parallel_config.data_parallel_rank == 0: + worker_ports = [get_open_port() for _ in range(parallel_config.world_size)] + engine_port = get_open_port() + engine.dp_store.set(worker_key, json.dumps(worker_ports).encode()) + engine.dp_store.set(engine_key, str(engine_port).encode()) + else: + worker_ports = json.loads(engine.dp_store.get(worker_key).decode()) + engine_port = int(engine.dp_store.get(engine_key).decode()) + + stateless_destroy_torch_distributed_process_group(engine.dp_group) + engine.dp_group, engine.dp_store = ( + stateless_init_torch_distributed_process_group( + parallel_config.data_parallel_master_ip, + engine_port, + parallel_config.data_parallel_rank, + parallel_config.data_parallel_size, + backend="gloo", + return_store=True, + ) + ) + return {"new_stateless_dp_group_ports": worker_ports} + + +def fault_tolerant_wrapper(busy_loop_func: Callable): + """Wrap the busy loop to catch faults and delegate recovery.""" + + def run_with_fault_tolerance(self: "EngineCoreProc"): + while True: + try: + busy_loop_func(self) + except SystemExit: + raise + except Exception as exc: + if not self.enable_fault_tolerance: + raise + self.ft_sentinel.on_fault(exc) + recovered = self.ft_sentinel.resumed.wait( + timeout=self.ft_sentinel.engine_recovery_timeout_sec + ) + if recovered: + continue + logger.error( + "[FT] No recovery within %ds timeout.", + self.ft_sentinel.engine_recovery_timeout_sec, + ) + raise + + return run_with_fault_tolerance diff --git a/vllm/v1/fault_tolerance/utils.py b/vllm/v1/fault_tolerance/utils.py new file mode 100644 index 000000000000..0c1b1689b01a --- /dev/null +++ b/vllm/v1/fault_tolerance/utils.py @@ -0,0 +1,17 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from typing import Any + +import msgspec + + +class FaultToleranceResult(msgspec.Struct): + request_id: str + success: bool + reason: str | None = None + + +class FaultToleranceRequest(msgspec.Struct): + instruction: str + params: dict[str, Any] + request_id: str = "" diff --git a/vllm/v1/worker/gpu/async_utils.py b/vllm/v1/worker/gpu/async_utils.py index e4659104f49e..4570d7267344 100644 --- a/vllm/v1/worker/gpu/async_utils.py +++ b/vllm/v1/worker/gpu/async_utils.py @@ -5,6 +5,7 @@ import numpy as np import torch +from vllm.model_executor.layers.fused_moe.all2all_utils import get_ep_all2all_manager from vllm.v1.outputs import AsyncModelRunnerOutput, LogprobsTensors, ModelRunnerOutput from vllm.v1.worker.gpu.sample.output import SamplerOutput @@ -17,6 +18,7 @@ def __init__( num_sampled_tokens: torch.Tensor, main_stream: torch.cuda.Stream, copy_stream: torch.cuda.Stream, + check_ep_fault: bool = False, ): # NOTE(woosuk): We must retain references to the GPU tensors, # as the copy operations are performed on a different CUDA stream than @@ -26,6 +28,7 @@ def __init__( self.num_sampled_tokens = num_sampled_tokens # Blocking (sleep) event to avoid busy-polling the CUDA driver lock. self.copy_event = torch.cuda.Event(blocking=True) + self._has_fault: torch.Tensor | None = None with stream(copy_stream, main_stream): copy_stream.wait_stream(main_stream) @@ -44,6 +47,9 @@ def __init__( k: v.to_cpu_nonblocking() if v is not None else None for k, v in self.model_runner_output.prompt_logprobs_dict.items() } + if check_ep_fault: + has_fault = get_ep_all2all_manager().query_fault() + self._has_fault = has_fault.to("cpu", non_blocking=True) self.copy_event.record(copy_stream) def get_output(self) -> ModelRunnerOutput: @@ -67,6 +73,15 @@ def get_output(self) -> ModelRunnerOutput: if self.logprobs_tensors is not None: self.model_runner_output.logprobs = self.logprobs_tensors.tolists() self.model_runner_output.prompt_logprobs_dict = self.prompt_logprobs_dict + + if self._has_fault is not None and self._has_fault.item(): + mask = get_ep_all2all_manager().query_active_mask() + raise RuntimeError( + "Fault detected in EP all2all communication: " + "one or more ranks timed out during dispatch/combine. " + f"Mask: {mask.cpu().tolist()}" + ) + return self.model_runner_output diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index 3f86b5595cb9..fdedbfb86d03 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -38,6 +38,7 @@ ) from vllm.forward_context import BatchDescriptor, set_forward_context from vllm.logger import init_logger +from vllm.model_executor.layers.fused_moe.all2all_utils import get_ep_all2all_manager from vllm.model_executor.layers.mamba.ops.ssu_dispatch import ( initialize_mamba_ssu_backend, ) @@ -175,6 +176,12 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device): self.dp_size = self.parallel_config.data_parallel_size self.dp_rank = self.parallel_config.data_parallel_rank + # Detect EP all2all peer faults to prevent emitting corrupted output. + # Only meaningful for MoE + DP with an FT-capable all2all backend. + self.check_ep_fault = False + if self.dp_size > 1 and self.model_config.is_moe: + self.check_ep_fault = get_ep_all2all_manager().support_fault_tolerance + # Decode context parallelism. self.dcp_size = self.parallel_config.decode_context_parallel_size self.use_dcp = self.dcp_size > 1 @@ -1488,6 +1495,7 @@ def sample_tokens( num_sampled_tokens=num_sampled, main_stream=self.main_stream, copy_stream=self.output_copy_stream, + check_ep_fault=self.check_ep_fault, ) mm_inputs: tuple[list[torch.Tensor], torch.Tensor] | None = None diff --git a/vllm/v1/worker/gpu_worker.py b/vllm/v1/worker/gpu_worker.py index ca63e0a117c7..7c3d03366881 100644 --- a/vllm/v1/worker/gpu_worker.py +++ b/vllm/v1/worker/gpu_worker.py @@ -70,6 +70,7 @@ ModelRunnerOutput, ) from vllm.v1.utils import compute_iteration_details, report_usage_stats +from vllm.v1.worker.sentinel.gpu_worker_sentinel import WorkerSentinel from vllm.v1.worker.startup_plan import ( maybe_apply_startup_plan, maybe_save_startup_plan, @@ -146,7 +147,9 @@ def __init__( from vllm.distributed.elastic_ep.elastic_execute import ElasticEPScalingExecutor self.elastic_ep_executor = ElasticEPScalingExecutor(self) - + self.worker_sentinel: WorkerSentinel | None = None + if self.parallel_config.enable_fault_tolerance: + self.worker_sentinel = WorkerSentinel(worker=self) # Buffers saved before sleep self._sleep_saved_buffers: dict[str, torch.Tensor] = {} self._sleep_rebuild_draft_metadata_buffers = False @@ -414,6 +417,10 @@ def init_device(self): # If usage stat is enabled, collect relevant info. report_usage_stats(self.vllm_config) + def handle_ft_command(self, ft_request): + assert self.worker_sentinel is not None + return self.worker_sentinel.handle_command(ft_request) + # FIXME(youkaichao & ywang96): Use TorchDispatchMode instead of memory pool # to hijack tensor allocation. def load_model(self, *, load_dummy_weights: bool = False) -> None: diff --git a/vllm/v1/worker/sentinel/__init__.py b/vllm/v1/worker/sentinel/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/vllm/v1/worker/sentinel/gpu_worker_sentinel.py b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py new file mode 100644 index 000000000000..80050cf9d391 --- /dev/null +++ b/vllm/v1/worker/sentinel/gpu_worker_sentinel.py @@ -0,0 +1,89 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from typing import TYPE_CHECKING, cast + +import torch + +from vllm.config import set_current_vllm_config +from vllm.distributed import ( + get_dp_group, + stateless_destroy_torch_distributed_process_group, + stateless_init_torch_distributed_process_group, +) +from vllm.logger import init_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.serial_utils import run_method + +if TYPE_CHECKING: + from vllm.v1.worker.gpu.model_runner import GPUModelRunner as GPUModelRunnerV2 + from vllm.v1.worker.gpu_worker import Worker + +logger = init_logger(__name__) + +# All2all backends that support fault-tolerant timeout + rank masking, +# required for FT under DP+EP MoE deployments. +FT_BACKEND_SET = frozenset({"deepep_low_latency", "nixl_ep"}) + + +class WorkerSentinel: + """Holds FT state for a single worker (mask tensors, DP config). + + Methods are called via collective_rpc from EngineCoreSentinel. + """ + + def __init__(self, worker: "Worker"): + self.worker = worker + self.dp_rank = worker.parallel_config.data_parallel_rank + self.dp_size = worker.parallel_config.data_parallel_size + self.data_parallel_master_ip = worker.parallel_config.data_parallel_master_ip + all2all_backend = worker.parallel_config.all2all_backend + if all2all_backend not in FT_BACKEND_SET: + raise ValueError( + f"Fault tolerance requires an FT-capable all2all backend " + f"(one of {sorted(FT_BACKEND_SET)}), but got '{all2all_backend}'." + ) + + def handle_command(self, ft_request: FaultToleranceRequest): + """Dispatch an FT command by instruction name.""" + with set_current_vllm_config(self.worker.vllm_config): + return run_method(self, ft_request.instruction, (ft_request,), {}) + + def retry(self, ft_request: FaultToleranceRequest): + torch.accelerator.synchronize() + params = ft_request.params + self._clean_worker_state() + if self.dp_size > 1: + get_ep_all2all_manager().clean_buffers() + old_cpu_group = get_dp_group().cpu_group + stateless_destroy_torch_distributed_process_group(old_cpu_group) + world_size = self.worker.parallel_config.world_size + port = params["new_stateless_dp_group_ports"][self.worker.rank % world_size] + get_dp_group().cpu_group = stateless_init_torch_distributed_process_group( + self.data_parallel_master_ip, + port, + self.dp_rank, + self.dp_size, + backend="gloo", + ) + + def _clean_worker_state(self): + model_runner = self.worker.model_runner + model_runner.execute_model_state = None + if self.worker.use_v2_model_runner: + runner = cast("GPUModelRunnerV2", model_runner) + for req_id in list(runner.req_states.req_id_to_index): + runner._remove_request(req_id) + else: + model_runner.kv_connector_output = None + + input_batch = model_runner.input_batch + cached_req_ids = list(input_batch.req_id_to_index) + for req_id in cached_req_ids: + model_runner.requests.pop(req_id, None) + model_runner.num_prompt_logprobs.pop(req_id, None) + input_batch.remove_request(req_id) + + input_batch.condense() + input_batch.refresh_metadata() + input_batch.req_prompt_embeds.clear()