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 logging
import time import time
from dataclasses import dataclass from dataclasses import dataclass
from typing import Iterator, List, Optional from typing import TYPE_CHECKING, Iterator, List, Optional
import torch import torch
from sglang.srt.distributed import get_world_group, parallel_state 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.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__) logger = logging.getLogger(__name__)
@@ -221,3 +225,67 @@ def get_healthy_expert_location_src_rank(
"No healthy rank found for broadcasting expert location metadata. " "No healthy rank found for broadcasting expert location metadata. "
"All ranks are marked as elastic_ep_rejoin." "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.configs.update_config import adjust_config_with_unaligned_cpu_tp
from sglang.srt.debug_utils.dumper import dumper from sglang.srt.debug_utils.dumper import dumper
from sglang.srt.distributed import ( from sglang.srt.distributed import bootstrap
bootstrap,
get_world_group,
)
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import ( from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
maybe_init_shared_mooncake_transfer_engine, maybe_init_shared_mooncake_transfer_engine,
) )
@@ -45,7 +42,8 @@ from sglang.srt.elastic_ep.elastic_ep import (
ElasticEPStateManager, ElasticEPStateManager,
get_healthy_expert_location_src_rank, get_healthy_expert_location_src_rank,
join_process_groups, 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.elastic_ep.expert_backup_client import ExpertBackupClient
from sglang.srt.environ import envs from sglang.srt.environ import envs
@@ -178,7 +176,6 @@ from sglang.srt.state_capturer.routed_experts import (
set_global_experts_capturer, set_global_experts_capturer,
) )
from sglang.srt.utils import ( from sglang.srt.utils import (
broadcast_pyobj,
cpu_has_amx_support, cpu_has_amx_support,
enable_show_time_cost, enable_show_time_cost,
get_available_gpu_memory, get_available_gpu_memory,
@@ -476,43 +473,83 @@ class ModelRunner:
) )
def initialize(self): 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( self.memory_saver_adapter = TorchMemorySaverAdapter.create(
enable=self.server_args.enable_memory_saver 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(): if self.server_args.remote_instance_weight_loader_use_transfer_engine():
self.remote_instance_weight_transporter.init_engine() self.remote_instance_weight_transporter.init_engine()
if not self.is_draft_worker: def maybe_init_expert_location_metadata(self):
set_global_expert_location_metadata( if self.is_draft_worker:
compute_initial_expert_location_metadata( return
server_args=server_args, set_global_expert_location_metadata(
model_config=self.model_config, compute_initial_expert_location_metadata(
moe_ep_rank=self.ps.moe_ep_rank, 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( if self.ps.tp_rank == 0 and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get():
"Initial expert_location_metadata:\n%s", logger.info(
format_expert_location_layout( "Initial expert_location_metadata:\n%s",
get_global_expert_location_metadata() 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,
)
) )
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: if self.server_args.ep_dispatch_algorithm == "lp" and not self.is_draft_worker:
init_lplb_solvers(model_config=self.model_config) init_lplb_solvers(model_config=self.model_config)
# Expert parallelism def maybe_init_eplb_manager(self):
self.eplb_manager = ( self.eplb_manager = (
EPLBManager( EPLBManager(
server_args=self.server_args, server_args=self.server_args,
@@ -526,31 +563,18 @@ class ModelRunner:
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
) )
self.expert_location_updater = ExpertLocationUpdater()
def maybe_init_elastic_ep(self):
if self.server_args.elastic_ep_backend: if self.server_args.elastic_ep_backend:
ElasticEPStateManager.init(self.server_args) ElasticEPStateManager.init(self.server_args)
def init_token_oracle(self):
self._token_oracle_manager = install_token_oracle_from_env( 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, 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 def maybe_init_expert_backup_client(self):
# routed-experts capture-write kernel.
if self.is_draft_worker:
disable_routed_experts_capture_for_draft(self.model)
# Load the expert backup client
self.expert_backup_client = ( self.expert_backup_client = (
ExpertBackupClient( ExpertBackupClient(
server_args=self.server_args, server_args=self.server_args,
@@ -566,45 +590,25 @@ class ModelRunner:
else None else None
) )
self.remote_instance_weight_transporter.maybe_register_and_publish_weight_info() def maybe_apply_post_load_model_transforms(self):
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)
# In layered loading, torchao may have been applied # In layered loading, torchao may have been applied
torchao_applied = getattr(self.model, "torchao_applied", False)
if not torchao_applied: if not torchao_applied:
apply_torchao_config_to_model(self.model, get_server_args().torchao_config) 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) supports_torch_tp = getattr(self.model, "supports_torch_tp", False)
if self.ps.tp_size > 1 and supports_torch_tp: if self.ps.tp_size > 1 and supports_torch_tp:
self.apply_torch_tp() self.apply_torch_tp()
# Init lora def maybe_init_lora_manager(self):
if server_args.enable_lora: if self.server_args.enable_lora:
self.init_lora_manager() self.init_lora_manager()
# Enable batch invariant mode def maybe_enable_batch_invariant_mode(self):
if server_args.enable_deterministic_inference: if self.server_args.enable_deterministic_inference:
from sglang.srt.batch_invariant_ops import enable_batch_invariant_mode from sglang.srt.batch_invariant_ops import enable_batch_invariant_mode
enable_batch_invariant_mode() enable_batch_invariant_mode()
self.configure_kv_cache_dtype()
def get_pp_proxy_topk_size(self) -> Optional[int]: def get_pp_proxy_topk_size(self) -> Optional[int]:
return misc_utils.resolve_pp_proxy_topk_size( return misc_utils.resolve_pp_proxy_topk_size(
model_config=self.model_config, model_config=self.model_config,
@@ -665,34 +669,38 @@ class ModelRunner:
# Init ngram embedding token table # Init ngram embedding token table
self.init_ngram_embedding_manager() self.init_ngram_embedding_manager()
if self.enable_hisparse: self.maybe_init_hisparse_coordinator()
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.init_routed_experts_capturer() self.init_routed_experts_capturer()
self.init_indexer_capturer() self.init_indexer_capturer()
self.graph_shared_output = None 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): def post_capture_resize_kv_pool(self):
resize = compute_post_capture_kv_resize(self) resize = compute_post_capture_kv_resize(self)
self.max_total_num_tokens = resize.max_total_num_tokens self.max_total_num_tokens = resize.max_total_num_tokens
@@ -911,47 +919,6 @@ class ModelRunner:
tp_rank=self.ps.tp_rank, 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): def init_lora_manager(self):
self.lora_manager = LoRAManager( self.lora_manager = LoRAManager(
base_model=self.model, base_model=self.model,
@@ -1249,8 +1216,14 @@ class ModelRunner:
self.msprobe_debugger.stop() self.msprobe_debugger.stop()
self.msprobe_debugger.step() self.msprobe_debugger.step()
if self.enable_elastic_ep: if self.server_args.elastic_ep_backend is not None:
self.maybe_recover_ep_ranks() 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 return output
@@ -1499,17 +1472,7 @@ class ModelRunner:
reinit_attn_backend: bool, reinit_attn_backend: bool,
split_forward_count: int, split_forward_count: int,
) -> ModelRunnerOutput: ) -> ModelRunnerOutput:
elastic_ep_state = ElasticEPStateManager.instance() if maybe_rebalance_after_rank_fault(eplb_manager=self.eplb_manager):
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
output = self._forward_raw( output = self._forward_raw(
forward_batch, forward_batch,
pp_proxy_tensors, pp_proxy_tensors,