[6/N] (Elastic EP) Recover failed ranks (#15771)

This commit is contained in:
Xun Sun
2026-04-28 00:44:26 -07:00
committed by GitHub
parent b8a2dcd300
commit 9a53ab3d6d
8 changed files with 337 additions and 9 deletions
@@ -332,6 +332,7 @@ Please consult the documentation below and [server_args.py](https://github.com/s
| `--elastic-ep-backend` | Specify the collective communication backend for elastic EP. Currently supports 'mooncake'. | `none` | `none`, `mooncake` | | `--elastic-ep-backend` | Specify the collective communication backend for elastic EP. Currently supports 'mooncake'. | `none` | `none`, `mooncake` |
| `--enable-elastic-expert-backup` | Enable elastic EP backend to backup expert weights in DRAM feature. Currently supports 'mooncake'.| `False` | bool flag (set to enable) | | `--enable-elastic-expert-backup` | Enable elastic EP backend to backup expert weights in DRAM feature. Currently supports 'mooncake'.| `False` | bool flag (set to enable) |
| `--mooncake-ib-device` | The InfiniBand devices for Mooncake Backend transfer, accepts multiple comma-separated devices (e.g., --mooncake-ib-device mlx5_0,mlx5_1). Default is None, which triggers automatic device detection when Mooncake Backend is enabled. | `None` | Type: str | | `--mooncake-ib-device` | The InfiniBand devices for Mooncake Backend transfer, accepts multiple comma-separated devices (e.g., --mooncake-ib-device mlx5_0,mlx5_1). Default is None, which triggers automatic device detection when Mooncake Backend is enabled. | `None` | Type: str |
| `--elastic-ep-rejoin` | Indicates that this process is a relaunched elastic EP rank that should rejoin an existing process group during rank recovery. | `False` | bool flag (set to enable) |
## Mamba Cache ## Mamba Cache
| Argument | Description | Defaults | Options | | Argument | Description | Defaults | Options |
@@ -250,6 +250,7 @@ class GroupCoordinator:
use_message_queue_broadcaster: bool = False, use_message_queue_broadcaster: bool = False,
group_name: Optional[str] = None, group_name: Optional[str] = None,
gloo_timeout: timedelta = timedelta(seconds=120 * 60), gloo_timeout: timedelta = timedelta(seconds=120 * 60),
recovered_rank: bool = False,
): ):
# Set group info # Set group info
group_name = group_name or "anonymous" group_name = group_name or "anonymous"
@@ -288,13 +289,13 @@ class GroupCoordinator:
device_group = torch.distributed.new_group( device_group = torch.distributed.new_group(
ranks, ranks,
backend="mooncake", backend="mooncake",
pg_options=MooncakeBackendOptions(active_ranks), pg_options=MooncakeBackendOptions(active_ranks, recovered_rank),
timeout=subgroup_timeout, timeout=subgroup_timeout,
) )
cpu_group = torch.distributed.new_group( cpu_group = torch.distributed.new_group(
ranks, ranks,
backend="mooncake-cpu", backend="mooncake-cpu",
pg_options=MooncakeBackendOptions(active_ranks_cpu), pg_options=MooncakeBackendOptions(active_ranks_cpu, recovered_rank),
timeout=subgroup_timeout, timeout=subgroup_timeout,
) )
else: else:
@@ -443,7 +444,8 @@ class GroupCoordinator:
) )
self.mq_broadcaster: Optional[MessageQueue] = None self.mq_broadcaster: Optional[MessageQueue] = None
if use_message_queue_broadcaster and self.world_size > 1: if use_message_queue_broadcaster and self.world_size > 1 and not recovered_rank:
# Recovered ranks create their mq_broadcaster in elastic_ep.py
self.mq_broadcaster = MessageQueue.create_from_process_group( self.mq_broadcaster = MessageQueue.create_from_process_group(
self.cpu_group, 1 << 22, 6 self.cpu_group, 1 << 22, 6
) )
@@ -1385,7 +1387,7 @@ def get_world_group() -> GroupCoordinator:
def init_world_group( def init_world_group(
ranks: List[int], local_rank: int, backend: str ranks: List[int], local_rank: int, backend: str, recovered_rank: bool = False
) -> GroupCoordinator: ) -> GroupCoordinator:
return GroupCoordinator( return GroupCoordinator(
group_ranks=[ranks], group_ranks=[ranks],
@@ -1399,6 +1401,7 @@ def init_world_group(
use_xpu_communicator=False, use_xpu_communicator=False,
use_npu_communicator=False, use_npu_communicator=False,
group_name="world", group_name="world",
recovered_rank=recovered_rank,
) )
@@ -1412,6 +1415,7 @@ def init_model_parallel_group(
group_name: Optional[str] = None, group_name: Optional[str] = None,
use_mscclpp_allreduce: Optional[bool] = None, use_mscclpp_allreduce: Optional[bool] = None,
use_torch_symm_mem_allreduce: Optional[bool] = None, use_torch_symm_mem_allreduce: Optional[bool] = None,
recovered_rank: bool = False,
) -> GroupCoordinator: ) -> GroupCoordinator:
if use_custom_allreduce is None: if use_custom_allreduce is None:
use_custom_allreduce = _ENABLE_CUSTOM_ALL_REDUCE use_custom_allreduce = _ENABLE_CUSTOM_ALL_REDUCE
@@ -1436,6 +1440,7 @@ def init_model_parallel_group(
use_npu_communicator=True, use_npu_communicator=True,
use_message_queue_broadcaster=use_message_queue_broadcaster, use_message_queue_broadcaster=use_message_queue_broadcaster,
group_name=group_name, group_name=group_name,
recovered_rank=recovered_rank,
) )
@@ -1651,6 +1656,7 @@ def init_distributed_environment(
backend: str = "nccl", backend: str = "nccl",
timeout: Optional[int] = None, timeout: Optional[int] = None,
moe_a2a_backend: Optional[str] = None, moe_a2a_backend: Optional[str] = None,
recovered_rank: bool = False,
): ):
logger.debug( logger.debug(
"world_size=%d rank=%d local_rank=%d " "distributed_init_method=%s backend=%s", "world_size=%d rank=%d local_rank=%d " "distributed_init_method=%s backend=%s",
@@ -1681,8 +1687,17 @@ def init_distributed_environment(
assert isinstance(timeout, (int)), "timeout must be a number" assert isinstance(timeout, (int)), "timeout must be a number"
assert timeout > 0, "timeout must be positive" assert timeout > 0, "timeout must be positive"
timeout = timedelta(seconds=timeout) timeout = timedelta(seconds=timeout)
_MODEL_PARALLEL_GROUP_TIMEOUT = timeout _MODEL_PARALLEL_GROUP_TIMEOUT = timeout
pg_options = get_torch_distributed_pg_options()
if backend == "mooncake":
from mooncake.ep import MooncakeBackendOptions
# Setting "cuda" as device here is safe, as it is guarded under the mooncake case
active_ranks = torch.ones(world_size, dtype=torch.int32, device="cuda")
pg_options = MooncakeBackendOptions(active_ranks, recovered_rank)
else:
pg_options = get_torch_distributed_pg_options()
# this backend is used for WORLD # this backend is used for WORLD
torch.distributed.init_process_group( torch.distributed.init_process_group(
@@ -1711,7 +1726,9 @@ def init_distributed_environment(
global _WORLD global _WORLD
if _WORLD is None: if _WORLD is None:
ranks = list(range(torch.distributed.get_world_size())) ranks = list(range(torch.distributed.get_world_size()))
_WORLD = init_world_group(ranks, local_rank, backend) _WORLD = init_world_group(
ranks, local_rank, backend, recovered_rank=recovered_rank
)
else: else:
assert ( assert (
_WORLD.world_size == torch.distributed.get_world_size() _WORLD.world_size == torch.distributed.get_world_size()
@@ -1728,6 +1745,7 @@ def initialize_model_parallel(
backend: Optional[str] = None, backend: Optional[str] = None,
duplicate_tp_group: bool = False, duplicate_tp_group: bool = False,
enable_symm_mem: bool = False, enable_symm_mem: bool = False,
recovered_rank: bool = False,
) -> None: ) -> None:
""" """
Initialize model parallel groups. Initialize model parallel groups.
@@ -1806,6 +1824,7 @@ def initialize_model_parallel(
backend, backend,
use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(), use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(),
group_name="tp", group_name="tp",
recovered_rank=recovered_rank,
) )
if duplicate_tp_group: if duplicate_tp_group:
@@ -1819,6 +1838,7 @@ def initialize_model_parallel(
backend, backend,
use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(), use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(),
group_name="pdmux_prefill_tp", group_name="pdmux_prefill_tp",
recovered_rank=recovered_rank,
) )
if _TP.pynccl_comm: if _TP.pynccl_comm:
_TP.pynccl_comm.disabled = False _TP.pynccl_comm.disabled = False
@@ -1857,6 +1877,7 @@ def initialize_model_parallel(
backend, backend,
use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(), use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(),
group_name="attn_cp", group_name="attn_cp",
recovered_rank=recovered_rank,
) )
from sglang.srt.layers.sampler import SYNC_TOKEN_IDS_ACROSS_TP from sglang.srt.layers.sampler import SYNC_TOKEN_IDS_ACROSS_TP
@@ -1892,6 +1913,7 @@ def initialize_model_parallel(
use_torch_symm_mem_allreduce=False, use_torch_symm_mem_allreduce=False,
use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(), use_message_queue_broadcaster=envs.SGLANG_USE_MESSAGE_QUEUE_BROADCASTER.get(),
group_name="attention_tp", group_name="attention_tp",
recovered_rank=recovered_rank,
) )
moe_ep_size = expert_model_parallel_size moe_ep_size = expert_model_parallel_size
@@ -1922,6 +1944,7 @@ def initialize_model_parallel(
get_world_group().local_rank, get_world_group().local_rank,
backend, backend,
group_name="moe_dp", group_name="moe_dp",
recovered_rank=recovered_rank,
) )
global _MOE_EP global _MOE_EP
@@ -1948,6 +1971,7 @@ def initialize_model_parallel(
use_pynccl=False, use_pynccl=False,
use_custom_allreduce=False, use_custom_allreduce=False,
group_name="moe_ep", group_name="moe_ep",
recovered_rank=recovered_rank,
) )
global _MOE_TP global _MOE_TP
@@ -1975,6 +1999,7 @@ def initialize_model_parallel(
use_pynccl=False, use_pynccl=False,
use_custom_allreduce=False, use_custom_allreduce=False,
group_name="moe_tp", group_name="moe_tp",
recovered_rank=recovered_rank,
) )
# Build the pipeline model-parallel groups. # Build the pipeline model-parallel groups.
@@ -1994,6 +2019,7 @@ def initialize_model_parallel(
backend, backend,
use_custom_allreduce=False, use_custom_allreduce=False,
group_name="pp", group_name="pp",
recovered_rank=recovered_rank,
) )
+131 -1
View File
@@ -1,13 +1,18 @@
from __future__ import annotations from __future__ import annotations
import logging
import time
from dataclasses import dataclass from dataclasses import dataclass
from typing import Optional from typing import Iterator, List, Optional
import torch import torch
from sglang.srt.distributed import parallel_state
from sglang.srt.managers.schedule_batch import ServerArgs from sglang.srt.managers.schedule_batch import ServerArgs
from sglang.srt.utils import is_cpu, is_cuda from sglang.srt.utils import is_cpu, is_cuda
logger = logging.getLogger(__name__)
@dataclass @dataclass
class ElasticEPState: class ElasticEPState:
@@ -26,6 +31,12 @@ class ElasticEPState:
if self.active_ranks is not None: if self.active_ranks is not None:
self.last_active_ranks = self.active_ranks.clone() self.last_active_ranks = self.active_ranks.clone()
def reset(self):
if self.active_ranks is not None:
self.active_ranks.fill_(1)
self.snapshot_active_to_last()
self.sync_active_to_cpu()
class ElasticEPStateManager: class ElasticEPStateManager:
_instance: Optional[ElasticEPState] = None _instance: Optional[ElasticEPState] = None
@@ -41,6 +52,13 @@ class ElasticEPStateManager:
if server_args.elastic_ep_backend is not None: if server_args.elastic_ep_backend is not None:
cls._instance = cls._build_state(ep_size=None, device=None) cls._instance = cls._build_state(ep_size=None, device=None)
if server_args.elastic_ep_rejoin:
# Mask out peer ranks to perform cuda graph capture on its own
cls._instance.active_ranks.zero_()
cls._instance.active_ranks[torch.distributed.get_rank()] = 1
cls._instance.snapshot_active_to_last()
cls._instance.sync_active_to_cpu()
return cls._instance return cls._instance
@staticmethod @staticmethod
@@ -71,3 +89,115 @@ class ElasticEPStateManager:
dev = device if device is not None else cls._select_device() dev = device if device is not None else cls._select_device()
return torch.ones(size, dtype=torch.int32, device=dev) return torch.ones(size, dtype=torch.int32, device=dev)
# ---------------------------------------------------------------------------
# Helpers for elastic EP recovery
# ---------------------------------------------------------------------------
_PEER_STATE_POLL_INTERVAL_SEC = 0.01
def _get_process_group_backend(process_group, device: str):
return process_group._get_backend(torch.device(device))
def _iter_live_parallel_groups() -> Iterator[parallel_state.GroupCoordinator]:
groups = []
for group_ref in parallel_state._groups.values():
group = group_ref()
if group is not None:
groups.append(group)
for group in sorted(groups, key=lambda x: x.unique_name):
yield group
def _map_global_to_group_local_ranks(
group_ranks: List[int], global_ranks: List[int]
) -> List[int]:
rank_to_local = {rank: idx for idx, rank in enumerate(group_ranks)}
return [rank_to_local[rank] for rank in global_ranks if rank in rank_to_local]
def _wait_for_peer_state(mooncake_ep, backend, ranks: List[int]) -> None:
# Relaunched ranks become recoverable asynchronously, so we poll until the
# target backend reports all requested peers as ready.
while not all(mooncake_ep.get_peer_state(backend, ranks)):
time.sleep(_PEER_STATE_POLL_INTERVAL_SEC)
def _maybe_create_message_queue(group) -> None:
if not group.use_message_queue_broadcaster or group.world_size <= 1:
return
from sglang.srt.distributed.device_communicators.shm_broadcast import MessageQueue
group.mq_broadcaster = MessageQueue.create_from_process_group(
group.cpu_group, 1 << 22, 6
)
def _refresh_ep_members() -> None:
from sglang.srt.layers.moe.token_dispatcher.mooncake import EPBuffer
EPBuffer._buffer.update_ep_member()
def try_recover_ranks(global_ranks: List[int]) -> bool:
from mooncake import ep as mooncake_ep
world_backend = _get_process_group_backend(torch.distributed.group.WORLD, "cuda")
if not all(mooncake_ep.get_peer_state(world_backend, global_ranks)):
# The relaunched ranks have not finished initializing yet.
return False
# Recover the world backend first, then recover each derived process group
# using ranks mapped into that group's local rank space.
mooncake_ep.recover_ranks(world_backend, global_ranks)
for group in _iter_live_parallel_groups():
group_local_ranks = _map_global_to_group_local_ranks(group.ranks, global_ranks)
if not group_local_ranks:
continue
device_backend = _get_process_group_backend(group.device_group, "cuda")
_wait_for_peer_state(mooncake_ep, device_backend, group_local_ranks)
mooncake_ep.recover_ranks(device_backend, group_local_ranks)
cpu_backend = _get_process_group_backend(group.cpu_group, "cpu")
_wait_for_peer_state(mooncake_ep, cpu_backend, group_local_ranks)
mooncake_ep.recover_ranks(cpu_backend, group_local_ranks)
_maybe_create_message_queue(group)
_refresh_ep_members()
return True
def join_process_groups():
from mooncake import ep as mooncake_ep
def join_backend(label: str, backend) -> None:
logger.info("Recovered rank joining Mooncake backend %s", label)
mooncake_ep.join_group(backend)
join_backend(
"default_world",
_get_process_group_backend(torch.distributed.group.WORLD, "cuda"),
)
for group in _iter_live_parallel_groups():
if group.world_size <= 1:
continue
join_backend(
f"{group.unique_name}:device",
_get_process_group_backend(group.device_group, "cuda"),
)
join_backend(
f"{group.unique_name}:cpu",
_get_process_group_backend(group.cpu_group, "cpu"),
)
_maybe_create_message_queue(group)
_refresh_ep_members()
+3
View File
@@ -41,6 +41,9 @@ class EPLBManager:
def on_forward_pass_end(self): def on_forward_pass_end(self):
next(self._main_generator) next(self._main_generator)
def reset_generator(self):
self._main_generator = self._entrypoint()
# can be more complex if needed # can be more complex if needed
def _entrypoint(self): def _entrypoint(self):
while True: while True:
+46
View File
@@ -318,6 +318,52 @@ def set_global_expert_location_metadata(value):
_global_expert_location_metadata = value _global_expert_location_metadata = value
def broadcast_global_expert_location_metadata(
src_rank: int = 0, group: Optional[torch.distributed.ProcessGroup] = None
):
"""Broadcast the global ExpertLocationMetadata from src_rank to all ranks.
This is used in Elastic EP rank recovery to ensure that all ranks (including
newly recovered ones) share exactly the same expert location metadata.
Note: The caller must ensure src_rank is a healthy rank. In recovery scenarios,
this function is called after try_recover_ranks succeeds, at which point all
ranks (including src_rank=0) have recovered and are ready.
"""
metadata = get_global_expert_location_metadata()
assert metadata is not None
# Ensure device tensors are contiguous before broadcasting in-place
metadata.physical_to_logical_map = metadata.physical_to_logical_map.contiguous()
metadata.logical_to_all_physical_map = (
metadata.logical_to_all_physical_map.contiguous()
)
metadata.logical_to_all_physical_map_num_valid = (
metadata.logical_to_all_physical_map_num_valid.contiguous()
)
if metadata.logical_to_rank_dispatch_physical_map is not None:
metadata.logical_to_rank_dispatch_physical_map = (
metadata.logical_to_rank_dispatch_physical_map.contiguous()
)
device_tensors = [
metadata.physical_to_logical_map,
metadata.logical_to_all_physical_map,
metadata.logical_to_all_physical_map_num_valid,
]
if metadata.logical_to_rank_dispatch_physical_map is not None:
device_tensors.append(metadata.logical_to_rank_dispatch_physical_map)
for tensor in device_tensors:
torch.distributed.broadcast(tensor, src=src_rank, group=group)
# After broadcasting device tensors, refresh corresponding CPU copies
metadata.physical_to_logical_map_cpu = metadata.physical_to_logical_map.cpu()
metadata.logical_to_all_physical_map_cpu = (
metadata.logical_to_all_physical_map.cpu()
)
def _compute_logical_to_all_physical_map( def _compute_logical_to_all_physical_map(
server_args: ServerArgs, server_args: ServerArgs,
physical_to_logical_map: torch.Tensor, physical_to_logical_map: torch.Tensor,
@@ -361,7 +361,33 @@ class DataParallelController:
logger.debug("Worker port broadcast completed") logger.debug("Worker port broadcast completed")
return worker_ports return worker_ports
finally: finally:
rep_socket.close() if self.server_args.elastic_ep_backend is None:
rep_socket.close()
else:
threading.Thread(
target=self._reply_ports_as_server,
args=(rep_socket, worker_ports),
daemon=True,
).start()
def _reply_ports_as_server(self, rep_socket: zmq.Socket, worker_ports: List[int]):
"""
Runs as a background thread to broadcast worker ports for recovered EP ranks
"""
while True:
# Wait for client handshake
try:
client_rank = rep_socket.recv().decode()
except Exception:
logger.exception(
"Failed to recv/decode handshake in reply thread; continue"
)
continue
logger.debug(f"Received handshake from node {client_rank}")
# Send worker ports to client
rep_socket.send_pyobj(worker_ports)
logger.debug(f"Sent worker ports to node {client_rank}")
def _receive_ports_as_client(self, endpoint: str, node_rank: int) -> List[int]: def _receive_ports_as_client(self, endpoint: str, node_rank: int) -> List[int]:
"""Receive worker ports from the server node.""" """Receive worker ports from the server node."""
@@ -75,7 +75,11 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory, use_symmetric_memory,
) )
from sglang.srt.distributed.parallel_state import monkey_patch_vllm_parallel_state from sglang.srt.distributed.parallel_state import monkey_patch_vllm_parallel_state
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager from sglang.srt.elastic_ep.elastic_ep import (
ElasticEPStateManager,
join_process_groups,
try_recover_ranks,
)
from sglang.srt.elastic_ep.expert_backup_client import ExpertBackupClient from sglang.srt.elastic_ep.expert_backup_client import ExpertBackupClient
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.eplb.eplb_manager import EPLBManager from sglang.srt.eplb.eplb_manager import EPLBManager
@@ -87,6 +91,7 @@ from sglang.srt.eplb.expert_distribution import (
) )
from sglang.srt.eplb.expert_location import ( from sglang.srt.eplb.expert_location import (
ExpertLocationMetadata, ExpertLocationMetadata,
broadcast_global_expert_location_metadata,
compute_initial_expert_location_metadata, compute_initial_expert_location_metadata,
get_global_expert_location_metadata, get_global_expert_location_metadata,
set_global_expert_location_metadata, set_global_expert_location_metadata,
@@ -165,6 +170,7 @@ from sglang.srt.server_args import (
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils import ( from sglang.srt.utils import (
MultiprocessingSerializer, MultiprocessingSerializer,
broadcast_pyobj,
cpu_has_amx_support, cpu_has_amx_support,
dynamic_import, dynamic_import,
empty_context, empty_context,
@@ -493,6 +499,18 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.initialize(pre_model_load_memory) self.initialize(pre_model_load_memory)
self.check_quantized_moe_compatibility() self.check_quantized_moe_compatibility()
if (
self.server_args.elastic_ep_backend is not None
and self.server_args.elastic_ep_rejoin
):
join_process_groups()
broadcast_global_expert_location_metadata(
src_rank=self._get_healthy_expert_location_src_rank(
invoked_in_elastic_ep_rejoin_path=True
)
)
ElasticEPStateManager.instance().reset()
if self.is_multimodal: if self.is_multimodal:
sanity_check_mm_pad_shift_value(self.model_config.vocab_size) sanity_check_mm_pad_shift_value(self.model_config.vocab_size)
@@ -1081,6 +1099,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
distributed_init_method=dist_init_method, distributed_init_method=dist_init_method,
timeout=self.server_args.dist_timeout, timeout=self.server_args.dist_timeout,
moe_a2a_backend=self.server_args.moe_a2a_backend, moe_a2a_backend=self.server_args.moe_a2a_backend,
recovered_rank=self.server_args.elastic_ep_rejoin,
) )
initialize_model_parallel( initialize_model_parallel(
tensor_model_parallel_size=self.tp_size, tensor_model_parallel_size=self.tp_size,
@@ -1091,6 +1110,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
moe_data_model_parallel_size=self.moe_dp_size, moe_data_model_parallel_size=self.moe_dp_size,
duplicate_tp_group=self.server_args.enable_pdmux, duplicate_tp_group=self.server_args.enable_pdmux,
enable_symm_mem=self.server_args.enable_symm_mem, enable_symm_mem=self.server_args.enable_symm_mem,
recovered_rank=self.server_args.elastic_ep_rejoin,
) )
initialize_dp_attention( initialize_dp_attention(
server_args=self.server_args, server_args=self.server_args,
@@ -1464,6 +1484,66 @@ class ModelRunner(ModelRunnerKVCacheMixin):
weight_name_filter=weight_name_filter, weight_name_filter=weight_name_filter,
) )
def maybe_recover_ep_ranks(self):
# TODO(perf): `active_ranks.all()` on a CUDA tensor triggers host-device
# synchronization, and this function is on the forward-path.
# This check only runs when `--elastic-ep-backend` is enabled, so the
# synchronization overhead does not propagate to other configs.
# Leave for future optimization of the elastic EP path.
if self.tp_group.active_ranks.all() and self.tp_group.active_ranks_cpu.all():
return
tp_active_ranks = self.tp_group.active_ranks.detach().cpu().numpy()
tp_active_ranks_cpu = self.tp_group.active_ranks_cpu.detach().numpy()
tp_active_ranks &= tp_active_ranks_cpu
# NOTE: `ranks_to_recover` uses indices in `tp_group`. For the current
# Mooncake elastic EP implementation we assume `--pp-size=1`, so the
# tp-group index is the same as the global rank index.
ranks_to_recover = [
i for i in range(len(tp_active_ranks)) if not tp_active_ranks[i]
]
# try_recover_ranks polls peer state via Mooncake EP backend.
# Mooncake's internal semantics guarantee that all ranks observe
# consistent peer readiness state, so collective operations below
# are safe even though polling appears local.
if ranks_to_recover and try_recover_ranks(ranks_to_recover):
self.forward_pass_id = 0
self.eplb_manager.reset_generator()
broadcast_global_expert_location_metadata(
src_rank=self._get_healthy_expert_location_src_rank(
invoked_in_elastic_ep_rejoin_path=False
)
)
ElasticEPStateManager.instance().reset()
broadcast_pyobj(
[self.server_args.random_seed],
get_world_group().rank,
get_world_group().cpu_group,
src=get_world_group().ranks[0],
)
logger.info(f"recover ranks {ranks_to_recover} done")
def _get_healthy_expert_location_src_rank(
self, invoked_in_elastic_ep_rejoin_path: bool
) -> int:
world_group = get_world_group()
# NOTE: do not key off `self.server_args.elastic_ep_rejoin` here.
# A rank that was started as a rejoin rank may later act as a healthy
# rank in a subsequent recovery cycle.
local_rejoin_flag = bool(invoked_in_elastic_ep_rejoin_path)
gathered_rejoin_flags = world_group.all_gather_object(local_rejoin_flag)
for rank_in_group, is_rejoin_rank in enumerate(gathered_rejoin_flags):
if not is_rejoin_rank:
return world_group.ranks[rank_in_group]
raise RuntimeError(
"No healthy rank found for broadcasting expert location metadata. "
"All ranks are marked as elastic_ep_rejoin."
)
def update_weights_from_disk( def update_weights_from_disk(
self, self,
model_path: str, model_path: str,
@@ -3008,6 +3088,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.msprobe_debugger.stop() self.msprobe_debugger.stop()
self.msprobe_debugger.step() self.msprobe_debugger.step()
if self.server_args.elastic_ep_backend is not None:
self.maybe_recover_ep_ranks()
return output return output
def _forward_raw( def _forward_raw(
+13
View File
@@ -562,6 +562,7 @@ class ServerArgs:
elastic_ep_backend: Literal[None, "mooncake", "nixl"] = None elastic_ep_backend: Literal[None, "mooncake", "nixl"] = None
enable_elastic_expert_backup: bool = False enable_elastic_expert_backup: bool = False
mooncake_ib_device: Optional[str] = None mooncake_ib_device: Optional[str] = None
elastic_ep_rejoin: bool = False
# Mamba cache # Mamba cache
max_mamba_cache_size: Optional[int] = None max_mamba_cache_size: Optional[int] = None
@@ -3082,10 +3083,16 @@ class ServerArgs:
"elasticity_aware_hierarchical", "elasticity_aware_hierarchical",
], "Elastic EP requires eplb_algorithm to be set to 'auto' or 'elasticity_aware(_hierarchical)'." ], "Elastic EP requires eplb_algorithm to be set to 'auto' or 'elasticity_aware(_hierarchical)'."
assert self.pp_size == 1, "PP size should be set to 1 under elastic EP"
if self.elastic_ep_backend == "mooncake": if self.elastic_ep_backend == "mooncake":
self.mooncake_ib_device = self._validate_ib_devices( self.mooncake_ib_device = self._validate_ib_devices(
self.mooncake_ib_device self.mooncake_ib_device
) )
if self.elastic_ep_rejoin:
assert (
self.elastic_ep_backend is not None
), "Elastic EP rejoin requires elastic_ep_backend to be set."
def _handle_expert_distribution_metrics(self): def _handle_expert_distribution_metrics(self):
if self.enable_expert_distribution_metrics and ( if self.enable_expert_distribution_metrics and (
@@ -5581,6 +5588,12 @@ class ServerArgs:
"(e.g., --mooncake-ib-device mlx5_0,mlx5_1). " "(e.g., --mooncake-ib-device mlx5_0,mlx5_1). "
"Default is None, which triggers automatic device detection when Mooncake Backend is enabled.", "Default is None, which triggers automatic device detection when Mooncake Backend is enabled.",
) )
parser.add_argument(
"--elastic-ep-rejoin",
action="store_true",
default=ServerArgs.elastic_ep_rejoin,
help="Indicates that this process is a relaunched elastic EP rank that should rejoin an existing process group.",
)
# Mamba Cache # Mamba Cache
parser.add_argument( parser.add_argument(