[Elastic EP] Fix recovery lifecycle and add manual coverage (#31744)
Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
@@ -11,9 +11,10 @@ from sglang.srt.distributed import get_world_group, parallel_state
|
||||
from sglang.srt.distributed.utils import get_global_tcp_store
|
||||
from sglang.srt.eplb.expert_location import broadcast_global_expert_location_metadata
|
||||
from sglang.srt.managers.schedule_batch import ServerArgs
|
||||
from sglang.srt.utils import broadcast_pyobj, is_cpu, is_cuda
|
||||
from sglang.srt.utils import is_cpu, is_cuda
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
from sglang.srt.eplb.eplb_manager import EPLBManager
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -462,7 +463,8 @@ def maybe_recover_ep_ranks(
|
||||
*,
|
||||
tp_group: parallel_state.GroupCoordinator,
|
||||
eplb_manager: EPLBManager,
|
||||
random_seed: int,
|
||||
model_config: ModelConfig,
|
||||
moe_ep_rank: int,
|
||||
) -> bool:
|
||||
# TODO(perf): `active_ranks.all()` on a CUDA tensor triggers host-device
|
||||
# synchronization, and this function is on the forward-path.
|
||||
@@ -489,17 +491,13 @@ def maybe_recover_ep_ranks(
|
||||
if ranks_to_recover and try_recover_ranks(ranks_to_recover):
|
||||
eplb_manager.reset_generator()
|
||||
broadcast_global_expert_location_metadata(
|
||||
model_config=model_config,
|
||||
moe_ep_rank=moe_ep_rank,
|
||||
src_rank=get_healthy_expert_location_src_rank(
|
||||
invoked_in_elastic_ep_rejoin_path=False
|
||||
)
|
||||
),
|
||||
)
|
||||
ElasticEPStateManager.instance().reset()
|
||||
broadcast_pyobj(
|
||||
[random_seed],
|
||||
parallel_state.get_world_group().rank,
|
||||
parallel_state.get_world_group().cpu_group,
|
||||
src=parallel_state.get_world_group().ranks[0],
|
||||
)
|
||||
logger.info(f"recover ranks {ranks_to_recover} done")
|
||||
return True
|
||||
|
||||
|
||||
@@ -885,6 +885,12 @@ class Scheduler(
|
||||
if model_runner.token_to_kv_pool.post_capture_active:
|
||||
model_runner.post_capture_resize_kv_pool()
|
||||
|
||||
if (
|
||||
self.server_args.elastic_ep_backend is not None
|
||||
and self.server_args.ep_join_mode == "recover"
|
||||
):
|
||||
model_runner.post_capture_elastic_ep_recover()
|
||||
|
||||
# Dispatch the model worker
|
||||
if self.spec_algorithm.is_none():
|
||||
self.model_worker = self.tp_worker
|
||||
|
||||
@@ -342,8 +342,8 @@ class TpModelWorker(BaseTpWorker):
|
||||
self.world_group = get_world_group()
|
||||
|
||||
# Sync random seed across TP workers.
|
||||
# Scale joiners cannot enter the launch-time WORLD broadcast.
|
||||
if server_args.is_ep_scale_joiner:
|
||||
# Elastic joiners cannot enter the launch-time WORLD broadcast.
|
||||
if server_args.is_ep_joiner:
|
||||
self.random_seed = server_args.random_seed
|
||||
else:
|
||||
self.random_seed = broadcast_pyobj(
|
||||
|
||||
@@ -399,39 +399,27 @@ class ModelRunner:
|
||||
def _initialize_elastic_ep_joiner(self) -> None:
|
||||
if not (
|
||||
self.server_args.elastic_ep_backend is not None
|
||||
and self.server_args.is_ep_joiner
|
||||
and self.server_args.is_ep_scale_joiner
|
||||
):
|
||||
return
|
||||
|
||||
is_scale_join = self.server_args.ep_join_mode == "scale"
|
||||
if is_scale_join:
|
||||
join_effective_ep_size = (
|
||||
self.server_args.ep_join_rank_offset + self.ps.tp_size
|
||||
join_effective_ep_size = self.server_args.ep_join_rank_offset + self.ps.tp_size
|
||||
dist.barrier(group=self.tp_group.cpu_group)
|
||||
if self.ps.tp_rank == 0:
|
||||
register_scale_cohort(
|
||||
self.server_args.ep_join_rank_offset,
|
||||
join_effective_ep_size,
|
||||
)
|
||||
dist.barrier(group=self.tp_group.cpu_group)
|
||||
if self.ps.tp_rank == 0:
|
||||
register_scale_cohort(
|
||||
self.server_args.ep_join_rank_offset,
|
||||
join_effective_ep_size,
|
||||
)
|
||||
join_scale_process_group()
|
||||
self.server_args.override(
|
||||
"elastic_ep.scale_join", ep_size=join_effective_ep_size
|
||||
)
|
||||
else:
|
||||
join_process_groups()
|
||||
join_scale_process_group()
|
||||
self.server_args.override(
|
||||
"elastic_ep.scale_join", ep_size=join_effective_ep_size
|
||||
)
|
||||
|
||||
global_ep_rank = self.ps.tp_rank + self.server_args.ep_join_rank_offset
|
||||
broadcast_global_expert_location_metadata(
|
||||
model_config=self.model_config,
|
||||
moe_ep_rank=global_ep_rank,
|
||||
src_rank=(
|
||||
0
|
||||
if is_scale_join
|
||||
else get_healthy_expert_location_src_rank(
|
||||
invoked_in_elastic_ep_rejoin_path=True
|
||||
)
|
||||
),
|
||||
src_rank=0,
|
||||
)
|
||||
set_global_expert_distribution_recorder(
|
||||
ExpertDistributionRecorder.init_new(
|
||||
@@ -441,10 +429,6 @@ class ModelRunner:
|
||||
)
|
||||
)
|
||||
|
||||
if not is_scale_join:
|
||||
ElasticEPStateManager.instance().reset()
|
||||
return
|
||||
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
enable_joiner_all_gather,
|
||||
update_dp_attention_post_scale,
|
||||
@@ -839,6 +823,27 @@ class ModelRunner:
|
||||
resize.capped_max_running_requests
|
||||
)
|
||||
|
||||
def post_capture_elastic_ep_recover(self):
|
||||
join_process_groups()
|
||||
|
||||
global_ep_rank = self.ps.tp_rank + self.server_args.ep_join_rank_offset
|
||||
broadcast_global_expert_location_metadata(
|
||||
model_config=self.model_config,
|
||||
moe_ep_rank=global_ep_rank,
|
||||
src_rank=get_healthy_expert_location_src_rank(
|
||||
invoked_in_elastic_ep_rejoin_path=True
|
||||
),
|
||||
)
|
||||
set_global_expert_distribution_recorder(
|
||||
ExpertDistributionRecorder.init_new(
|
||||
self.server_args,
|
||||
get_global_expert_location_metadata(),
|
||||
rank=global_ep_rank,
|
||||
)
|
||||
)
|
||||
|
||||
ElasticEPStateManager.instance().reset()
|
||||
|
||||
def init_attention_backends(self):
|
||||
"""Initialize attention backends only (no cuda graph capture)."""
|
||||
# Must be called BEFORE init_decode_cuda_graph() so CUDA graph capture
|
||||
@@ -1089,7 +1094,7 @@ class ModelRunner:
|
||||
dist_barrier_after_load(
|
||||
elastic_ep_backend=self.server_args.elastic_ep_backend,
|
||||
tp_rank=self.ps.tp_rank,
|
||||
is_ep_scale_joiner=self.server_args.is_ep_scale_joiner,
|
||||
is_ep_joiner=self.server_args.is_ep_joiner,
|
||||
)
|
||||
|
||||
def maybe_init_dwdp(self):
|
||||
@@ -1822,7 +1827,8 @@ class ModelRunner:
|
||||
recovered = maybe_recover_ep_ranks(
|
||||
tp_group=self.tp_group,
|
||||
eplb_manager=self.eplb_manager,
|
||||
random_seed=self.server_args.random_seed,
|
||||
model_config=self.model_config,
|
||||
moe_ep_rank=self._elastic_global_rank(),
|
||||
)
|
||||
if recovered:
|
||||
self.forward_pass_id = 0
|
||||
|
||||
@@ -301,11 +301,11 @@ def dist_barrier_after_load(
|
||||
*,
|
||||
elastic_ep_backend: Optional[str],
|
||||
tp_rank: int,
|
||||
is_ep_scale_joiner: bool = False,
|
||||
is_ep_joiner: bool = False,
|
||||
) -> None:
|
||||
if elastic_ep_backend == "mooncake":
|
||||
# Mooncake does not support `monitored_barrier`
|
||||
if not is_ep_scale_joiner:
|
||||
if not is_ep_joiner:
|
||||
dist.barrier(group=get_tp_group().cpu_group)
|
||||
else:
|
||||
# Handle the case where some ranks do not finish loading.
|
||||
|
||||
@@ -8984,7 +8984,7 @@ class PortArgs:
|
||||
# (no availability-based search). If incrementing would
|
||||
# overflow the valid TCP range, decrement instead.
|
||||
NUM_DERIVED_PORTS = 5
|
||||
if server_args.is_ep_scale_joiner:
|
||||
if server_args.is_ep_joiner:
|
||||
port_base = server_args.port + ZMQ_TCP_PORT_DELTA
|
||||
if port_base + NUM_DERIVED_PORTS > 65535:
|
||||
port_base = server_args.port - ZMQ_TCP_PORT_DELTA
|
||||
@@ -9004,7 +9004,7 @@ class PortArgs:
|
||||
assert worker_ports is not None
|
||||
scheduler_input_port = worker_ports[dp_rank]
|
||||
|
||||
is_joiner = server_args.is_ep_scale_joiner
|
||||
is_joiner = server_args.is_ep_joiner
|
||||
# Under SGLANG_DISTRIBUTED_INIT_METHOD_OVERRIDE, SGLang never binds
|
||||
# dist_init_port / nccl_port (rendezvous uses the externally-managed
|
||||
# store; see distributed/bootstrap.py:_resolve_dist_init_method), so
|
||||
|
||||
Reference in New Issue
Block a user