Narrow component dependencies to injected fields instead of ModelRunner (#31166)
This commit is contained in:
@@ -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(
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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):
|
||||||
|
|||||||
Reference in New Issue
Block a user