diff --git a/python/sglang/srt/eplb/eplb_manager.py b/python/sglang/srt/eplb/eplb_manager.py index 360ac93b7..4dc97b23e 100644 --- a/python/sglang/srt/eplb/eplb_manager.py +++ b/python/sglang/srt/eplb/eplb_manager.py @@ -5,8 +5,10 @@ import time from typing import TYPE_CHECKING, Any, Callable, List import torch.cuda +import torch.distributed as dist from torch import nn +from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager from sglang.srt.environ import envs from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ( @@ -81,6 +83,11 @@ class EPLBManager: self._rebalance_disabled_logged = False self.reset_generator() + def enable_rebalance(self): + self._rebalance_disabled_reason = None + self._rebalance_disabled_logged = False + self.reset_generator() + # can be more complex if needed def _entrypoint(self): while True: @@ -99,6 +106,15 @@ class EPLBManager: self._rebalance_disabled_logged = True return + elastic_state = ElasticEPStateManager.instance() + is_post_scale_rebalance = elastic_state is not None and elastic_state.has_scaled + # A failed later scale leaves the previously committed world serving. + if is_post_scale_rebalance and ( + elastic_state.pending_ep_size is not None + or elastic_state.scale_phase not in ("serving_expanded", "failed") + ): + return + logger.info("[EPLBManager] rebalance start") enable_timing = self._rebalance_layers_per_chunk is None @@ -119,8 +135,9 @@ class EPLBManager: if not self._check_rebalance_needed(average_utilization_rate_over_window): return - expert_location_metadata = ExpertLocationMetadata.init_by_eplb( - self._server_args, self._model_config, logical_count + expert_location_metadata = self._compute_expert_location_metadata( + logical_count, + broadcast_over_world=is_post_scale_rebalance, ) from sglang.srt.model_executor.model_runner_components.moe_ep_setup import ( @@ -144,7 +161,12 @@ class EPLBManager: new_expert_location_metadata=expert_location_metadata, update_layer_ids=chunk_layer_ids, nnodes=self._server_args.nnodes, - tp_rank=self._ps.tp_rank, + tp_rank=( + self._elastic_global_rank() + if is_post_scale_rebalance + else self._ps.tp_rank + ), + use_flat_topology=is_post_scale_rebalance, expert_backup_client=self._get_expert_backup_client(), update_weights_from_disk_callable=self._get_weight_updater().update_weights_from_disk, ep_dispatch_algorithm=self._server_args.ep_dispatch_algorithm, @@ -152,6 +174,10 @@ class EPLBManager: model_config=self._model_config ), ) + if is_post_scale_rebalance: + # P2P waits only synchronize participating peers. Ranks without + # moves must also install this chunk before NIXL resumes. + dist.barrier() self._log_rebalance_layout_after_update(update_layer_ids=all_update_layer_ids) @@ -162,6 +188,47 @@ class EPLBManager: msg += f" time={time_end - time_start:.3f}s" logger.info(msg) + def _compute_expert_location_metadata( + self, logical_count, *, broadcast_over_world: bool + ) -> ExpertLocationMetadata: + if not broadcast_over_world: + return ExpertLocationMetadata.init_by_eplb( + self._server_args, + self._model_config, + logical_count, + ) + + current_metadata = get_global_expert_location_metadata() + assert current_metadata is not None + # One owner prevents process-local launch topology from influencing + # the mapping chosen for the expanded world. + if dist.get_rank() == 0: + computed_metadata = ExpertLocationMetadata.init_by_eplb( + self._server_args, + self._model_config, + logical_count, + # Arbitrary append topologies may not preserve node divisibility. + use_flat_topology=True, + ) + physical_to_logical_map = ( + computed_metadata.physical_to_logical_map.contiguous() + ) + else: + physical_to_logical_map = torch.empty_like( + current_metadata.physical_to_logical_map + ) + + dist.broadcast(physical_to_logical_map, src=0) + return ExpertLocationMetadata.init_by_mapping( + self._server_args, + self._model_config, + physical_to_logical_map, + moe_ep_rank=self._elastic_global_rank(), + ) + + def _elastic_global_rank(self) -> int: + return self._ps.tp_rank + self._server_args.ep_join_rank_offset + def _check_rebalance_needed(self, average_utilization_rate_over_window): if average_utilization_rate_over_window is None: return True @@ -240,6 +307,7 @@ def update_expert_location_with_recovery( update_layer_ids: List[int], nnodes: int, tp_rank: int, + use_flat_topology: bool = False, expert_backup_client, update_weights_from_disk_callable, ep_dispatch_algorithm: str, @@ -251,6 +319,7 @@ def update_expert_location_with_recovery( update_layer_ids=update_layer_ids, nnodes=nnodes, rank=tp_rank, + use_flat_topology=use_flat_topology, ) if len(p2p_missing_logical_experts) > 0: diff --git a/python/sglang/srt/eplb/expert_location.py b/python/sglang/srt/eplb/expert_location.py index 04ce28c1d..7602fdb1f 100644 --- a/python/sglang/srt/eplb/expert_location.py +++ b/python/sglang/srt/eplb/expert_location.py @@ -173,7 +173,11 @@ class ExpertLocationMetadata: @staticmethod def init_by_eplb( - server_args: ServerArgs, model_config: ModelConfig, logical_count: torch.Tensor + server_args: ServerArgs, + model_config: ModelConfig, + logical_count: torch.Tensor, + *, + use_flat_topology: bool = False, ): if not isinstance(logical_count, torch.Tensor): logical_count = torch.tensor(logical_count) @@ -189,7 +193,7 @@ class ExpertLocationMetadata: model_config_for_expert_location = common["model_config_for_expert_location"] num_physical_experts = common["num_physical_experts"] num_groups = model_config_for_expert_location.num_groups - num_nodes = server_args.nnodes + num_nodes = 1 if use_flat_topology else server_args.nnodes from sglang.srt.eplb import eplb_algorithms diff --git a/python/sglang/srt/eplb/expert_location_updater.py b/python/sglang/srt/eplb/expert_location_updater.py index 7873223f0..46bd5cdf2 100644 --- a/python/sglang/srt/eplb/expert_location_updater.py +++ b/python/sglang/srt/eplb/expert_location_updater.py @@ -46,6 +46,7 @@ class ExpertLocationUpdater: update_layer_ids: List[int], nnodes: int, rank: int, + use_flat_topology: bool = False, ): """ Update experts' physical location after EPLB. @@ -59,13 +60,14 @@ class ExpertLocationUpdater: old_expert_location_metadata = get_global_expert_location_metadata() assert old_expert_location_metadata is not None + topology_num_nodes = 1 if use_flat_topology else nnodes missing_logical_experts_by_layers = _update_expert_weights( routed_experts_weights_of_layer=routed_experts_weights_of_layer, old_expert_location_metadata=old_expert_location_metadata, new_expert_location_metadata=new_expert_location_metadata, update_layer_ids=update_layer_ids, - nnodes=nnodes, + nnodes=topology_num_nodes, rank=rank, ) old_expert_location_metadata.update( diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index b212c9441..862ba08f5 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -4453,6 +4453,8 @@ class Scheduler( pending_ep_size=ElasticEPStateManager.get_pending_ep_size(), scale_phase=ElasticEPStateManager.get_scale_phase(), ) + if (eplb_manager := self.tp_worker.model_runner.eplb_manager) is not None: + eplb_manager.disable_rebalance("elastic EP scale-up is pending") logger.debug( "[Elastic EP][scale] scale requested: target_ep_size=%d; " "waiting for a joining cohort", diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index c919cd394..856aec750 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -444,7 +444,8 @@ class ModelRunner: ) if self.eplb_manager is not None: self.eplb_manager.disable_rebalance( - "EPLB rebalance is disabled after elastic EP scale-up" + "EPLB rebalance is disabled while elastic EP scale-up " + "is being finalized" ) state = ElasticEPStateManager.instance() @@ -460,6 +461,7 @@ class ModelRunner: ) if state is not None: state.scale_phase = "serving_expanded" + self._rearm_eplb_after_elastic_scale() def init_msprobe(self): self.msprobe_debugger = misc_utils.create_msprobe_debugger(self.server_args) @@ -1703,6 +1705,26 @@ class ModelRunner: def _elastic_global_rank(self) -> int: return self.ps.tp_rank + self.server_args.ep_join_rank_offset + def _rearm_eplb_after_elastic_scale(self) -> None: + if self.eplb_manager is None: + return + recorder = get_global_expert_distribution_recorder() + if not recorder.recording: + recorder.start_record() + self.eplb_manager.enable_rebalance() + + def _reset_eplb_after_elastic_scale_failure(self) -> None: + if self.eplb_manager is None: + return + set_global_expert_distribution_recorder( + ExpertDistributionRecorder.init_new( + self.server_args, + get_global_expert_location_metadata(), + rank=self._elastic_global_rank(), + ) + ) + self._rearm_eplb_after_elastic_scale() + def _report_elastic_scale_failure(self, error: str, effective_size: int) -> None: if self.ps.tp_rank != 0 or self.server_args.is_ep_scale_joiner: return @@ -1769,7 +1791,8 @@ class ModelRunner: if self.eplb_manager is not None: self.eplb_manager.disable_rebalance( - "EPLB rebalance is disabled after elastic EP scale-up" + "EPLB rebalance is disabled while elastic EP scale-up " + "is being finalized" ) from sglang.srt.layers.dp_attention import update_dp_attention_post_scale @@ -1786,6 +1809,7 @@ class ModelRunner: log_tag="JOINER" if self.server_args.is_ep_scale_joiner else "PRIMARY", ) ElasticEPStateManager.commit_scale() + self._rearm_eplb_after_elastic_scale() if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner: from sglang.srt.managers.io_struct import ElasticScaleUpdateReq @@ -1844,6 +1868,7 @@ class ModelRunner: if timeout.item(): error = f"Timed out waiting for ranks to join target EP size {pending_size}" ElasticEPStateManager.fail_scale(error) + self._reset_eplb_after_elastic_scale_failure() self._report_elastic_scale_failure(error, effective_size) if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner: logger.error("[Elastic EP] %s", error) @@ -1859,6 +1884,7 @@ class ModelRunner: f"joining cohort target {cohort_target}" ) ElasticEPStateManager.fail_scale(error) + self._reset_eplb_after_elastic_scale_failure() self._report_elastic_scale_failure(error, effective_size) if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner: logger.error("[Elastic EP] %s", error) diff --git a/test/manual/ep/test_elastic_scale.py b/test/manual/ep/test_elastic_scale.py index 237526dd0..2bbed7ceb 100644 --- a/test/manual/ep/test_elastic_scale.py +++ b/test/manual/ep/test_elastic_scale.py @@ -407,6 +407,7 @@ class _ElasticScaleUpEndToEndBase(CustomTestCase): self._generate_logprob_ok("post-scale") self._run_post_scale_gsm8k() + self._generate_logprob_ok("after post-scale workload") @unittest.skipUnless( @@ -458,6 +459,7 @@ class TestElasticScaleUp4To5To6(_ElasticScaleUpEndToEndBase): self._generate_ok("after second scale") self._generate_logprob_ok("after second scale") self._run_post_scale_gsm8k() + self._generate_logprob_ok("after post-scale workload") @unittest.skipUnless(