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
@@ -2,6 +2,7 @@ import logging
import re
import threading
import time
from typing import Any, Callable
import torch
import zmq
@@ -29,17 +30,25 @@ def extract_layer_and_expert_id(param_name):
class ExpertBackupClient:
def __init__(self, server_args: ServerArgs, model_runner):
def __init__(
self,
*,
server_args: ServerArgs,
model_config,
moe_ep_size: int,
moe_ep_rank: int,
get_model: Callable[[], Any],
):
context = zmq.Context(2)
self.server_args = server_args
self.engine_num = server_args.nnodes
self.engine_rank = server_args.node_rank
self.recv_list = [None] * self.engine_num
self.ready_sockets = [None] * self.engine_num
self.model_runner = model_runner
self.moe_ep_size = model_runner.ps.moe_ep_size
self.model_config = model_runner.model_config
self.moe_ep_rank = model_runner.ps.moe_ep_rank
self._get_model = get_model
self.moe_ep_size = moe_ep_size
self.model_config = model_config
self.moe_ep_rank = moe_ep_rank
self.dram_map_list = [None] * self.engine_num
self.session_id_list = [None] * self.engine_num
self.transfer_engine = None
@@ -87,7 +96,7 @@ class ExpertBackupClient:
self.transfer_engine = get_mooncake_transfer_engine()
self.params_dict = dict(self.model_runner.model.named_parameters())
self.params_dict = dict(self._get_model().named_parameters())
for name, param in self.params_dict.items():
param_data = param.data
ret_value = self.transfer_engine.engine.register_memory(
+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,
@@ -344,7 +344,7 @@ class ModelRunner:
create_offloader_from_server_args(server_args, dp_rank=self.ps.dp_rank)
)
self._weight_checker = WeightChecker(model_runner=self)
self._weight_checker = WeightChecker(get_model=lambda: self.model, ps=self.ps)
if envs.SGLANG_DETECT_SLOW_RANK.get():
slow_rank_detector.execute()
@@ -526,7 +526,15 @@ class ModelRunner:
# Expert parallelism
self.eplb_manager = (
EPLBManager(self)
EPLBManager(
server_args=self.server_args,
model_config=self.model_config,
ps=self.ps,
get_model=lambda: self.model,
get_expert_location_updater=lambda: self.expert_location_updater,
get_expert_backup_client=lambda: self.expert_backup_client,
get_weight_updater=lambda: self.weight_updater,
)
if self.server_args.enable_eplb and (not self.is_draft_worker)
else None
)
@@ -556,7 +564,13 @@ class ModelRunner:
# Load the expert backup client
self.expert_backup_client = (
ExpertBackupClient(self.server_args, self)
ExpertBackupClient(
server_args=self.server_args,
model_config=self.model_config,
moe_ep_size=self.ps.moe_ep_size,
moe_ep_rank=self.ps.moe_ep_rank,
get_model=lambda: self.model,
)
if (
self.server_args.enable_elastic_expert_backup
and self.server_args.elastic_ep_backend is not None
+10 -8
View File
@@ -1,7 +1,7 @@
import hashlib
import logging
import time
from typing import Dict, Iterable, NamedTuple, Optional, Set
from typing import Any, Callable, Dict, Iterable, NamedTuple, Optional, Set
import torch
import torch.distributed as dist
@@ -64,8 +64,9 @@ def _is_non_persistent_buffer_name(name: str) -> bool:
class WeightChecker:
def __init__(self, model_runner):
self._model_runner = model_runner
def __init__(self, *, get_model: Callable[[], Any], ps: Any):
self._get_model = get_model
self._ps = ps
self._snapshot_tensors = None
def handle(self, action: str, allow_quant_error: bool = False) -> Optional[Dict]:
@@ -101,7 +102,7 @@ class WeightChecker:
def _compare(self, allow_quant_error: bool = False):
assert self._snapshot_tensors is not None
quantized_set = _build_quantized_set(self._model_runner.model)
quantized_set = _build_quantized_set(self._get_model())
skip_compare_names = {
name
for name, param in self._model_state()
@@ -121,7 +122,7 @@ class WeightChecker:
torch.cuda.synchronize()
start = time.perf_counter()
quantized_set = _build_quantized_set(self._model_runner.model)
quantized_set = _build_quantized_set(self._get_model())
skip_compare_names = {
name
for name, param in self._model_state()
@@ -157,7 +158,7 @@ class WeightChecker:
return info.model_dump()
def _parallelism_info(self) -> ParallelismInfo:
ps = self._model_runner.ps
ps = self._ps
return ParallelismInfo(
tp_rank=ps.tp_rank,
tp_size=ps.tp_size,
@@ -170,8 +171,9 @@ class WeightChecker:
)
def _model_state(self):
yield from self._model_runner.model.named_parameters()
yield from self._model_runner.model.named_buffers()
model = self._get_model()
yield from model.named_parameters()
yield from model.named_buffers()
def _hash_tensor(t: torch.Tensor) -> str: