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,
dp_size=model_runner.server_args.dp_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,
get_idle_batch=None,
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,
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()
logger.info("Init torch distributed begin.")
try:
torch.get_device_module(device).set_device(gpu_id)
torch.get_device_module(device).set_device(ps.gpu_id)
except Exception:
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
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:
monkey_patch_p2p_access_check()
@@ -94,7 +83,7 @@ def init_torch_distributed(
if not is_draft_worker:
if device == "cpu":
_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.
@@ -103,28 +92,30 @@ def init_torch_distributed(
dist_init_method=dist_init_method,
server_args=server_args,
model_config=model_config,
gpu_id=gpu_id,
tp_rank=tp_rank,
tp_size=tp_size,
pp_rank=pp_rank,
pp_size=pp_size,
dp_size=dp_size,
attn_cp_size=attn_cp_size,
moe_ep_size=moe_ep_size,
moe_dp_size=moe_dp_size,
dcp_size=dcp_size,
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,
attn_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,
)
# Pre-warm NCCL/RCCL/HCCL to eliminate cold-start latency in first request
# Controlled by --pre-warm-nccl flag (default: enabled on AMD GPUs)
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(
device,
gpu_id,
ps.gpu_id,
distributed=get_world_group().world_size > 1,
cpu_group=get_world_group().cpu_group,
)
@@ -133,8 +124,8 @@ def init_torch_distributed(
attention_tp_group = get_parallel().attn_tp_group
# Check memory for tensor parallelism
local_gpu_memory = get_available_gpu_memory(device, gpu_id)
if tp_size > 1 and not is_draft_worker:
local_gpu_memory = get_available_gpu_memory(device, ps.gpu_id)
if ps.tp_size > 1 and not is_draft_worker:
_check_tp_memory_balance(
pre_model_load_memory=pre_model_load_memory,
local_gpu_memory=local_gpu_memory,
@@ -228,7 +219,7 @@ def _init_parallel_groups(
tp_size: int,
pp_rank: int,
pp_size: int,
dp_size: int,
attn_dp_size: int,
attn_cp_size: int,
moe_ep_size: int,
moe_dp_size: int,
@@ -246,7 +237,7 @@ def _init_parallel_groups(
)
initialize_model_parallel(
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,
expert_model_parallel_size=moe_ep_size,
attention_context_model_parallel_size=attn_cp_size,
@@ -37,9 +37,9 @@ class ExpertBackupClient:
self.recv_list = [None] * self.engine_num
self.ready_sockets = [None] * self.engine_num
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.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.session_id_list = [None] * self.engine_num
self.transfer_engine = None
+2 -2
View File
@@ -107,7 +107,7 @@ class EPLBManager:
new_expert_location_metadata=expert_location_metadata,
update_layer_ids=chunk_layer_ids,
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,
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,
@@ -149,7 +149,7 @@ class EPLBManager:
def _should_log_expert_location_metadata(self) -> bool:
return (
self._model_runner.tp_rank == 0
self._model_runner.ps.tp_rank == 0
and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get()
)
@@ -367,7 +367,7 @@ class AscendAttnBackend(AttentionBackend):
self.is_dllm_model = True
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:
return (
@@ -262,7 +262,7 @@ class FlashAttentionBackend(AttentionBackend):
self.needs_cpu_seq_lens = False
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
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
# avoid the seq_lens_sum D2H sync (see get_verify_buffers_to_fill_after_draft).
self.cuda_graph_custom_mask = None
@@ -342,10 +342,10 @@ class FlashAttentionBackend(AttentionBackend):
self.head_dim = model_runner.model_config.head_dim
self.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(
model_runner.tp_size
model_runner.ps.tp_size
)
_softcapping = getattr(
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.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
@@ -58,7 +58,7 @@ class WaveAttnBackend(AttentionBackend):
import wave_lang.kernel.wave.cache as cache
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}")
cache.CACHE_BASE_DIR = new_dir
@@ -57,7 +57,7 @@ class XPUAttentionBackend(AttentionBackend):
self.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
self.num_local_heads = self.num_attention_heads // self.tp_size
self.device = model_runner.device
@@ -476,8 +476,8 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
logger.info(
"PP-parallel DeepGEMM warmup start "
"(pp_rank=%d, tp_rank=%d, batch_sizes=%s, disagg=%s).",
model_runner.pp_rank,
model_runner.tp_rank,
model_runner.ps.pp_rank,
model_runner.ps.tp_rank,
batch_sizes,
disagg_mode,
)
@@ -505,5 +505,5 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
logger.info(
"PP-parallel DeepGEMM warmup done in %.2fs (pp_rank=%d).",
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.device = server_args.device
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_rank = ps.tp_rank % self.dcp_size
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.dist_port = nccl_port
self.server_args = server_args
@@ -395,7 +383,7 @@ class ModelRunner:
"pp_proxy_tensors" in inspect.signature(self.model.forward).parameters
)
if self.pp_size > 1:
if self.ps.pp_size > 1:
assert (
self.support_pp
), "Pipeline Parallel is not compatible with this model."
@@ -409,7 +397,7 @@ class ModelRunner:
def init_weight_updater(self):
self.weight_updater = WeightUpdater(
tp_rank=self.tp_rank,
tp_rank=self.ps.tp_rank,
device=self.device,
gpu_id=self.gpu_id,
model_config=self.model_config,
@@ -432,8 +420,8 @@ class ModelRunner:
def init_weight_exporter(self):
self.weight_exporter = WeightExporter(
tp_rank=self.tp_rank,
tp_size=self.tp_size,
tp_rank=self.ps.tp_rank,
tp_size=self.ps.tp_size,
gpu_id=self.gpu_id,
get_model_path=lambda: self.model_config.model_path,
get_model=lambda: self.model,
@@ -443,7 +431,7 @@ class ModelRunner:
self.remote_instance_weight_transporter = RemoteInstanceWeightTransporter(
server_args=self.server_args,
get_model=lambda: self.model,
tp_rank=self.tp_rank,
tp_rank=self.ps.tp_rank,
gpu_id=self.gpu_id,
)
@@ -492,8 +480,8 @@ class ModelRunner:
from sglang.srt.model_executor.mindspore_runner import init_ms_distributed
init_ms_distributed(
world_size=self.tp_size * self.pp_size,
rank=self.tp_size * self.pp_rank + self.tp_rank,
world_size=self.ps.tp_size * self.ps.pp_size,
rank=self.ps.tp_size * self.ps.pp_rank + self.ps.tp_rank,
local_rank=self.gpu_id,
server_args=self.server_args,
port=self.dist_port,
@@ -514,10 +502,10 @@ class ModelRunner:
compute_initial_expert_location_metadata(
server_args=server_args,
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(
"Initial expert_location_metadata:\n%s",
format_expert_location_layout(
@@ -529,7 +517,7 @@ class ModelRunner:
ExpertDistributionRecorder.init_new(
server_args,
get_global_expert_location_metadata(),
rank=self.tp_rank,
rank=self.ps.tp_rank,
)
)
@@ -557,8 +545,8 @@ class ModelRunner:
model=self.model,
model_config=self.model_config,
server_args=self.server_args,
moe_ep_size=self.moe_ep_size,
moe_ep_rank=self.moe_ep_rank,
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
@@ -600,7 +588,7 @@ class ModelRunner:
# Apply torch TP if the model supports it
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()
# Init lora
@@ -618,8 +606,8 @@ class ModelRunner:
def get_pp_proxy_topk_size(self) -> Optional[int]:
return misc_utils.resolve_pp_proxy_topk_size(
model_config=self.model_config,
pp_size=self.pp_size,
pp_rank=self.pp_rank,
pp_size=self.ps.pp_size,
pp_rank=self.ps.pp_rank,
start_layer=self.layer_info.start_layer,
)
@@ -864,9 +852,9 @@ class ModelRunner:
def check_quantized_moe_compatibility(self):
check_quantized_moe_compatibility(
model_config=self.model_config,
tp_size=self.tp_size,
moe_ep_size=self.moe_ep_size,
moe_dp_size=self.moe_dp_size,
tp_size=self.ps.tp_size,
moe_ep_size=self.ps.moe_ep_size,
moe_dp_size=self.ps.moe_dp_size,
)
def init_torch_distributed(self):
@@ -908,18 +896,18 @@ class ModelRunner:
self.load_config = build_load_config(
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_session_id=self.remote_instance_weight_transporter.session_id,
draft_model_idx=self.draft_model_idx,
)
if self.device == "cpu":
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(
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(
@@ -981,9 +969,9 @@ class ModelRunner:
server_args=self.server_args,
spec_algorithm=self.spec_algorithm,
is_draft_worker=self.is_draft_worker,
tp_size=self.tp_size,
tp_rank=self.tp_rank,
pp_rank=self.pp_rank,
tp_size=self.ps.tp_size,
tp_rank=self.ps.tp_rank,
pp_rank=self.ps.pp_rank,
)
if dumper.may_enable:
@@ -1000,7 +988,7 @@ class ModelRunner:
dist_barrier_after_load(
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):
@@ -1053,8 +1041,8 @@ class ModelRunner:
dtype=self.dtype,
server_args=self.server_args,
lora_backend=self.server_args.lora_backend,
tp_size=self.tp_size,
tp_rank=self.tp_rank,
tp_size=self.ps.tp_size,
tp_rank=self.ps.tp_rank,
max_lora_rank=self.server_args.max_lora_rank,
target_modules=self.server_args.lora_target_modules,
lora_paths=self.server_args.lora_paths,
@@ -1401,12 +1389,12 @@ class ModelRunner:
def init_threads_binding(self):
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):
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):
@@ -1523,7 +1511,7 @@ class ModelRunner:
if self.msprobe_debugger is not None:
rank_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
)
self.msprobe_debugger.start(model=self.model, rank_id=rank_id)
@@ -229,7 +229,7 @@ class BaseRunner(ABC):
if (
envs.SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP.get()
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()
):
from sglang.srt.layers.deep_gemm_wrapper.compile_utils import (
@@ -458,10 +458,10 @@ class BaseRunner(ABC):
pp_hidden_tokens = num_tokens
if (
capture_forward_mode == ForwardMode.EXTEND
and mr.pp_rank != 0
and mr.attn_cp_size > 1
and mr.ps.pp_rank != 0
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(
{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(server_args.quantization),
str(server_args.moe_runner_backend),
str(mr.tp_size),
str(mr.pp_size),
str(mr.dp_size),
str(mr.moe_ep_size),
str(mr.ps.tp_size),
str(mr.ps.pp_size),
str(mr.ps.attn_dp_size),
str(mr.ps.moe_ep_size),
str(mr.model_config.hf_config.__class__.__name__),
]
if mr.is_draft_worker:
@@ -144,7 +144,10 @@ def flashinfer_autotune_cache_path(model_runner: ModelRunner) -> Path:
/ cache_key
)
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
@@ -301,7 +301,7 @@ class DFlashWorkerV2(BaseSpecWorker):
def _maybe_build_draft_sampler(self):
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)
return None
@@ -325,7 +325,7 @@ class DFlashWorkerV2(BaseSpecWorker):
return _eager("added vocab")
num_org = int(shard.num_org_elements)
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.")
return _DflashDraftSampler(
weight=lm_head.weight,
@@ -347,7 +347,7 @@ class DFlashWorkerV2(BaseSpecWorker):
fused_disable_reason = "draft model does not support fused context KV"
if fused_disable_reason is not None:
if self.tp_rank == 0:
if self.ps.tp_rank == 0:
logger.info(
"DFLASH fused KV materialization disabled: %s",
fused_disable_reason,
@@ -106,8 +106,8 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
# Fields the parent's capture() reads:
self.device = model_runner.device
self.device_module = torch.get_device_module(self.device)
self.tp_size = model_runner.tp_size
self.attn_dp_size = model_runner.attn_dp_size
self.tp_size = model_runner.ps.tp_size
self.attn_dp_size = model_runner.ps.attn_dp_size
self.pp_size = model_runner.server_args.pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
@@ -95,8 +95,8 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
# Fields the parent's capture() reads:
self.device = model_runner.device
self.device_module = torch.get_device_module(self.device)
self.tp_size = model_runner.tp_size
self.attn_dp_size = model_runner.attn_dp_size
self.tp_size = model_runner.ps.tp_size
self.attn_dp_size = model_runner.ps.attn_dp_size
self.pp_size = model_runner.server_args.pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile
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_sync = require_mlp_sync(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.attn_dp_size = self.model_runner.attn_dp_size
self.tp_size = self.model_runner.ps.tp_size
self.attn_dp_size = self.model_runner.ps.attn_dp_size
self.pp_size = model_runner.server_args.pp_size
self.speculative_num_steps = model_runner.server_args.speculative_num_steps
self.topk = model_runner.server_args.speculative_eagle_topk
@@ -22,10 +22,12 @@ start of the next draft.
from __future__ import annotations
import logging
from dataclasses import replace
from typing import Optional
import torch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.moe.utils import (
speculative_moe_a2a_backend_context,
speculative_moe_backend_context,
@@ -90,11 +92,7 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
self,
server_args: ServerArgs,
gpu_id: int,
tp_rank: int,
dp_rank: Optional[int],
moe_ep_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
ps: ParallelState,
nccl_port: int,
target_worker: TpModelWorker,
):
@@ -102,6 +100,7 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.ps = ps
self.gpu_id = gpu_id
self.device = server_args.device
self.target_worker = target_worker
@@ -132,12 +131,8 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
self,
server_args=server_args,
gpu_id=gpu_id,
tp_rank=tp_rank,
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,
# spec workers don't support pipeline parallelism
ps=replace(ps, pp_rank=0),
nccl_port=nccl_port,
is_draft_worker=True,
)
@@ -649,11 +644,7 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
self,
server_args: ServerArgs,
gpu_id: int,
tp_rank: int,
dp_rank: Optional[int],
moe_ep_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
ps: ParallelState,
nccl_port: int,
target_worker: TpModelWorker,
):
@@ -664,7 +655,7 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
self.topk = server_args.speculative_eagle_topk
self.speculative_num_steps = server_args.speculative_num_steps
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.device = server_args.device
self._target_worker = target_worker
@@ -685,11 +676,7 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
self._draft_worker = FrozenKVMTPDraftWorker(
server_args,
gpu_id,
tp_rank,
dp_rank,
moe_ep_rank,
attn_cp_rank,
moe_dp_rank,
ps,
nccl_port,
target_worker,
)
@@ -125,7 +125,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
# Fields the parent's capture() reads:
self.device = model_runner.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.pp_size = model_runner.server_args.pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile
+7 -7
View File
@@ -157,14 +157,14 @@ class WeightChecker:
return info.model_dump()
def _parallelism_info(self) -> ParallelismInfo:
mr = self._model_runner
ps = self._model_runner.ps
return ParallelismInfo(
tp_rank=mr.tp_rank,
tp_size=mr.tp_size,
dp_rank=mr.dp_rank if mr.dp_rank is not None else 0,
dp_size=mr.dp_size,
pp_rank=mr.pp_rank,
pp_size=mr.pp_size,
tp_rank=ps.tp_rank,
tp_size=ps.tp_size,
dp_rank=ps.dp_rank if ps.dp_rank is not None else 0,
dp_size=ps.attn_dp_size,
pp_rank=ps.pp_rank,
pp_size=ps.pp_size,
rank=dist.get_rank() if dist.is_initialized() else 0,
size=dist.get_world_size() if dist.is_initialized() else 1,
)