Split initialize() into orchestration helpers (#31169)

This commit is contained in:
fzyzcjy
2026-07-14 16:09:00 +08:00
committed by GitHub
parent 64a70c9097
commit 0fe2dbd42c
2 changed files with 182 additions and 151 deletions
+70 -2
View File
@@ -3,13 +3,17 @@ from __future__ import annotations
import logging
import time
from dataclasses import dataclass
from typing import Iterator, List, Optional
from typing import TYPE_CHECKING, Iterator, List, Optional
import torch
from sglang.srt.distributed import get_world_group, parallel_state
from sglang.srt.eplb.expert_location import broadcast_global_expert_location_metadata
from sglang.srt.managers.schedule_batch import ServerArgs
from sglang.srt.utils import is_cpu, is_cuda
from sglang.srt.utils import broadcast_pyobj, is_cpu, is_cuda
if TYPE_CHECKING:
from sglang.srt.eplb.eplb_manager import EPLBManager
logger = logging.getLogger(__name__)
@@ -221,3 +225,67 @@ def get_healthy_expert_location_src_rank(
"No healthy rank found for broadcasting expert location metadata. "
"All ranks are marked as elastic_ep_rejoin."
)
def maybe_recover_ep_ranks(
*,
tp_group: parallel_state.GroupCoordinator,
eplb_manager: EPLBManager,
random_seed: int,
) -> bool:
# TODO(perf): `active_ranks.all()` on a CUDA tensor triggers host-device
# synchronization, and this function is on the forward-path.
# This check only runs when `--elastic-ep-backend` is enabled, so the
# synchronization overhead does not propagate to other configs.
# Leave for future optimization of the elastic EP path.
if tp_group.active_ranks.all() and tp_group.active_ranks_cpu.all():
return False
tp_active_ranks = tp_group.active_ranks.detach().cpu().numpy()
tp_active_ranks_cpu = tp_group.active_ranks_cpu.detach().numpy()
tp_active_ranks &= tp_active_ranks_cpu
# NOTE: `ranks_to_recover` uses indices in `tp_group`. For the current
# Mooncake elastic EP implementation we assume `--pp-size=1`, so the
# tp-group index is the same as the global rank index.
ranks_to_recover = [
i for i in range(len(tp_active_ranks)) if not tp_active_ranks[i]
]
# try_recover_ranks polls peer state via Mooncake EP backend.
# Mooncake's internal semantics guarantee that all ranks observe
# consistent peer readiness state, so collective operations below
# are safe even though polling appears local.
if ranks_to_recover and try_recover_ranks(ranks_to_recover):
eplb_manager.reset_generator()
broadcast_global_expert_location_metadata(
src_rank=get_healthy_expert_location_src_rank(
invoked_in_elastic_ep_rejoin_path=False
)
)
ElasticEPStateManager.instance().reset()
broadcast_pyobj(
[random_seed],
parallel_state.get_world_group().rank,
parallel_state.get_world_group().cpu_group,
src=parallel_state.get_world_group().ranks[0],
)
logger.info(f"recover ranks {ranks_to_recover} done")
return True
return False
def maybe_rebalance_after_rank_fault(*, eplb_manager: EPLBManager) -> bool:
elastic_ep_state = ElasticEPStateManager.instance()
if elastic_ep_state is None or elastic_ep_state.is_active_equal_last():
return False
elastic_ep_state.snapshot_active_to_last()
elastic_ep_state.sync_active_to_cpu()
logger.info("EPLB due to rank faults")
gen = eplb_manager.rebalance()
while True:
try:
next(gen)
except StopIteration:
break
return True
+112 -149
View File
@@ -32,10 +32,7 @@ from sglang.srt.configs.model_config import (
)
from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp
from sglang.srt.debug_utils.dumper import dumper
from sglang.srt.distributed import (
bootstrap,
get_world_group,
)
from sglang.srt.distributed import bootstrap
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
maybe_init_shared_mooncake_transfer_engine,
)
@@ -45,7 +42,8 @@ from sglang.srt.elastic_ep.elastic_ep import (
ElasticEPStateManager,
get_healthy_expert_location_src_rank,
join_process_groups,
try_recover_ranks,
maybe_rebalance_after_rank_fault,
maybe_recover_ep_ranks,
)
from sglang.srt.elastic_ep.expert_backup_client import ExpertBackupClient
from sglang.srt.environ import envs
@@ -178,7 +176,6 @@ from sglang.srt.state_capturer.routed_experts import (
set_global_experts_capturer,
)
from sglang.srt.utils import (
broadcast_pyobj,
cpu_has_amx_support,
enable_show_time_cost,
get_available_gpu_memory,
@@ -476,43 +473,83 @@ class ModelRunner:
)
def initialize(self):
server_args = self.server_args
self.init_memory_saver_adapter()
self.maybe_init_remote_instance_transfer_engine()
self.maybe_init_expert_location_metadata()
self.maybe_init_lplb_solvers()
self.maybe_init_eplb_manager()
self.expert_location_updater = ExpertLocationUpdater()
self.maybe_init_elastic_ep()
self.init_token_oracle()
self.sampler = create_sampler()
self.load_model()
prepare_moe_topk(
model=self.model,
model_config=self.model_config,
server_args=self.server_args,
moe_ep_size=self.ps.moe_ep_size,
moe_ep_rank=self.ps.moe_ep_rank,
)
# Must run before backend/graph init so no draft graph records a
# routed-experts capture-write kernel.
if self.is_draft_worker:
disable_routed_experts_capture_for_draft(self.model)
self.maybe_init_expert_backup_client()
self.remote_instance_weight_transporter.maybe_register_and_publish_weight_info()
self.layer_info: ModelLayerInfo = resolve_layer_indices(
model=self.model,
model_config=self.model_config,
is_draft_worker=self.is_draft_worker,
spec_algorithm=self.spec_algorithm,
)
adjust_hybrid_swa_layer_ids(
model_config=self.model_config,
start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer,
is_hybrid_swa=self.is_hybrid_swa,
)
self.maybe_apply_post_load_model_transforms()
self.maybe_init_lora_manager()
self.maybe_enable_batch_invariant_mode()
self.configure_kv_cache_dtype()
def init_memory_saver_adapter(self):
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=self.server_args.enable_memory_saver
)
def maybe_init_remote_instance_transfer_engine(self):
if self.server_args.remote_instance_weight_loader_use_transfer_engine():
self.remote_instance_weight_transporter.init_engine()
if not self.is_draft_worker:
set_global_expert_location_metadata(
compute_initial_expert_location_metadata(
server_args=server_args,
model_config=self.model_config,
moe_ep_rank=self.ps.moe_ep_rank,
)
def maybe_init_expert_location_metadata(self):
if self.is_draft_worker:
return
set_global_expert_location_metadata(
compute_initial_expert_location_metadata(
server_args=self.server_args,
model_config=self.model_config,
moe_ep_rank=self.ps.moe_ep_rank,
)
if self.ps.tp_rank == 0 and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get():
logger.info(
"Initial expert_location_metadata:\n%s",
format_expert_location_layout(
get_global_expert_location_metadata()
),
)
set_global_expert_distribution_recorder(
ExpertDistributionRecorder.init_new(
server_args,
get_global_expert_location_metadata(),
rank=self.ps.tp_rank,
)
)
if self.ps.tp_rank == 0 and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get():
logger.info(
"Initial expert_location_metadata:\n%s",
format_expert_location_layout(get_global_expert_location_metadata()),
)
set_global_expert_distribution_recorder(
ExpertDistributionRecorder.init_new(
self.server_args,
get_global_expert_location_metadata(),
rank=self.ps.tp_rank,
)
)
def maybe_init_lplb_solvers(self):
if self.server_args.ep_dispatch_algorithm == "lp" and not self.is_draft_worker:
init_lplb_solvers(model_config=self.model_config)
# Expert parallelism
def maybe_init_eplb_manager(self):
self.eplb_manager = (
EPLBManager(
server_args=self.server_args,
@@ -526,31 +563,18 @@ class ModelRunner:
if self.server_args.enable_eplb and (not self.is_draft_worker)
else None
)
self.expert_location_updater = ExpertLocationUpdater()
def maybe_init_elastic_ep(self):
if self.server_args.elastic_ep_backend:
ElasticEPStateManager.init(self.server_args)
def init_token_oracle(self):
self._token_oracle_manager = install_token_oracle_from_env(
server_args=server_args,
server_args=self.server_args,
vocab_size=self.model_config.vocab_size,
)
# Load the model
self.sampler = create_sampler()
self.load_model()
prepare_moe_topk(
model=self.model,
model_config=self.model_config,
server_args=self.server_args,
moe_ep_size=self.ps.moe_ep_size,
moe_ep_rank=self.ps.moe_ep_rank,
)
# Must run before backend/graph init so no draft graph records a
# routed-experts capture-write kernel.
if self.is_draft_worker:
disable_routed_experts_capture_for_draft(self.model)
# Load the expert backup client
def maybe_init_expert_backup_client(self):
self.expert_backup_client = (
ExpertBackupClient(
server_args=self.server_args,
@@ -566,45 +590,25 @@ class ModelRunner:
else None
)
self.remote_instance_weight_transporter.maybe_register_and_publish_weight_info()
self.layer_info: ModelLayerInfo = resolve_layer_indices(
model=self.model,
model_config=self.model_config,
is_draft_worker=self.is_draft_worker,
spec_algorithm=self.spec_algorithm,
)
adjust_hybrid_swa_layer_ids(
model_config=self.model_config,
start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer,
is_hybrid_swa=self.is_hybrid_swa,
)
# Apply torchao quantization
torchao_applied = getattr(self.model, "torchao_applied", False)
def maybe_apply_post_load_model_transforms(self):
# In layered loading, torchao may have been applied
torchao_applied = getattr(self.model, "torchao_applied", False)
if not torchao_applied:
apply_torchao_config_to_model(self.model, get_server_args().torchao_config)
# Apply torch TP if the model supports it
supports_torch_tp = getattr(self.model, "supports_torch_tp", False)
if self.ps.tp_size > 1 and supports_torch_tp:
self.apply_torch_tp()
# Init lora
if server_args.enable_lora:
def maybe_init_lora_manager(self):
if self.server_args.enable_lora:
self.init_lora_manager()
# Enable batch invariant mode
if server_args.enable_deterministic_inference:
def maybe_enable_batch_invariant_mode(self):
if self.server_args.enable_deterministic_inference:
from sglang.srt.batch_invariant_ops import enable_batch_invariant_mode
enable_batch_invariant_mode()
self.configure_kv_cache_dtype()
def get_pp_proxy_topk_size(self) -> Optional[int]:
return misc_utils.resolve_pp_proxy_topk_size(
model_config=self.model_config,
@@ -665,34 +669,38 @@ class ModelRunner:
# Init ngram embedding token table
self.init_ngram_embedding_manager()
if self.enable_hisparse:
from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator
from sglang.srt.mem_cache.sparsity import parse_hisparse_config
hisparse_cfg = parse_hisparse_config(self.server_args)
hisparse_top_k = getattr(
self.model_config.hf_text_config, "index_topk", hisparse_cfg.top_k
)
self.hisparse_coordinator = HiSparseCoordinator(
req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
top_k=hisparse_top_k,
device_buffer_size=hisparse_cfg.device_buffer_size,
device=self.device,
tp_group=(
self.attention_tp_group.cpu_group
if self.server_args.enable_dp_attention
else self.tp_group.cpu_group
),
host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
swap_in_block_size=hisparse_cfg.swap_in_block_size,
)
self.maybe_init_hisparse_coordinator()
self.init_routed_experts_capturer()
self.init_indexer_capturer()
self.graph_shared_output = None
def maybe_init_hisparse_coordinator(self):
if not self.enable_hisparse:
return
from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator
from sglang.srt.mem_cache.sparsity import parse_hisparse_config
hisparse_cfg = parse_hisparse_config(self.server_args)
hisparse_top_k = getattr(
self.model_config.hf_text_config, "index_topk", hisparse_cfg.top_k
)
self.hisparse_coordinator = HiSparseCoordinator(
req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
top_k=hisparse_top_k,
device_buffer_size=hisparse_cfg.device_buffer_size,
device=self.device,
tp_group=(
self.attention_tp_group.cpu_group
if self.server_args.enable_dp_attention
else self.tp_group.cpu_group
),
host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
swap_in_block_size=hisparse_cfg.swap_in_block_size,
)
def post_capture_resize_kv_pool(self):
resize = compute_post_capture_kv_resize(self)
self.max_total_num_tokens = resize.max_total_num_tokens
@@ -911,47 +919,6 @@ class ModelRunner:
tp_rank=self.ps.tp_rank,
)
def maybe_recover_ep_ranks(self):
# TODO(perf): `active_ranks.all()` on a CUDA tensor triggers host-device
# synchronization, and this function is on the forward-path.
# This check only runs when `--elastic-ep-backend` is enabled, so the
# synchronization overhead does not propagate to other configs.
# Leave for future optimization of the elastic EP path.
if self.tp_group.active_ranks.all() and self.tp_group.active_ranks_cpu.all():
return
tp_active_ranks = self.tp_group.active_ranks.detach().cpu().numpy()
tp_active_ranks_cpu = self.tp_group.active_ranks_cpu.detach().numpy()
tp_active_ranks &= tp_active_ranks_cpu
# NOTE: `ranks_to_recover` uses indices in `tp_group`. For the current
# Mooncake elastic EP implementation we assume `--pp-size=1`, so the
# tp-group index is the same as the global rank index.
ranks_to_recover = [
i for i in range(len(tp_active_ranks)) if not tp_active_ranks[i]
]
# try_recover_ranks polls peer state via Mooncake EP backend.
# Mooncake's internal semantics guarantee that all ranks observe
# consistent peer readiness state, so collective operations below
# are safe even though polling appears local.
if ranks_to_recover and try_recover_ranks(ranks_to_recover):
self.forward_pass_id = 0
self.eplb_manager.reset_generator()
broadcast_global_expert_location_metadata(
src_rank=get_healthy_expert_location_src_rank(
invoked_in_elastic_ep_rejoin_path=False
)
)
ElasticEPStateManager.instance().reset()
broadcast_pyobj(
[self.server_args.random_seed],
get_world_group().rank,
get_world_group().cpu_group,
src=get_world_group().ranks[0],
)
logger.info(f"recover ranks {ranks_to_recover} done")
def init_lora_manager(self):
self.lora_manager = LoRAManager(
base_model=self.model,
@@ -1249,8 +1216,14 @@ class ModelRunner:
self.msprobe_debugger.stop()
self.msprobe_debugger.step()
if self.enable_elastic_ep:
self.maybe_recover_ep_ranks()
if self.server_args.elastic_ep_backend is not None:
recovered = maybe_recover_ep_ranks(
tp_group=self.tp_group,
eplb_manager=self.eplb_manager,
random_seed=self.server_args.random_seed,
)
if recovered:
self.forward_pass_id = 0
return output
@@ -1499,17 +1472,7 @@ class ModelRunner:
reinit_attn_backend: bool,
split_forward_count: int,
) -> ModelRunnerOutput:
elastic_ep_state = ElasticEPStateManager.instance()
if elastic_ep_state is not None and not elastic_ep_state.is_active_equal_last():
elastic_ep_state.snapshot_active_to_last()
elastic_ep_state.sync_active_to_cpu()
logging.info("EPLB due to rank faults")
gen = self.eplb_manager.rebalance()
while True:
try:
next(gen)
except StopIteration:
break
if maybe_rebalance_after_rank_fault(eplb_manager=self.eplb_manager):
output = self._forward_raw(
forward_batch,
pp_proxy_tensors,