[2/N] elastic-ep: Enable EPLB after scale-up (#30553)
This commit is contained in:
@@ -5,8 +5,10 @@ import time
|
|||||||
from typing import TYPE_CHECKING, Any, Callable, List
|
from typing import TYPE_CHECKING, Any, Callable, List
|
||||||
|
|
||||||
import torch.cuda
|
import torch.cuda
|
||||||
|
import torch.distributed as dist
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
|
||||||
|
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
||||||
from sglang.srt.eplb.expert_location import (
|
from sglang.srt.eplb.expert_location import (
|
||||||
@@ -81,6 +83,11 @@ class EPLBManager:
|
|||||||
self._rebalance_disabled_logged = False
|
self._rebalance_disabled_logged = False
|
||||||
self.reset_generator()
|
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
|
# can be more complex if needed
|
||||||
def _entrypoint(self):
|
def _entrypoint(self):
|
||||||
while True:
|
while True:
|
||||||
@@ -99,6 +106,15 @@ class EPLBManager:
|
|||||||
self._rebalance_disabled_logged = True
|
self._rebalance_disabled_logged = True
|
||||||
return
|
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")
|
logger.info("[EPLBManager] rebalance start")
|
||||||
|
|
||||||
enable_timing = self._rebalance_layers_per_chunk is None
|
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):
|
if not self._check_rebalance_needed(average_utilization_rate_over_window):
|
||||||
return
|
return
|
||||||
|
|
||||||
expert_location_metadata = ExpertLocationMetadata.init_by_eplb(
|
expert_location_metadata = self._compute_expert_location_metadata(
|
||||||
self._server_args, self._model_config, logical_count
|
logical_count,
|
||||||
|
broadcast_over_world=is_post_scale_rebalance,
|
||||||
)
|
)
|
||||||
|
|
||||||
from sglang.srt.model_executor.model_runner_components.moe_ep_setup import (
|
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,
|
new_expert_location_metadata=expert_location_metadata,
|
||||||
update_layer_ids=chunk_layer_ids,
|
update_layer_ids=chunk_layer_ids,
|
||||||
nnodes=self._server_args.nnodes,
|
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(),
|
expert_backup_client=self._get_expert_backup_client(),
|
||||||
update_weights_from_disk_callable=self._get_weight_updater().update_weights_from_disk,
|
update_weights_from_disk_callable=self._get_weight_updater().update_weights_from_disk,
|
||||||
ep_dispatch_algorithm=self._server_args.ep_dispatch_algorithm,
|
ep_dispatch_algorithm=self._server_args.ep_dispatch_algorithm,
|
||||||
@@ -152,6 +174,10 @@ class EPLBManager:
|
|||||||
model_config=self._model_config
|
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)
|
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"
|
msg += f" time={time_end - time_start:.3f}s"
|
||||||
logger.info(msg)
|
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):
|
def _check_rebalance_needed(self, average_utilization_rate_over_window):
|
||||||
if average_utilization_rate_over_window is None:
|
if average_utilization_rate_over_window is None:
|
||||||
return True
|
return True
|
||||||
@@ -240,6 +307,7 @@ def update_expert_location_with_recovery(
|
|||||||
update_layer_ids: List[int],
|
update_layer_ids: List[int],
|
||||||
nnodes: int,
|
nnodes: int,
|
||||||
tp_rank: int,
|
tp_rank: int,
|
||||||
|
use_flat_topology: bool = False,
|
||||||
expert_backup_client,
|
expert_backup_client,
|
||||||
update_weights_from_disk_callable,
|
update_weights_from_disk_callable,
|
||||||
ep_dispatch_algorithm: str,
|
ep_dispatch_algorithm: str,
|
||||||
@@ -251,6 +319,7 @@ def update_expert_location_with_recovery(
|
|||||||
update_layer_ids=update_layer_ids,
|
update_layer_ids=update_layer_ids,
|
||||||
nnodes=nnodes,
|
nnodes=nnodes,
|
||||||
rank=tp_rank,
|
rank=tp_rank,
|
||||||
|
use_flat_topology=use_flat_topology,
|
||||||
)
|
)
|
||||||
|
|
||||||
if len(p2p_missing_logical_experts) > 0:
|
if len(p2p_missing_logical_experts) > 0:
|
||||||
|
|||||||
@@ -173,7 +173,11 @@ class ExpertLocationMetadata:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def init_by_eplb(
|
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):
|
if not isinstance(logical_count, torch.Tensor):
|
||||||
logical_count = torch.tensor(logical_count)
|
logical_count = torch.tensor(logical_count)
|
||||||
@@ -189,7 +193,7 @@ class ExpertLocationMetadata:
|
|||||||
model_config_for_expert_location = common["model_config_for_expert_location"]
|
model_config_for_expert_location = common["model_config_for_expert_location"]
|
||||||
num_physical_experts = common["num_physical_experts"]
|
num_physical_experts = common["num_physical_experts"]
|
||||||
num_groups = model_config_for_expert_location.num_groups
|
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
|
from sglang.srt.eplb import eplb_algorithms
|
||||||
|
|
||||||
|
|||||||
@@ -46,6 +46,7 @@ class ExpertLocationUpdater:
|
|||||||
update_layer_ids: List[int],
|
update_layer_ids: List[int],
|
||||||
nnodes: int,
|
nnodes: int,
|
||||||
rank: int,
|
rank: int,
|
||||||
|
use_flat_topology: bool = False,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Update experts' physical location after EPLB.
|
Update experts' physical location after EPLB.
|
||||||
@@ -59,13 +60,14 @@ class ExpertLocationUpdater:
|
|||||||
|
|
||||||
old_expert_location_metadata = get_global_expert_location_metadata()
|
old_expert_location_metadata = get_global_expert_location_metadata()
|
||||||
assert old_expert_location_metadata is not None
|
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(
|
missing_logical_experts_by_layers = _update_expert_weights(
|
||||||
routed_experts_weights_of_layer=routed_experts_weights_of_layer,
|
routed_experts_weights_of_layer=routed_experts_weights_of_layer,
|
||||||
old_expert_location_metadata=old_expert_location_metadata,
|
old_expert_location_metadata=old_expert_location_metadata,
|
||||||
new_expert_location_metadata=new_expert_location_metadata,
|
new_expert_location_metadata=new_expert_location_metadata,
|
||||||
update_layer_ids=update_layer_ids,
|
update_layer_ids=update_layer_ids,
|
||||||
nnodes=nnodes,
|
nnodes=topology_num_nodes,
|
||||||
rank=rank,
|
rank=rank,
|
||||||
)
|
)
|
||||||
old_expert_location_metadata.update(
|
old_expert_location_metadata.update(
|
||||||
|
|||||||
@@ -4453,6 +4453,8 @@ class Scheduler(
|
|||||||
pending_ep_size=ElasticEPStateManager.get_pending_ep_size(),
|
pending_ep_size=ElasticEPStateManager.get_pending_ep_size(),
|
||||||
scale_phase=ElasticEPStateManager.get_scale_phase(),
|
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(
|
logger.debug(
|
||||||
"[Elastic EP][scale] scale requested: target_ep_size=%d; "
|
"[Elastic EP][scale] scale requested: target_ep_size=%d; "
|
||||||
"waiting for a joining cohort",
|
"waiting for a joining cohort",
|
||||||
|
|||||||
@@ -444,7 +444,8 @@ class ModelRunner:
|
|||||||
)
|
)
|
||||||
if self.eplb_manager is not None:
|
if self.eplb_manager is not None:
|
||||||
self.eplb_manager.disable_rebalance(
|
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()
|
state = ElasticEPStateManager.instance()
|
||||||
@@ -460,6 +461,7 @@ class ModelRunner:
|
|||||||
)
|
)
|
||||||
if state is not None:
|
if state is not None:
|
||||||
state.scale_phase = "serving_expanded"
|
state.scale_phase = "serving_expanded"
|
||||||
|
self._rearm_eplb_after_elastic_scale()
|
||||||
|
|
||||||
def init_msprobe(self):
|
def init_msprobe(self):
|
||||||
self.msprobe_debugger = misc_utils.create_msprobe_debugger(self.server_args)
|
self.msprobe_debugger = misc_utils.create_msprobe_debugger(self.server_args)
|
||||||
@@ -1703,6 +1705,26 @@ class ModelRunner:
|
|||||||
def _elastic_global_rank(self) -> int:
|
def _elastic_global_rank(self) -> int:
|
||||||
return self.ps.tp_rank + self.server_args.ep_join_rank_offset
|
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:
|
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:
|
if self.ps.tp_rank != 0 or self.server_args.is_ep_scale_joiner:
|
||||||
return
|
return
|
||||||
@@ -1769,7 +1791,8 @@ class ModelRunner:
|
|||||||
|
|
||||||
if self.eplb_manager is not None:
|
if self.eplb_manager is not None:
|
||||||
self.eplb_manager.disable_rebalance(
|
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
|
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",
|
log_tag="JOINER" if self.server_args.is_ep_scale_joiner else "PRIMARY",
|
||||||
)
|
)
|
||||||
ElasticEPStateManager.commit_scale()
|
ElasticEPStateManager.commit_scale()
|
||||||
|
self._rearm_eplb_after_elastic_scale()
|
||||||
|
|
||||||
if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner:
|
if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner:
|
||||||
from sglang.srt.managers.io_struct import ElasticScaleUpdateReq
|
from sglang.srt.managers.io_struct import ElasticScaleUpdateReq
|
||||||
@@ -1844,6 +1868,7 @@ class ModelRunner:
|
|||||||
if timeout.item():
|
if timeout.item():
|
||||||
error = f"Timed out waiting for ranks to join target EP size {pending_size}"
|
error = f"Timed out waiting for ranks to join target EP size {pending_size}"
|
||||||
ElasticEPStateManager.fail_scale(error)
|
ElasticEPStateManager.fail_scale(error)
|
||||||
|
self._reset_eplb_after_elastic_scale_failure()
|
||||||
self._report_elastic_scale_failure(error, effective_size)
|
self._report_elastic_scale_failure(error, effective_size)
|
||||||
if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner:
|
if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner:
|
||||||
logger.error("[Elastic EP] %s", error)
|
logger.error("[Elastic EP] %s", error)
|
||||||
@@ -1859,6 +1884,7 @@ class ModelRunner:
|
|||||||
f"joining cohort target {cohort_target}"
|
f"joining cohort target {cohort_target}"
|
||||||
)
|
)
|
||||||
ElasticEPStateManager.fail_scale(error)
|
ElasticEPStateManager.fail_scale(error)
|
||||||
|
self._reset_eplb_after_elastic_scale_failure()
|
||||||
self._report_elastic_scale_failure(error, effective_size)
|
self._report_elastic_scale_failure(error, effective_size)
|
||||||
if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner:
|
if self.ps.tp_rank == 0 and not self.server_args.is_ep_scale_joiner:
|
||||||
logger.error("[Elastic EP] %s", error)
|
logger.error("[Elastic EP] %s", error)
|
||||||
|
|||||||
@@ -407,6 +407,7 @@ class _ElasticScaleUpEndToEndBase(CustomTestCase):
|
|||||||
self._generate_logprob_ok("post-scale")
|
self._generate_logprob_ok("post-scale")
|
||||||
|
|
||||||
self._run_post_scale_gsm8k()
|
self._run_post_scale_gsm8k()
|
||||||
|
self._generate_logprob_ok("after post-scale workload")
|
||||||
|
|
||||||
|
|
||||||
@unittest.skipUnless(
|
@unittest.skipUnless(
|
||||||
@@ -458,6 +459,7 @@ class TestElasticScaleUp4To5To6(_ElasticScaleUpEndToEndBase):
|
|||||||
self._generate_ok("after second scale")
|
self._generate_ok("after second scale")
|
||||||
self._generate_logprob_ok("after second scale")
|
self._generate_logprob_ok("after second scale")
|
||||||
self._run_post_scale_gsm8k()
|
self._run_post_scale_gsm8k()
|
||||||
|
self._generate_logprob_ok("after post-scale workload")
|
||||||
|
|
||||||
|
|
||||||
@unittest.skipUnless(
|
@unittest.skipUnless(
|
||||||
|
|||||||
Reference in New Issue
Block a user