Drop ModelRunner's duplicated parallel-degree fields and read them via self.ps (#31165)

This commit is contained in:
fzyzcjy
2026-07-14 16:02:40 +08:00
committed by GitHub
parent cfd17301a8
commit 6999007a13
20 changed files with 107 additions and 138 deletions
+1 -1
View File
@@ -530,7 +530,7 @@ def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
batch, batch,
dp_size=model_runner.server_args.dp_size, dp_size=model_runner.server_args.dp_size,
attn_tp_size=get_parallel().attn_tp_size, attn_tp_size=get_parallel().attn_tp_size,
attn_cp_size=model_runner.attn_cp_size, attn_cp_size=model_runner.ps.attn_cp_size,
tp_group=model_runner.tp_group, tp_group=model_runner.tp_group,
get_idle_batch=None, get_idle_batch=None,
disable_cuda_graph=model_runner.server_args.disable_cuda_graph, disable_cuda_graph=model_runner.server_args.disable_cuda_graph,
+24 -33
View File
@@ -58,31 +58,20 @@ def init_torch_distributed(
is_draft_worker: bool, is_draft_worker: bool,
local_omp_cpuid: Optional[List[int]], local_omp_cpuid: Optional[List[int]],
): ):
gpu_id = ps.gpu_id
tp_rank = ps.tp_rank
tp_size = ps.tp_size
pp_rank = ps.pp_rank
pp_size = ps.pp_size
dp_size = ps.attn_dp_size
attn_cp_size = ps.attn_cp_size
moe_ep_size = ps.moe_ep_size
moe_dp_size = ps.moe_dp_size
dcp_size = ps.dcp_size
tic = time.perf_counter() tic = time.perf_counter()
logger.info("Init torch distributed begin.") logger.info("Init torch distributed begin.")
try: try:
torch.get_device_module(device).set_device(gpu_id) torch.get_device_module(device).set_device(ps.gpu_id)
except Exception: except Exception:
logger.warning( logger.warning(
f"Context: {device=} {gpu_id=} {os.environ.get('CUDA_VISIBLE_DEVICES')=} {tp_rank=} {tp_size=}" f"Context: {device=} {ps.gpu_id=} {os.environ.get('CUDA_VISIBLE_DEVICES')=} {ps.tp_rank=} {ps.tp_size=}"
) )
raise raise
backend = _resolve_backend(device=device, server_args=server_args, gpu_id=gpu_id) backend = _resolve_backend(device=device, server_args=server_args, gpu_id=ps.gpu_id)
before_avail_memory = get_available_gpu_memory(device, gpu_id) before_avail_memory = get_available_gpu_memory(device, ps.gpu_id)
if not server_args.enable_p2p_check: if not server_args.enable_p2p_check:
monkey_patch_p2p_access_check() monkey_patch_p2p_access_check()
@@ -94,7 +83,7 @@ def init_torch_distributed(
if not is_draft_worker: if not is_draft_worker:
if device == "cpu": if device == "cpu":
_init_cpu_threads_env( _init_cpu_threads_env(
tp_size=tp_size, tp_rank=tp_rank, local_omp_cpuid=local_omp_cpuid tp_size=ps.tp_size, tp_rank=ps.tp_rank, local_omp_cpuid=local_omp_cpuid
) )
# Only initialize the distributed environment on the target model worker. # Only initialize the distributed environment on the target model worker.
@@ -103,28 +92,30 @@ def init_torch_distributed(
dist_init_method=dist_init_method, dist_init_method=dist_init_method,
server_args=server_args, server_args=server_args,
model_config=model_config, model_config=model_config,
gpu_id=gpu_id, gpu_id=ps.gpu_id,
tp_rank=tp_rank, tp_rank=ps.tp_rank,
tp_size=tp_size, tp_size=ps.tp_size,
pp_rank=pp_rank, pp_rank=ps.pp_rank,
pp_size=pp_size, pp_size=ps.pp_size,
dp_size=dp_size, attn_dp_size=ps.attn_dp_size,
attn_cp_size=attn_cp_size, attn_cp_size=ps.attn_cp_size,
moe_ep_size=moe_ep_size, moe_ep_size=ps.moe_ep_size,
moe_dp_size=moe_dp_size, moe_dp_size=ps.moe_dp_size,
dcp_size=dcp_size, dcp_size=ps.dcp_size,
) )
# Pre-warm NCCL/RCCL/HCCL to eliminate cold-start latency in first request # Pre-warm NCCL/RCCL/HCCL to eliminate cold-start latency in first request
# Controlled by --pre-warm-nccl flag (default: enabled on AMD GPUs) # Controlled by --pre-warm-nccl flag (default: enabled on AMD GPUs)
if server_args.pre_warm_nccl and ( if server_args.pre_warm_nccl and (
tp_size > 1 or pp_size > 1 or moe_ep_size > 1 ps.tp_size > 1 or ps.pp_size > 1 or ps.moe_ep_size > 1
): ):
_prewarm_nccl(tp_size=tp_size, pp_size=pp_size, moe_ep_size=moe_ep_size) _prewarm_nccl(
tp_size=ps.tp_size, pp_size=ps.pp_size, moe_ep_size=ps.moe_ep_size
)
pre_model_load_memory = get_available_gpu_memory( pre_model_load_memory = get_available_gpu_memory(
device, device,
gpu_id, ps.gpu_id,
distributed=get_world_group().world_size > 1, distributed=get_world_group().world_size > 1,
cpu_group=get_world_group().cpu_group, cpu_group=get_world_group().cpu_group,
) )
@@ -133,8 +124,8 @@ def init_torch_distributed(
attention_tp_group = get_parallel().attn_tp_group attention_tp_group = get_parallel().attn_tp_group
# Check memory for tensor parallelism # Check memory for tensor parallelism
local_gpu_memory = get_available_gpu_memory(device, gpu_id) local_gpu_memory = get_available_gpu_memory(device, ps.gpu_id)
if tp_size > 1 and not is_draft_worker: if ps.tp_size > 1 and not is_draft_worker:
_check_tp_memory_balance( _check_tp_memory_balance(
pre_model_load_memory=pre_model_load_memory, pre_model_load_memory=pre_model_load_memory,
local_gpu_memory=local_gpu_memory, local_gpu_memory=local_gpu_memory,
@@ -228,7 +219,7 @@ def _init_parallel_groups(
tp_size: int, tp_size: int,
pp_rank: int, pp_rank: int,
pp_size: int, pp_size: int,
dp_size: int, attn_dp_size: int,
attn_cp_size: int, attn_cp_size: int,
moe_ep_size: int, moe_ep_size: int,
moe_dp_size: int, moe_dp_size: int,
@@ -246,7 +237,7 @@ def _init_parallel_groups(
) )
initialize_model_parallel( initialize_model_parallel(
tensor_model_parallel_size=tp_size, tensor_model_parallel_size=tp_size,
attention_data_parallel_size=dp_size, attention_data_parallel_size=attn_dp_size,
pipeline_model_parallel_size=pp_size, pipeline_model_parallel_size=pp_size,
expert_model_parallel_size=moe_ep_size, expert_model_parallel_size=moe_ep_size,
attention_context_model_parallel_size=attn_cp_size, attention_context_model_parallel_size=attn_cp_size,
@@ -37,9 +37,9 @@ class ExpertBackupClient:
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.model_runner = model_runner
self.moe_ep_size = model_runner.moe_ep_size self.moe_ep_size = model_runner.ps.moe_ep_size
self.model_config = model_runner.model_config self.model_config = model_runner.model_config
self.moe_ep_rank = model_runner.moe_ep_rank self.moe_ep_rank = model_runner.ps.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
+2 -2
View File
@@ -107,7 +107,7 @@ class EPLBManager:
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._model_runner.server_args.nnodes,
tp_rank=self._model_runner.tp_rank, tp_rank=self._model_runner.ps.tp_rank,
expert_backup_client=self._model_runner.expert_backup_client, expert_backup_client=self._model_runner.expert_backup_client,
update_weights_from_disk_callable=self._model_runner.weight_updater.update_weights_from_disk, 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, ep_dispatch_algorithm=self._model_runner.server_args.ep_dispatch_algorithm,
@@ -149,7 +149,7 @@ class EPLBManager:
def _should_log_expert_location_metadata(self) -> bool: def _should_log_expert_location_metadata(self) -> bool:
return ( return (
self._model_runner.tp_rank == 0 self._model_runner.ps.tp_rank == 0
and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get() and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get()
) )
@@ -367,7 +367,7 @@ class AscendAttnBackend(AttentionBackend):
self.is_dllm_model = True self.is_dllm_model = True
self.dllm_block_size = self.dllm_config.block_size self.dllm_block_size = self.dllm_config.block_size
self.attn_cp_size = model_runner.attn_cp_size self.attn_cp_size = model_runner.ps.attn_cp_size
def _is_swa_layer(self, layer: RadixAttention) -> bool: def _is_swa_layer(self, layer: RadixAttention) -> bool:
return ( return (
@@ -262,7 +262,7 @@ class FlashAttentionBackend(AttentionBackend):
self.needs_cpu_seq_lens = False self.needs_cpu_seq_lens = False
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
self.skip_prefill = skip_prefill self.skip_prefill = skip_prefill
self.attn_cp_size = model_runner.attn_cp_size self.attn_cp_size = model_runner.ps.attn_cp_size
# Preallocated FULL_MASK tree-mask scratch; lets build_tree_kernel_efficient # Preallocated FULL_MASK tree-mask scratch; lets build_tree_kernel_efficient
# avoid the seq_lens_sum D2H sync (see get_verify_buffers_to_fill_after_draft). # avoid the seq_lens_sum D2H sync (see get_verify_buffers_to_fill_after_draft).
self.cuda_graph_custom_mask = None self.cuda_graph_custom_mask = None
@@ -342,10 +342,10 @@ class FlashAttentionBackend(AttentionBackend):
self.head_dim = model_runner.model_config.head_dim self.head_dim = model_runner.model_config.head_dim
self.num_attention_heads = ( self.num_attention_heads = (
model_runner.model_config.hf_text_config.num_attention_heads model_runner.model_config.hf_text_config.num_attention_heads
// model_runner.tp_size // model_runner.ps.tp_size
) )
self.num_kv_heads = model_runner.model_config.get_num_kv_heads( self.num_kv_heads = model_runner.model_config.get_num_kv_heads(
model_runner.tp_size model_runner.ps.tp_size
) )
_softcapping = getattr( _softcapping = getattr(
model_runner.model_config.hf_text_config, "attn_logit_softcapping", None model_runner.model_config.hf_text_config, "attn_logit_softcapping", None
@@ -38,7 +38,7 @@ class IntelAMXAttnBackend(AttentionBackend):
self.swa_out_cache_loc = None self.swa_out_cache_loc = None
self.num_head = ( self.num_head = (
model_runner.model_config.num_attention_heads // model_runner.tp_size model_runner.model_config.num_attention_heads // model_runner.ps.tp_size
) )
# [NB]: `layer_id` set to 0 for qwen3-next models, as not all attn layers require kv pool # [NB]: `layer_id` set to 0 for qwen3-next models, as not all attn layers require kv pool
@@ -58,7 +58,7 @@ class WaveAttnBackend(AttentionBackend):
import wave_lang.kernel.wave.cache as cache import wave_lang.kernel.wave.cache as cache
base_cache_dir = cache.CACHE_BASE_DIR base_cache_dir = cache.CACHE_BASE_DIR
new_dir = base_cache_dir / f"worker_{model_runner.tp_rank}" new_dir = base_cache_dir / f"worker_{model_runner.ps.tp_rank}"
logger.info(f"Setting Wave cache dir: {new_dir}") logger.info(f"Setting Wave cache dir: {new_dir}")
cache.CACHE_BASE_DIR = new_dir cache.CACHE_BASE_DIR = new_dir
@@ -57,7 +57,7 @@ class XPUAttentionBackend(AttentionBackend):
self.num_attention_heads = ( self.num_attention_heads = (
model_runner.model_config.hf_text_config.num_attention_heads model_runner.model_config.hf_text_config.num_attention_heads
) )
self.tp_size = model_runner.tp_size self.tp_size = model_runner.ps.tp_size
assert self.num_attention_heads % self.tp_size == 0 assert self.num_attention_heads % self.tp_size == 0
self.num_local_heads = self.num_attention_heads // self.tp_size self.num_local_heads = self.num_attention_heads // self.tp_size
self.device = model_runner.device self.device = model_runner.device
@@ -476,8 +476,8 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
logger.info( logger.info(
"PP-parallel DeepGEMM warmup start " "PP-parallel DeepGEMM warmup start "
"(pp_rank=%d, tp_rank=%d, batch_sizes=%s, disagg=%s).", "(pp_rank=%d, tp_rank=%d, batch_sizes=%s, disagg=%s).",
model_runner.pp_rank, model_runner.ps.pp_rank,
model_runner.tp_rank, model_runner.ps.tp_rank,
batch_sizes, batch_sizes,
disagg_mode, disagg_mode,
) )
@@ -505,5 +505,5 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
logger.info( logger.info(
"PP-parallel DeepGEMM warmup done in %.2fs (pp_rank=%d).", "PP-parallel DeepGEMM warmup done in %.2fs (pp_rank=%d).",
time.perf_counter() - t0, time.perf_counter() - t0,
model_runner.pp_rank, model_runner.ps.pp_rank,
) )
@@ -264,21 +264,9 @@ class ModelRunner:
self.memory_pool_config = memory_pool_config self.memory_pool_config = memory_pool_config
self.device = server_args.device self.device = server_args.device
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.tp_rank = ps.tp_rank
self.tp_size = ps.tp_size
self.dcp_size = server_args.dcp_size self.dcp_size = server_args.dcp_size
self.dcp_rank = ps.tp_rank % self.dcp_size self.dcp_rank = ps.tp_rank % self.dcp_size
self.ps = ps self.ps = ps
self.moe_ep_rank = ps.moe_ep_rank
self.moe_ep_size = ps.moe_ep_size
self.dp_rank = ps.dp_rank
self.attn_dp_size = ps.attn_dp_size
self.pp_rank = ps.pp_rank
self.pp_size = ps.pp_size
self.attn_cp_rank = ps.attn_cp_rank
self.attn_cp_size = ps.attn_cp_size
self.moe_dp_rank = ps.moe_dp_rank
self.moe_dp_size = ps.moe_dp_size
self.model_config = model_config self.model_config = model_config
self.dist_port = nccl_port self.dist_port = nccl_port
self.server_args = server_args self.server_args = server_args
@@ -395,7 +383,7 @@ class ModelRunner:
"pp_proxy_tensors" in inspect.signature(self.model.forward).parameters "pp_proxy_tensors" in inspect.signature(self.model.forward).parameters
) )
if self.pp_size > 1: if self.ps.pp_size > 1:
assert ( assert (
self.support_pp self.support_pp
), "Pipeline Parallel is not compatible with this model." ), "Pipeline Parallel is not compatible with this model."
@@ -409,7 +397,7 @@ class ModelRunner:
def init_weight_updater(self): def init_weight_updater(self):
self.weight_updater = WeightUpdater( self.weight_updater = WeightUpdater(
tp_rank=self.tp_rank, tp_rank=self.ps.tp_rank,
device=self.device, device=self.device,
gpu_id=self.gpu_id, gpu_id=self.gpu_id,
model_config=self.model_config, model_config=self.model_config,
@@ -432,8 +420,8 @@ class ModelRunner:
def init_weight_exporter(self): def init_weight_exporter(self):
self.weight_exporter = WeightExporter( self.weight_exporter = WeightExporter(
tp_rank=self.tp_rank, tp_rank=self.ps.tp_rank,
tp_size=self.tp_size, tp_size=self.ps.tp_size,
gpu_id=self.gpu_id, gpu_id=self.gpu_id,
get_model_path=lambda: self.model_config.model_path, get_model_path=lambda: self.model_config.model_path,
get_model=lambda: self.model, get_model=lambda: self.model,
@@ -443,7 +431,7 @@ class ModelRunner:
self.remote_instance_weight_transporter = RemoteInstanceWeightTransporter( self.remote_instance_weight_transporter = RemoteInstanceWeightTransporter(
server_args=self.server_args, server_args=self.server_args,
get_model=lambda: self.model, get_model=lambda: self.model,
tp_rank=self.tp_rank, tp_rank=self.ps.tp_rank,
gpu_id=self.gpu_id, gpu_id=self.gpu_id,
) )
@@ -492,8 +480,8 @@ class ModelRunner:
from sglang.srt.model_executor.mindspore_runner import init_ms_distributed from sglang.srt.model_executor.mindspore_runner import init_ms_distributed
init_ms_distributed( init_ms_distributed(
world_size=self.tp_size * self.pp_size, world_size=self.ps.tp_size * self.ps.pp_size,
rank=self.tp_size * self.pp_rank + self.tp_rank, rank=self.ps.tp_size * self.ps.pp_rank + self.ps.tp_rank,
local_rank=self.gpu_id, local_rank=self.gpu_id,
server_args=self.server_args, server_args=self.server_args,
port=self.dist_port, port=self.dist_port,
@@ -514,10 +502,10 @@ class ModelRunner:
compute_initial_expert_location_metadata( compute_initial_expert_location_metadata(
server_args=server_args, server_args=server_args,
model_config=self.model_config, model_config=self.model_config,
moe_ep_rank=self.moe_ep_rank, moe_ep_rank=self.ps.moe_ep_rank,
) )
) )
if self.tp_rank == 0 and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get(): if self.ps.tp_rank == 0 and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get():
logger.info( logger.info(
"Initial expert_location_metadata:\n%s", "Initial expert_location_metadata:\n%s",
format_expert_location_layout( format_expert_location_layout(
@@ -529,7 +517,7 @@ class ModelRunner:
ExpertDistributionRecorder.init_new( ExpertDistributionRecorder.init_new(
server_args, server_args,
get_global_expert_location_metadata(), get_global_expert_location_metadata(),
rank=self.tp_rank, rank=self.ps.tp_rank,
) )
) )
@@ -557,8 +545,8 @@ class ModelRunner:
model=self.model, model=self.model,
model_config=self.model_config, model_config=self.model_config,
server_args=self.server_args, server_args=self.server_args,
moe_ep_size=self.moe_ep_size, moe_ep_size=self.ps.moe_ep_size,
moe_ep_rank=self.moe_ep_rank, moe_ep_rank=self.ps.moe_ep_rank,
) )
# Must run before backend/graph init so no draft graph records a # Must run before backend/graph init so no draft graph records a
@@ -600,7 +588,7 @@ class ModelRunner:
# Apply torch TP if the model supports it # 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.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 # Init lora
@@ -618,8 +606,8 @@ class ModelRunner:
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,
pp_size=self.pp_size, pp_size=self.ps.pp_size,
pp_rank=self.pp_rank, pp_rank=self.ps.pp_rank,
start_layer=self.layer_info.start_layer, start_layer=self.layer_info.start_layer,
) )
@@ -864,9 +852,9 @@ class ModelRunner:
def check_quantized_moe_compatibility(self): def check_quantized_moe_compatibility(self):
check_quantized_moe_compatibility( check_quantized_moe_compatibility(
model_config=self.model_config, model_config=self.model_config,
tp_size=self.tp_size, tp_size=self.ps.tp_size,
moe_ep_size=self.moe_ep_size, moe_ep_size=self.ps.moe_ep_size,
moe_dp_size=self.moe_dp_size, moe_dp_size=self.ps.moe_dp_size,
) )
def init_torch_distributed(self): def init_torch_distributed(self):
@@ -908,18 +896,18 @@ class ModelRunner:
self.load_config = build_load_config( self.load_config = build_load_config(
server_args=self.server_args, server_args=self.server_args,
tp_rank=self.tp_rank, tp_rank=self.ps.tp_rank,
remote_instance_weight_transporter_engine=self.remote_instance_weight_transporter.engine, remote_instance_weight_transporter_engine=self.remote_instance_weight_transporter.engine,
remote_instance_weight_transporter_session_id=self.remote_instance_weight_transporter.session_id, remote_instance_weight_transporter_session_id=self.remote_instance_weight_transporter.session_id,
draft_model_idx=self.draft_model_idx, draft_model_idx=self.draft_model_idx,
) )
if self.device == "cpu": if self.device == "cpu":
self.model_config = adjust_config_with_unaligned_cpu_tp( self.model_config = adjust_config_with_unaligned_cpu_tp(
self.model_config, self.load_config, self.tp_size self.model_config, self.load_config, self.ps.tp_size
) )
maybe_trigger_remote_instance_nccl_send_group( maybe_trigger_remote_instance_nccl_send_group(
server_args=self.server_args, tp_rank=self.tp_rank server_args=self.server_args, tp_rank=self.ps.tp_rank
) )
loaded = load_model_with_memory_saver( loaded = load_model_with_memory_saver(
@@ -981,9 +969,9 @@ class ModelRunner:
server_args=self.server_args, server_args=self.server_args,
spec_algorithm=self.spec_algorithm, spec_algorithm=self.spec_algorithm,
is_draft_worker=self.is_draft_worker, is_draft_worker=self.is_draft_worker,
tp_size=self.tp_size, tp_size=self.ps.tp_size,
tp_rank=self.tp_rank, tp_rank=self.ps.tp_rank,
pp_rank=self.pp_rank, pp_rank=self.ps.pp_rank,
) )
if dumper.may_enable: if dumper.may_enable:
@@ -1000,7 +988,7 @@ class ModelRunner:
dist_barrier_after_load( dist_barrier_after_load(
elastic_ep_backend=self.server_args.elastic_ep_backend, elastic_ep_backend=self.server_args.elastic_ep_backend,
tp_rank=self.tp_rank, tp_rank=self.ps.tp_rank,
) )
def maybe_recover_ep_ranks(self): def maybe_recover_ep_ranks(self):
@@ -1053,8 +1041,8 @@ class ModelRunner:
dtype=self.dtype, dtype=self.dtype,
server_args=self.server_args, server_args=self.server_args,
lora_backend=self.server_args.lora_backend, lora_backend=self.server_args.lora_backend,
tp_size=self.tp_size, tp_size=self.ps.tp_size,
tp_rank=self.tp_rank, tp_rank=self.ps.tp_rank,
max_lora_rank=self.server_args.max_lora_rank, max_lora_rank=self.server_args.max_lora_rank,
target_modules=self.server_args.lora_target_modules, target_modules=self.server_args.lora_target_modules,
lora_paths=self.server_args.lora_paths, lora_paths=self.server_args.lora_paths,
@@ -1401,12 +1389,12 @@ class ModelRunner:
def init_threads_binding(self): def init_threads_binding(self):
self.local_omp_cpuid = numa_utils.init_threads_binding( self.local_omp_cpuid = numa_utils.init_threads_binding(
tp_rank=self.tp_rank, tp_size=self.tp_size tp_rank=self.ps.tp_rank, tp_size=self.ps.tp_size
) )
def apply_torch_tp(self): def apply_torch_tp(self):
model_parallel.apply_torch_tp( model_parallel.apply_torch_tp(
model=self.model, device=self.device, tp_size=self.tp_size model=self.model, device=self.device, tp_size=self.ps.tp_size
) )
def update_decode_attn_backend(self, stream_idx: int): def update_decode_attn_backend(self, stream_idx: int):
@@ -1523,7 +1511,7 @@ class ModelRunner:
if self.msprobe_debugger is not None: if self.msprobe_debugger is not None:
rank_id = ( rank_id = (
self.gpu_id self.gpu_id
if self.attn_dp_size is not None and self.attn_dp_size > 1 if self.ps.attn_dp_size is not None and self.ps.attn_dp_size > 1
else None else None
) )
self.msprobe_debugger.start(model=self.model, rank_id=rank_id) self.msprobe_debugger.start(model=self.model, rank_id=rank_id)
@@ -229,7 +229,7 @@ class BaseRunner(ABC):
if ( if (
envs.SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP.get() envs.SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP.get()
and deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM and deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM
and mr.pp_size > 1 and mr.ps.pp_size > 1
and not mr.spec_algorithm.is_speculative() and not mr.spec_algorithm.is_speculative()
): ):
from sglang.srt.layers.deep_gemm_wrapper.compile_utils import ( from sglang.srt.layers.deep_gemm_wrapper.compile_utils import (
@@ -458,10 +458,10 @@ class BaseRunner(ABC):
pp_hidden_tokens = num_tokens pp_hidden_tokens = num_tokens
if ( if (
capture_forward_mode == ForwardMode.EXTEND capture_forward_mode == ForwardMode.EXTEND
and mr.pp_rank != 0 and mr.ps.pp_rank != 0
and mr.attn_cp_size > 1 and mr.ps.attn_cp_size > 1
): ):
pp_hidden_tokens = num_tokens // mr.attn_cp_size pp_hidden_tokens = num_tokens // mr.ps.attn_cp_size
pp_proxy_tensors = PPProxyTensors( pp_proxy_tensors = PPProxyTensors(
{k: v[:pp_hidden_tokens] for k, v in buffers.pp_proxy_tensors.items()} {k: v[:pp_hidden_tokens] for k, v in buffers.pp_proxy_tensors.items()}
) )
@@ -125,10 +125,10 @@ def flashinfer_autotune_cache_path(model_runner: ModelRunner) -> Path:
str(mr.dtype), str(mr.dtype),
str(server_args.quantization), str(server_args.quantization),
str(server_args.moe_runner_backend), str(server_args.moe_runner_backend),
str(mr.tp_size), str(mr.ps.tp_size),
str(mr.pp_size), str(mr.ps.pp_size),
str(mr.dp_size), str(mr.ps.attn_dp_size),
str(mr.moe_ep_size), str(mr.ps.moe_ep_size),
str(mr.model_config.hf_config.__class__.__name__), str(mr.model_config.hf_config.__class__.__name__),
] ]
if mr.is_draft_worker: if mr.is_draft_worker:
@@ -144,7 +144,10 @@ def flashinfer_autotune_cache_path(model_runner: ModelRunner) -> Path:
/ cache_key / cache_key
) )
cache_dir.mkdir(parents=True, exist_ok=True) cache_dir.mkdir(parents=True, exist_ok=True)
return cache_dir / f"rank_tp{mr.tp_rank}_pp{mr.pp_rank}_dp{mr.dp_rank or 0}.json" return (
cache_dir
/ f"rank_tp{mr.ps.tp_rank}_pp{mr.ps.pp_rank}_dp{mr.ps.dp_rank or 0}.json"
)
@contextlib.contextmanager @contextlib.contextmanager
@@ -301,7 +301,7 @@ class DFlashWorkerV2(BaseSpecWorker):
def _maybe_build_draft_sampler(self): def _maybe_build_draft_sampler(self):
def _eager(reason): def _eager(reason):
if self.tp_rank == 0: if self.ps.tp_rank == 0:
logger.info("DFLASH draft greedy head kept eager (reason=%s).", reason) logger.info("DFLASH draft greedy head kept eager (reason=%s).", reason)
return None return None
@@ -325,7 +325,7 @@ class DFlashWorkerV2(BaseSpecWorker):
return _eager("added vocab") return _eager("added vocab")
num_org = int(shard.num_org_elements) num_org = int(shard.num_org_elements)
org_vocab_start = int(shard.org_vocab_start_index) org_vocab_start = int(shard.org_vocab_start_index)
if self.tp_rank == 0: if self.ps.tp_rank == 0:
logger.info("DFLASH draft greedy head folded into the draft cuda graph.") logger.info("DFLASH draft greedy head folded into the draft cuda graph.")
return _DflashDraftSampler( return _DflashDraftSampler(
weight=lm_head.weight, weight=lm_head.weight,
@@ -347,7 +347,7 @@ class DFlashWorkerV2(BaseSpecWorker):
fused_disable_reason = "draft model does not support fused context KV" fused_disable_reason = "draft model does not support fused context KV"
if fused_disable_reason is not None: if fused_disable_reason is not None:
if self.tp_rank == 0: if self.ps.tp_rank == 0:
logger.info( logger.info(
"DFLASH fused KV materialization disabled: %s", "DFLASH fused KV materialization disabled: %s",
fused_disable_reason, fused_disable_reason,
@@ -106,8 +106,8 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
# Fields the parent's capture() reads: # Fields the parent's capture() reads:
self.device = model_runner.device self.device = model_runner.device
self.device_module = torch.get_device_module(self.device) self.device_module = torch.get_device_module(self.device)
self.tp_size = model_runner.tp_size self.tp_size = model_runner.ps.tp_size
self.attn_dp_size = model_runner.attn_dp_size self.attn_dp_size = model_runner.ps.attn_dp_size
self.pp_size = model_runner.server_args.pp_size self.pp_size = model_runner.server_args.pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
@@ -95,8 +95,8 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
# Fields the parent's capture() reads: # Fields the parent's capture() reads:
self.device = model_runner.device self.device = model_runner.device
self.device_module = torch.get_device_module(self.device) self.device_module = torch.get_device_module(self.device)
self.tp_size = model_runner.tp_size self.tp_size = model_runner.ps.tp_size
self.attn_dp_size = model_runner.attn_dp_size self.attn_dp_size = model_runner.ps.attn_dp_size
self.pp_size = model_runner.server_args.pp_size self.pp_size = model_runner.server_args.pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
@@ -91,8 +91,8 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args) self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
self.require_mlp_sync = require_mlp_sync(model_runner.server_args) self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args) self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.tp_size = self.model_runner.tp_size self.tp_size = self.model_runner.ps.tp_size
self.attn_dp_size = self.model_runner.attn_dp_size self.attn_dp_size = self.model_runner.ps.attn_dp_size
self.pp_size = model_runner.server_args.pp_size self.pp_size = model_runner.server_args.pp_size
self.speculative_num_steps = model_runner.server_args.speculative_num_steps self.speculative_num_steps = model_runner.server_args.speculative_num_steps
self.topk = model_runner.server_args.speculative_eagle_topk self.topk = model_runner.server_args.speculative_eagle_topk
@@ -22,10 +22,12 @@ start of the next draft.
from __future__ import annotations from __future__ import annotations
import logging import logging
from dataclasses import replace
from typing import Optional from typing import Optional
import torch import torch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.moe.utils import ( from sglang.srt.layers.moe.utils import (
speculative_moe_a2a_backend_context, speculative_moe_a2a_backend_context,
speculative_moe_backend_context, speculative_moe_backend_context,
@@ -90,11 +92,7 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
tp_rank: int, ps: ParallelState,
dp_rank: Optional[int],
moe_ep_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -102,6 +100,7 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
self.topk = server_args.speculative_eagle_topk self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.ps = ps
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.device = server_args.device self.device = server_args.device
self.target_worker = target_worker self.target_worker = target_worker
@@ -132,12 +131,8 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
self, self,
server_args=server_args, server_args=server_args,
gpu_id=gpu_id, gpu_id=gpu_id,
tp_rank=tp_rank, # spec workers don't support pipeline parallelism
pp_rank=0, ps=replace(ps, pp_rank=0),
dp_rank=dp_rank,
moe_ep_rank=moe_ep_rank,
attn_cp_rank=attn_cp_rank,
moe_dp_rank=moe_dp_rank,
nccl_port=nccl_port, nccl_port=nccl_port,
is_draft_worker=True, is_draft_worker=True,
) )
@@ -649,11 +644,7 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
tp_rank: int, ps: ParallelState,
dp_rank: Optional[int],
moe_ep_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -664,7 +655,7 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
self.topk = server_args.speculative_eagle_topk self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.tp_rank = tp_rank self.ps = ps
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.device = server_args.device self.device = server_args.device
self._target_worker = target_worker self._target_worker = target_worker
@@ -685,11 +676,7 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
self._draft_worker = FrozenKVMTPDraftWorker( self._draft_worker = FrozenKVMTPDraftWorker(
server_args, server_args,
gpu_id, gpu_id,
tp_rank, ps,
dp_rank,
moe_ep_rank,
attn_cp_rank,
moe_dp_rank,
nccl_port, nccl_port,
target_worker, target_worker,
) )
@@ -125,7 +125,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
# Fields the parent's capture() reads: # Fields the parent's capture() reads:
self.device = model_runner.device self.device = model_runner.device
self.device_module = torch.get_device_module(self.device) self.device_module = torch.get_device_module(self.device)
self.tp_size = model_runner.tp_size self.tp_size = model_runner.ps.tp_size
self.dp_size = model_runner.server_args.dp_size self.dp_size = model_runner.server_args.dp_size
self.pp_size = model_runner.server_args.pp_size self.pp_size = model_runner.server_args.pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
+7 -7
View File
@@ -157,14 +157,14 @@ class WeightChecker:
return info.model_dump() return info.model_dump()
def _parallelism_info(self) -> ParallelismInfo: def _parallelism_info(self) -> ParallelismInfo:
mr = self._model_runner ps = self._model_runner.ps
return ParallelismInfo( return ParallelismInfo(
tp_rank=mr.tp_rank, tp_rank=ps.tp_rank,
tp_size=mr.tp_size, tp_size=ps.tp_size,
dp_rank=mr.dp_rank if mr.dp_rank is not None else 0, dp_rank=ps.dp_rank if ps.dp_rank is not None else 0,
dp_size=mr.dp_size, dp_size=ps.attn_dp_size,
pp_rank=mr.pp_rank, pp_rank=ps.pp_rank,
pp_size=mr.pp_size, pp_size=ps.pp_size,
rank=dist.get_rank() if dist.is_initialized() else 0, rank=dist.get_rank() if dist.is_initialized() else 0,
size=dist.get_world_size() if dist.is_initialized() else 1, size=dist.get_world_size() if dist.is_initialized() else 1,
) )