[Feature] Detect all2all peer fault with fault tolerance backend and prevent corrupted output (#43637)

Signed-off-by: fangyuchu <fangyuchu@qq.com>
Co-authored-by: Tyler Michael Smith <tlrmchlsmth@gmail.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
This commit is contained in:
fangyuchu
2026-06-30 09:00:25 -07:00
committed by GitHub
co-authored by Tyler Michael Smith mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
parent e840f0d3f5
commit 245888ff77
5 changed files with 82 additions and 4 deletions
@@ -261,8 +261,13 @@ class DeepEPLLAll2AllManager(DeepEPAll2AllManagerBase):
All2All communication based on DeepEP Low-Latency kernels.
"""
_buffer: Any = None
_mask: torch.Tensor | None = None
_last_mask: torch.Tensor | None = None
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.
def _make_all2all_kwargs(
self,
@@ -304,6 +309,7 @@ class DeepEPLLAll2AllManager(DeepEPAll2AllManagerBase):
allow_nvlink_for_low_latency_mode=True,
allow_mnnvl=envs.VLLM_DEEPEP_LOW_LATENCY_USE_MNNVL,
explicitly_destroy=True,
enable_shrink=self.support_fault_tolerance,
)
return kwargs
@@ -319,12 +325,30 @@ class DeepEPLLAll2AllManager(DeepEPAll2AllManagerBase):
handle: deep_ep.Buffer = self.handle_cache.get_or_create(
buffer_kwargs, deep_ep.Buffer
)
DeepEPLLAll2AllManager._buffer = handle
return handle
# DeepEP LL uses RDMA so no SMs are used for communication
def max_sms_used(self) -> int | None:
return 0
def query_active_mask(self) -> torch.Tensor:
buf = DeepEPLLAll2AllManager._buffer
assert buf is not None
if DeepEPLLAll2AllManager._mask is None:
DeepEPLLAll2AllManager._mask = torch.zeros(
self.world_size, device="cuda", dtype=torch.int32
)
buf.low_latency_query_mask_buffer(DeepEPLLAll2AllManager._mask)
return DeepEPLLAll2AllManager._mask
def query_fault(self) -> torch.Tensor:
current = self.query_active_mask()
if DeepEPLLAll2AllManager._last_mask is None:
DeepEPLLAll2AllManager._last_mask = torch.zeros_like(current)
has_fault = (current != DeepEPLLAll2AllManager._last_mask).any()
return has_fault
@dataclass
class _NixlEPBufferState:
@@ -341,6 +365,8 @@ class NixlEPAll2AllManager(All2AllManagerBase):
_buffer: _NixlEPBufferState | None = None
_lock = threading.RLock()
_mask: torch.Tensor | None = None
_last_mask: torch.Tensor | None = None
def __init__(self, cpu_group, tcp_store_group=None):
if tcp_store_group is None:
@@ -350,6 +376,7 @@ class NixlEPAll2AllManager(All2AllManagerBase):
store=dist.PrefixStore("nixl_ep", cpu_group.get_group_store()),
)
super().__init__(cpu_group, tcp_store_group)
self.support_fault_tolerance = True
self.max_num_ep_ranks = envs.VLLM_NIXL_EP_MAX_NUM_RANKS
@@ -508,6 +535,25 @@ class NixlEPAll2AllManager(All2AllManagerBase):
def max_sms_used(self) -> int | None:
return 0
def query_active_mask(self) -> torch.Tensor:
state = NixlEPAll2AllManager._buffer
assert state is not None
if NixlEPAll2AllManager._mask is None:
NixlEPAll2AllManager._mask = torch.zeros(
self.max_num_ep_ranks, device="cuda", dtype=torch.int32
)
state.buffer.query_mask_buffer(NixlEPAll2AllManager._mask)
return NixlEPAll2AllManager._mask[: state.active_ep_size]
def query_fault(self) -> torch.Tensor:
current = self.query_active_mask()
last = NixlEPAll2AllManager._last_mask
if last is None or last.shape != current.shape:
NixlEPAll2AllManager._last_mask = torch.zeros_like(current)
last = NixlEPAll2AllManager._last_mask
has_fault = (current != last).any()
return has_fault
class FlashInferNVLinkTwoSidedManager(All2AllManagerBase):
"""
@@ -62,6 +62,8 @@ class All2AllManagerBase:
in_the_same_node_as(tcp_store_group, source_rank=0)
)
self.support_fault_tolerance = False
def get_handle(self, kwargs):
# get a handle for the all2all communication,
# based on the kwargs.
@@ -102,6 +104,13 @@ class All2AllManagerBase:
# - raise a clear error if extra_tensors is not supported.
raise NotImplementedError
def query_active_mask(self) -> torch.Tensor:
raise NotImplementedError
def query_fault(self) -> torch.Tensor:
"""Returns has_fault scalar."""
raise NotImplementedError
def set_num_sms(self, num_sms: int):
pass
@@ -56,7 +56,7 @@ if current_platform.is_cuda_alike():
)
def _get_ep_all2all_manager(eep_stage: bool = False) -> Any:
def get_ep_all2all_manager(eep_stage: bool = False) -> Any:
if eep_stage:
from vllm.distributed.elastic_ep.standby_state import get_standby_ep_group
@@ -146,7 +146,7 @@ def maybe_make_prepare_finalize(
"Detected DP deployment with no --enable-expert-parallel. "
"Falling back to AllGather+ReduceScatter dispatch/combine."
)
all2all_manager = _get_ep_all2all_manager(eep_stage)
all2all_manager = get_ep_all2all_manager(eep_stage)
return make_moe_prepare_and_finalize_naive_dp_ep(
is_sequence_parallel=moe.moe_parallel_config.is_sequence_parallel,
num_dispatchers=all2all_manager.world_size,
@@ -155,7 +155,7 @@ def maybe_make_prepare_finalize(
else:
return make_moe_prepare_and_finalize_no_dp_ep(use_monolithic)
all2all_manager = _get_ep_all2all_manager(eep_stage)
all2all_manager = get_ep_all2all_manager(eep_stage)
prepare_finalize: FusedMoEPrepareAndFinalize | None = None
+5 -1
View File
@@ -939,7 +939,11 @@ class WorkerProc:
converted to a FAILURE response.
"""
if isinstance(output, AsyncModelRunnerOutput):
output = output.get_output()
try:
output = output.get_output()
except Exception as e:
logger.exception("Error getting async model runner output")
output = e
if isinstance(output, Exception):
result = (WorkerProc.ResponseStatus.FAILURE, str(output))
+19
View File
@@ -58,6 +58,7 @@ from vllm.logger import init_logger
from vllm.lora.layers import LoRAMapping, LoRAMappingType
from vllm.model_executor.layers.attention import Attention, MLAAttention
from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
from vllm.model_executor.layers.fused_moe.all2all_utils import get_ep_all2all_manager
from vllm.model_executor.layers.fused_moe.routed_experts_capturer import (
RoutedExpertsCapturer,
)
@@ -249,6 +250,7 @@ class AsyncGPUModelRunnerOutput(AsyncModelRunnerOutput):
async_output_copy_stream: torch.cuda.Stream,
vocab_size: int,
routed_experts: RoutedExpertsTensors | None = None,
check_ep_fault: bool = False,
):
self._model_runner_output = model_runner_output
self._invalid_req_indices = invalid_req_indices
@@ -262,6 +264,7 @@ class AsyncGPUModelRunnerOutput(AsyncModelRunnerOutput):
self.vocab_size = vocab_size
self._logprobs_tensors = logprobs_tensors
self._routed_experts = routed_experts
self._has_fault: torch.Tensor | None = None
# Initiate the copy on a separate stream, but do not synchronize it.
default_stream = torch.cuda.current_stream()
@@ -280,6 +283,9 @@ class AsyncGPUModelRunnerOutput(AsyncModelRunnerOutput):
if self._routed_experts is not None
else None
)
if check_ep_fault:
has_fault = get_ep_all2all_manager().query_fault()
self._has_fault = has_fault.to("cpu", non_blocking=True)
self.async_copy_ready_event.record()
def get_output(self) -> ModelRunnerOutput:
@@ -316,6 +322,14 @@ class AsyncGPUModelRunnerOutput(AsyncModelRunnerOutput):
output.routed_experts = self._routed_experts_cpu.tolists()
del self._routed_experts
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 output
@@ -445,6 +459,10 @@ class GPUModelRunner(
self.device = device
self.dtype = self.model_config.dtype
self.check_ep_fault = False
if parallel_config.data_parallel_size > 1 and self.model_config.is_moe:
self.check_ep_fault = get_ep_all2all_manager().support_fault_tolerance
self.kv_cache_dtype = kv_cache_dtype_str_to_dtype(
cache_config.cache_dtype, self.model_config
)
@@ -4695,6 +4713,7 @@ class GPUModelRunner(
async_output_copy_stream=self._get_or_create_async_output_copy_stream(),
vocab_size=self.input_batch.vocab_size,
routed_experts=routed_experts_snapshot,
check_ep_fault=self.check_ep_fault,
)
with record_function_or_nullcontext(
"gpu_model_runner: set_async_sampled_token_ids"