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 re
import threading import threading
import time import time
from typing import Any, Callable
import torch import torch
import zmq import zmq
@@ -29,17 +30,25 @@ def extract_layer_and_expert_id(param_name):
class ExpertBackupClient: 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) context = zmq.Context(2)
self.server_args = server_args self.server_args = server_args
self.engine_num = server_args.nnodes self.engine_num = server_args.nnodes
self.engine_rank = server_args.node_rank self.engine_rank = server_args.node_rank
self.recv_list = [None] * self.engine_num self.recv_list = [None] * self.engine_num
self.ready_sockets = [None] * self.engine_num self.ready_sockets = [None] * self.engine_num
self.model_runner = model_runner self._get_model = get_model
self.moe_ep_size = model_runner.ps.moe_ep_size self.moe_ep_size = moe_ep_size
self.model_config = model_runner.model_config self.model_config = model_config
self.moe_ep_rank = model_runner.ps.moe_ep_rank self.moe_ep_rank = moe_ep_rank
self.dram_map_list = [None] * self.engine_num self.dram_map_list = [None] * self.engine_num
self.session_id_list = [None] * self.engine_num self.session_id_list = [None] * self.engine_num
self.transfer_engine = None self.transfer_engine = None
@@ -87,7 +96,7 @@ class ExpertBackupClient:
self.transfer_engine = get_mooncake_transfer_engine() 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(): for name, param in self.params_dict.items():
param_data = param.data param_data = param.data
ret_value = self.transfer_engine.engine.register_memory( ret_value = self.transfer_engine.engine.register_memory(
+37 -19
View File
@@ -1,6 +1,8 @@
from __future__ import annotations
import logging import logging
import time import time
from typing import TYPE_CHECKING, List from typing import TYPE_CHECKING, Any, Callable, List
import torch.cuda import torch.cuda
from torch import nn 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 from sglang.srt.runtime_context import get_server_args
if TYPE_CHECKING: 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__) logger = logging.getLogger(__name__)
class EPLBManager: 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__() super().__init__()
self._model_runner = model_runner # These collaborators are set on ModelRunner AFTER EPLBManager is
self._server_args = model_runner.server_args # 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._rebalance_layers_per_chunk = (
self._server_args.eplb_rebalance_layers_per_chunk self._server_args.eplb_rebalance_layers_per_chunk
) )
@@ -83,7 +104,7 @@ class EPLBManager:
return return
expert_location_metadata = ExpertLocationMetadata.init_by_eplb( 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 ( 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: if len(update_layer_ids_chunks) > 1:
yield yield
update_expert_location_with_recovery( update_expert_location_with_recovery(
expert_location_updater=self._model_runner.expert_location_updater, expert_location_updater=self._get_expert_location_updater(),
model=self._model_runner.model, model=self._get_model(),
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._model_runner.server_args.nnodes, nnodes=self._server_args.nnodes,
tp_rank=self._model_runner.ps.tp_rank, tp_rank=self._ps.tp_rank,
expert_backup_client=self._model_runner.expert_backup_client, expert_backup_client=self._get_expert_backup_client(),
update_weights_from_disk_callable=self._model_runner.weight_updater.update_weights_from_disk, update_weights_from_disk_callable=self._get_weight_updater().update_weights_from_disk,
ep_dispatch_algorithm=self._model_runner.server_args.ep_dispatch_algorithm, ep_dispatch_algorithm=self._server_args.ep_dispatch_algorithm,
init_lplb_solvers_callable=lambda: init_lplb_solvers( 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]]: def _compute_update_layer_ids_chunks(self) -> List[List[int]]:
all_layer_ids = sorted( 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 chunk_size = self._rebalance_layers_per_chunk or 1000000
return list(_chunk_list(all_layer_ids, chunk_size=chunk_size)) return list(_chunk_list(all_layer_ids, chunk_size=chunk_size))
def _should_log_expert_location_metadata(self) -> bool: def _should_log_expert_location_metadata(self) -> bool:
return ( return self._ps.tp_rank == 0 and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get()
self._model_runner.ps.tp_rank == 0
and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get()
)
def _log_rebalance_layout_before_update( def _log_rebalance_layout_before_update(
self, self,
@@ -344,7 +344,7 @@ class ModelRunner:
create_offloader_from_server_args(server_args, dp_rank=self.ps.dp_rank) 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(): if envs.SGLANG_DETECT_SLOW_RANK.get():
slow_rank_detector.execute() slow_rank_detector.execute()
@@ -526,7 +526,15 @@ class ModelRunner:
# Expert parallelism # Expert parallelism
self.eplb_manager = ( 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) if self.server_args.enable_eplb and (not self.is_draft_worker)
else None else None
) )
@@ -556,7 +564,13 @@ class ModelRunner:
# Load the expert backup client # Load the expert backup client
self.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 ( if (
self.server_args.enable_elastic_expert_backup self.server_args.enable_elastic_expert_backup
and self.server_args.elastic_ep_backend is not None and self.server_args.elastic_ep_backend is not None
+10 -8
View File
@@ -1,7 +1,7 @@
import hashlib import hashlib
import logging import logging
import time import time
from typing import Dict, Iterable, NamedTuple, Optional, Set from typing import Any, Callable, Dict, Iterable, NamedTuple, Optional, Set
import torch import torch
import torch.distributed as dist import torch.distributed as dist
@@ -64,8 +64,9 @@ def _is_non_persistent_buffer_name(name: str) -> bool:
class WeightChecker: class WeightChecker:
def __init__(self, model_runner): def __init__(self, *, get_model: Callable[[], Any], ps: Any):
self._model_runner = model_runner self._get_model = get_model
self._ps = ps
self._snapshot_tensors = None self._snapshot_tensors = None
def handle(self, action: str, allow_quant_error: bool = False) -> Optional[Dict]: 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): def _compare(self, allow_quant_error: bool = False):
assert self._snapshot_tensors is not None 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 = { skip_compare_names = {
name name
for name, param in self._model_state() for name, param in self._model_state()
@@ -121,7 +122,7 @@ class WeightChecker:
torch.cuda.synchronize() torch.cuda.synchronize()
start = time.perf_counter() start = time.perf_counter()
quantized_set = _build_quantized_set(self._model_runner.model) quantized_set = _build_quantized_set(self._get_model())
skip_compare_names = { skip_compare_names = {
name name
for name, param in self._model_state() for name, param in self._model_state()
@@ -157,7 +158,7 @@ class WeightChecker:
return info.model_dump() return info.model_dump()
def _parallelism_info(self) -> ParallelismInfo: def _parallelism_info(self) -> ParallelismInfo:
ps = self._model_runner.ps ps = self._ps
return ParallelismInfo( return ParallelismInfo(
tp_rank=ps.tp_rank, tp_rank=ps.tp_rank,
tp_size=ps.tp_size, tp_size=ps.tp_size,
@@ -170,8 +171,9 @@ class WeightChecker:
) )
def _model_state(self): def _model_state(self):
yield from self._model_runner.model.named_parameters() model = self._get_model()
yield from self._model_runner.model.named_buffers() yield from model.named_parameters()
yield from model.named_buffers()
def _hash_tensor(t: torch.Tensor) -> str: def _hash_tensor(t: torch.Tensor) -> str:
@@ -20,6 +20,7 @@ from unittest.mock import patch
import torch import torch
from torch import nn from torch import nn
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.quantization.fp8_utils import ( from sglang.srt.layers.quantization.fp8_utils import (
quant_weight_ue8m0, quant_weight_ue8m0,
transform_scale_ue8m0, transform_scale_ue8m0,
@@ -129,14 +130,18 @@ class _FakeModelRunner:
dp_size: int = 1, dp_size: int = 1,
pp_rank: int = 0, pp_rank: int = 0,
pp_size: int = 1, pp_size: int = 1,
attn_dp_size: int | None = None,
): ):
self.model = model self.model = model
self.tp_rank = tp_rank self.ps = ParallelState.trivial(
self.tp_size = tp_size tp_rank=tp_rank,
self.dp_rank = dp_rank tp_size=tp_size,
self.dp_size = dp_size dp_rank=dp_rank,
self.pp_rank = pp_rank dp_size=dp_size,
self.pp_size = pp_size attn_dp_size=attn_dp_size if attn_dp_size is not None else dp_size,
pp_rank=pp_rank,
pp_size=pp_size,
)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -490,7 +495,8 @@ class _WeightCheckerTestBase(CustomTestCase):
def setUp(self): def setUp(self):
torch.manual_seed(0) torch.manual_seed(0)
self.model = _TinyModel().cuda() self.model = _TinyModel().cuda()
self.checker = WeightChecker(model_runner=_FakeModelRunner(self.model)) runner = _FakeModelRunner(self.model)
self.checker = WeightChecker(get_model=lambda: runner.model, ps=runner.ps)
class TestSnapshot(_WeightCheckerTestBase): class TestSnapshot(_WeightCheckerTestBase):
@@ -697,7 +703,9 @@ class _ChecksumTestBase(CustomTestCase):
pp_rank=0, pp_rank=0,
pp_size=1, pp_size=1,
) )
self.checker = WeightChecker(model_runner=self.runner) self.checker = WeightChecker(
get_model=lambda: self.runner.model, ps=self.runner.ps
)
class TestComputeChecksum(_ChecksumTestBase): class TestComputeChecksum(_ChecksumTestBase):