Narrow component dependencies to injected fields instead of ModelRunner (#31166)

This commit is contained in:
fzyzcjy
2026-07-14 16:03:07 +08:00
committed by GitHub
parent 6999007a13
commit 54f99a21d5
5 changed files with 95 additions and 44 deletions
+37 -19
View File
@@ -1,6 +1,8 @@
from __future__ import annotations
import logging
import time
from typing import TYPE_CHECKING, List
from typing import TYPE_CHECKING, Any, Callable, List
import torch.cuda
from torch import nn
@@ -17,16 +19,35 @@ from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater
from sglang.srt.runtime_context import get_server_args
if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__)
class EPLBManager:
def __init__(self, model_runner: "ModelRunner"):
def __init__(
self,
*,
server_args: ServerArgs,
model_config: ModelConfig,
ps: Any,
get_model: Callable[[], nn.Module],
get_expert_location_updater: Callable[[], ExpertLocationUpdater],
get_expert_backup_client: Callable[[], Any],
get_weight_updater: Callable[[], Any],
):
super().__init__()
self._model_runner = model_runner
self._server_args = model_runner.server_args
# These collaborators are set on ModelRunner AFTER EPLBManager is
# constructed (model load, expert_backup_client, weight_updater), so
# they are read through getters at rebalance time, not captured here.
self._server_args = server_args
self._model_config = model_config
self._ps = ps
self._get_model = get_model
self._get_expert_location_updater = get_expert_location_updater
self._get_expert_backup_client = get_expert_backup_client
self._get_weight_updater = get_weight_updater
self._rebalance_layers_per_chunk = (
self._server_args.eplb_rebalance_layers_per_chunk
)
@@ -83,7 +104,7 @@ class EPLBManager:
return
expert_location_metadata = ExpertLocationMetadata.init_by_eplb(
self._server_args, self._model_runner.model_config, logical_count
self._server_args, self._model_config, logical_count
)
from sglang.srt.model_executor.model_runner_components.moe_ep_setup import (
@@ -102,17 +123,17 @@ class EPLBManager:
if len(update_layer_ids_chunks) > 1:
yield
update_expert_location_with_recovery(
expert_location_updater=self._model_runner.expert_location_updater,
model=self._model_runner.model,
expert_location_updater=self._get_expert_location_updater(),
model=self._get_model(),
new_expert_location_metadata=expert_location_metadata,
update_layer_ids=chunk_layer_ids,
nnodes=self._model_runner.server_args.nnodes,
tp_rank=self._model_runner.ps.tp_rank,
expert_backup_client=self._model_runner.expert_backup_client,
update_weights_from_disk_callable=self._model_runner.weight_updater.update_weights_from_disk,
ep_dispatch_algorithm=self._model_runner.server_args.ep_dispatch_algorithm,
nnodes=self._server_args.nnodes,
tp_rank=self._ps.tp_rank,
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,
init_lplb_solvers_callable=lambda: init_lplb_solvers(
model_config=self._model_runner.model_config
model_config=self._model_config
),
)
@@ -142,16 +163,13 @@ class EPLBManager:
def _compute_update_layer_ids_chunks(self) -> List[List[int]]:
all_layer_ids = sorted(
list(self._model_runner.model.routed_experts_weights_of_layer.keys())
list(self._get_model().routed_experts_weights_of_layer.keys())
)
chunk_size = self._rebalance_layers_per_chunk or 1000000
return list(_chunk_list(all_layer_ids, chunk_size=chunk_size))
def _should_log_expert_location_metadata(self) -> bool:
return (
self._model_runner.ps.tp_rank == 0
and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get()
)
return self._ps.tp_rank == 0 and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get()
def _log_rebalance_layout_before_update(
self,