Retire the per-runner parallel record (#40343)

This commit is contained in:
Cheng Wan
2026-09-21 12:26:40 -07:00
committed by GitHub
parent 73f071db52
commit 970e946e4f
79 changed files with 395 additions and 463 deletions
+5 -28
View File
@@ -75,7 +75,6 @@ from sglang.srt.distributed.parallel_state import (
destroy_distributed_environment,
destroy_model_parallel,
)
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.entrypoints.engine import _set_envs_and_config
from sglang.srt.hardware_backend.mlx.runtime import use_mlx
from sglang.srt.layers.dp_attention import compute_dp_attention_world_info
@@ -329,32 +328,10 @@ def load_model(server_args, port_args, gpu_id, tp_rank):
cfg.attn_cp_size,
)
)
ps = ParallelState(
tp_rank=tp_rank,
tp_size=cfg.tp_size,
pp_rank=0,
pp_size=1,
dp_rank=None,
dp_size=cfg.dp_size,
attn_tp_rank=attn_tp_rank,
attn_tp_size=attn_tp_size,
attn_cp_rank=0,
attn_cp_size=cfg.attn_cp_size,
attn_dcp_rank=tp_rank % cfg.dcp_size,
attn_dcp_size=cfg.dcp_size,
attn_dp_rank=attn_dp_rank,
attn_dp_size=attn_dp_size,
moe_ep_rank=moe_ep_rank,
moe_ep_size=cfg.ep_size,
moe_dp_rank=None,
moe_dp_size=cfg.moe_dp_size,
gpu_id=gpu_id,
)
runner_kwargs = dict(
model_config=model_config,
mem_fraction_static=cfg.mem_fraction_static,
gpu_id=gpu_id,
ps=ps,
nccl_port=port_args.nccl_port,
server_args=server_args,
)
@@ -571,7 +548,7 @@ def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
model_runner=model_runner,
dp_size=get_parallel().dp_size,
attn_tp_size=get_parallel().attn_tp_size,
attn_cp_size=model_runner.ps.attn_cp_size,
attn_cp_size=model_runner.attn_cp_size,
tp_group=model_runner.tp_group,
get_idle_batch=None,
disable_cuda_graph=cuda_graph_fully_disabled(),
@@ -709,13 +686,12 @@ def correctness_test(
gpu_id,
tp_rank,
):
# With the placement this process was spawned with, so a rank read here
# does not need a process group -- the same bundle the runner is handed.
publish(
server_args,
role="scheduler",
ranks=SpawnRanks(
world_rank=spawn_world_rank(server_args, tp_rank=tp_rank, pp_rank=0)
world_rank=spawn_world_rank(server_args, tp_rank=tp_rank, pp_rank=0),
gpu_id=gpu_id,
),
)
@@ -926,7 +902,8 @@ def latency_test(
server_args,
role="scheduler",
ranks=SpawnRanks(
world_rank=spawn_world_rank(server_args, tp_rank=tp_rank, pp_rank=0)
world_rank=spawn_world_rank(server_args, tp_rank=tp_rank, pp_rank=0),
gpu_id=gpu_id,
),
)
initialize_moe_config()
@@ -92,6 +92,9 @@ class Arg(msgspec.Struct, frozen=True):
fallback: Any = None
_NO_DEFAULT = object()
class Derived(msgspec.Struct, frozen=True):
"""Metadata for a field the configuration implies, not one anyone types.
@@ -122,6 +125,12 @@ class Derived(msgspec.Struct, frozen=True):
doc: str = ""
fn: str = ""
# For a declaration with no ``fn`` whose absence is itself an answer:
# ``gpu_id`` is ``None`` in a process that runs on no device, and a reader
# wants that rather than an error. A rank has no such value -- the wrong
# one is a hang in a collective -- so it carries no default and a read
# before the write says so.
default: Any = _NO_DEFAULT
class NS(msgspec.Struct, frozen=True):
+12 -1
View File
@@ -17,7 +17,7 @@ from typing import (
import msgspec
from sglang.srt.arg_groups.arg_utils import A
from sglang.srt.arg_groups.arg_utils import A, Derived
class Device(msgspec.Struct):
@@ -40,6 +40,17 @@ class Device(msgspec.Struct):
int,
"The delta between consecutive GPU IDs that are used. For example, setting it to 2 will use GPU 0,2,4,...",
] = 1
gpu_id = Derived(
doc=(
"Which device this process runs on. Nobody types it and nothing "
"computes it from the configuration: the parent decides -- "
"reindexing narrows the visible devices before the spawn, and Ray "
"allocates from its own pool -- so the entry states it in the "
"bundle it hands `publish`. `None` is an answer rather than an "
"absence: most roles run on no device at all."
),
default=None,
)
random_seed: A[Optional[int], "The random seed."] = None
mlx_enable_sampling: A[
bool,
+4 -3
View File
@@ -111,6 +111,7 @@ from sglang.srt.observability.scheduler_stage_metrics import (
scheduler_stage_method,
)
from sglang.srt.runtime_context import (
get_device,
get_disagg,
get_memory,
get_parallel,
@@ -439,7 +440,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
if get_disagg().disaggregation_enable_kv_checksum:
kv_args = self.kv_manager.kv_args
self.scheduler.kv_checksum_computer = KvChecksumComputer(
device=torch.device(f"cuda:{self.scheduler.ps.gpu_id}"),
device=torch.device(f"cuda:{get_device().gpu_id}"),
kv_data_ptrs=kv_args.kv_data_ptrs,
kv_item_lens=kv_args.kv_item_lens,
state_data_ptrs=kv_args.state_data_ptrs,
@@ -565,7 +566,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
kv_args.engine_rank = self.tp_rank % (attn_tp_size)
kv_args.pp_rank = self.pp_rank
kv_args.system_dp_rank = self.scheduler.ps.dp_rank
kv_args.system_dp_rank = get_parallel().dp_rank
kv_args.kv_cache_dtype_str = (
self.scheduler.tp_worker.model_runner.kv_cache_dtype_str
)
@@ -633,7 +634,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
)
kv_args.ib_device = get_disagg().disaggregation_ib_device
kv_args.gpu_id = self.scheduler.ps.gpu_id
kv_args.gpu_id = get_device().gpu_id
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
kv_manager = kv_manager_class(
kv_args,
@@ -40,6 +40,7 @@ from sglang.srt.managers.schedule_batch import Modality, Req
from sglang.srt.multimodal.cache import media_preprocess_kwargs
from sglang.srt.multimodal.transport import determine_tensor_transport_mode
from sglang.srt.runtime_context import (
get_device,
get_disagg,
get_exec,
get_mm,
@@ -1921,7 +1922,7 @@ class MMReceiverBase(ABC):
self.scheduler_embedding_port,
)
self.scheduler = scheduler
self.gpu_id = scheduler.ps.gpu_id if scheduler is not None else 0
self.gpu_id = get_device().gpu_id if scheduler is not None else 0
self.wait_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get()
self.embedding_pool = None
+4 -3
View File
@@ -87,6 +87,7 @@ from sglang.srt.observability.scheduler_stage_metrics import (
scheduler_stage_method,
)
from sglang.srt.runtime_context import (
get_device,
get_disagg,
get_parallel,
get_schedule,
@@ -212,7 +213,7 @@ class PrefillBootstrapQueue:
if get_disagg().disaggregation_enable_kv_checksum:
kv_args = self.kv_manager.kv_args
self.scheduler.kv_checksum_computer = KvChecksumComputer(
device=torch.device(f"cuda:{self.scheduler.ps.gpu_id}"),
device=torch.device(f"cuda:{get_device().gpu_id}"),
kv_data_ptrs=kv_args.kv_data_ptrs,
kv_item_lens=kv_args.kv_item_lens,
state_data_ptrs=kv_args.state_data_ptrs,
@@ -226,7 +227,7 @@ class PrefillBootstrapQueue:
kv_args = kv_args_class()
kv_args.engine_rank = self.tp_rank
kv_args.pp_rank = self.pp_rank
kv_args.system_dp_rank = self.scheduler.ps.dp_rank
kv_args.system_dp_rank = get_parallel().dp_rank
kv_args.rust_http_port = (
self.scheduler.rust_server.http_port
if self.scheduler.rust_server is not None
@@ -303,7 +304,7 @@ class PrefillBootstrapQueue:
self.metadata_buffers.get_buf_infos()
)
kv_args.ib_device = get_disagg().disaggregation_ib_device
kv_args.gpu_id = self.scheduler.ps.gpu_id
kv_args.gpu_id = get_device().gpu_id
req_to_token_pool = getattr(self.scheduler, "req_to_token_pool", None)
setup_state_kv_args(
@@ -13,7 +13,7 @@ from typing import TYPE_CHECKING, Callable, Optional, Tuple
from sglang.srt.disaggregation.common.conn import CommonKVReceiver
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.managers.io_struct import PdRoleSwitchReqInput, PdRoleSwitchReqOutput
from sglang.srt.runtime_context import get_context, get_disagg
from sglang.srt.runtime_context import get_context, get_device, get_disagg
from sglang.srt.utils import get_available_gpu_memory
if TYPE_CHECKING:
@@ -95,7 +95,7 @@ def handle_pd_role_switch(
)
try:
available_graph_gb = get_available_gpu_memory(
scheduler.device, scheduler.ps.gpu_id
scheduler.device, get_device().gpu_id
)
except Exception as e:
return _fail(
+14 -12
View File
@@ -22,12 +22,12 @@ from sglang.srt.distributed.gated_launch import maybe_wait_for_gated_launch
from sglang.srt.distributed.parallel_state import (
_tag_groups_for_flashinfer_allreduce_only,
)
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import initialize_dp_attention
from sglang.srt.layers.layernorm_sp import initialize_layernorm_sp
from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import (
get_device,
get_disagg,
get_exec,
get_parallel,
@@ -63,17 +63,17 @@ def init_torch_distributed(
server_args: ServerArgs,
model_config: ModelConfig,
device: str,
ps: ParallelState,
dist_port: int,
is_draft_worker: bool,
local_omp_cpuid: Optional[List[int]],
):
tic = time.perf_counter()
logger.info("Init torch distributed begin.")
parallel = get_parallel()
backend = _resolve_backend(device=device)
before_avail_memory = get_available_gpu_memory(device, ps.gpu_id)
before_avail_memory = get_available_gpu_memory(device, get_device().gpu_id)
if not get_parallel().enable_p2p_check:
monkey_patch_p2p_access_check()
@@ -83,8 +83,8 @@ def init_torch_distributed(
if not is_draft_worker:
if device == "cpu":
_init_cpu_threads_env(
tp_size=ps.tp_size,
tp_rank=ps.tp_rank,
tp_size=parallel.tp_size,
tp_rank=parallel.tp_rank,
local_omp_cpuid=local_omp_cpuid,
dist_init_method=dist_init_method,
)
@@ -96,16 +96,18 @@ def init_torch_distributed(
dist_init_method=dist_init_method,
server_args=server_args,
model_config=model_config,
gpu_id=ps.gpu_id,
gpu_id=get_device().gpu_id,
)
# 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 get_exec().comm.pre_warm_nccl and (
ps.tp_size > 1 or ps.pp_size > 1 or ps.moe_ep_size > 1
parallel.tp_size > 1 or parallel.pp_size > 1 or parallel.moe_ep_size > 1
):
_prewarm_nccl(
tp_size=ps.tp_size, pp_size=ps.pp_size, moe_ep_size=ps.moe_ep_size
tp_size=parallel.tp_size,
pp_size=parallel.pp_size,
moe_ep_size=parallel.moe_ep_size,
)
# CUDA graph capture enables the PyNCCL communicator for TP LM-head
@@ -115,7 +117,7 @@ def init_torch_distributed(
if (
device == "cuda"
and get_parallel().enable_tp_lm_head_all_to_all
and ps.tp_size > 1
and parallel.tp_size > 1
):
_prewarm_tp_lm_head_all_to_all()
@@ -127,13 +129,13 @@ def init_torch_distributed(
# including them in this WORLD reduction would deadlock on absent peers.
pre_model_load_memory = get_available_gpu_memory(
device,
ps.gpu_id,
get_device().gpu_id,
distributed=get_world_group().world_size > 1 and not is_draft_worker,
cpu_group=get_world_group().cpu_group,
)
# Check memory for tensor parallelism
local_gpu_memory = get_available_gpu_memory(device, ps.gpu_id)
if ps.tp_size > 1 and not is_draft_worker:
local_gpu_memory = get_available_gpu_memory(device, get_device().gpu_id)
if parallel.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,
@@ -1,51 +0,0 @@
from dataclasses import dataclass
from typing import Optional
@dataclass(frozen=True, slots=True, kw_only=True)
class ParallelState:
tp_rank: int
tp_size: int
pp_rank: int
pp_size: int
dp_rank: Optional[int]
dp_size: int
attn_tp_rank: int
attn_tp_size: int
attn_cp_rank: int
attn_cp_size: int
attn_dcp_rank: int
attn_dcp_size: int
attn_dp_rank: int
attn_dp_size: int
moe_ep_rank: int
moe_ep_size: int
moe_dp_rank: Optional[int]
moe_dp_size: int
gpu_id: int
@staticmethod
def trivial(**overrides: Optional[int]) -> "ParallelState":
kwargs: dict[str, Optional[int]] = dict(
tp_rank=0,
tp_size=1,
pp_rank=0,
pp_size=1,
dp_rank=0,
dp_size=1,
attn_tp_rank=0,
attn_tp_size=1,
attn_cp_rank=0,
attn_cp_size=1,
attn_dcp_rank=0,
attn_dcp_size=1,
attn_dp_rank=0,
attn_dp_size=1,
moe_ep_rank=0,
moe_ep_size=1,
moe_dp_rank=0,
moe_dp_size=1,
gpu_id=0,
)
kwargs.update(overrides)
return ParallelState(**kwargs)
+6 -5
View File
@@ -31,7 +31,6 @@ class EPLBManager:
self,
*,
model_config: ModelConfig,
ps: Any,
get_model: Callable[[], nn.Module],
get_expert_location_updater: Callable[[], ExpertLocationUpdater],
get_expert_backup_client: Callable[[], Any],
@@ -42,7 +41,6 @@ class EPLBManager:
# constructed (model load, expert_backup_client, weight_updater), so
# they are read through getters at rebalance time, not captured here.
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
@@ -163,7 +161,7 @@ class EPLBManager:
tp_rank=(
self._elastic_global_rank()
if is_post_scale_rebalance
else self._ps.tp_rank
else get_parallel().tp_rank
),
use_flat_topology=is_post_scale_rebalance,
expert_backup_client=self._get_expert_backup_client(),
@@ -223,7 +221,7 @@ class EPLBManager:
)
def _elastic_global_rank(self) -> int:
return self._ps.tp_rank + get_parallel().ep_join_rank_offset
return get_parallel().tp_rank + get_parallel().ep_join_rank_offset
def _check_rebalance_needed(self, average_utilization_rate_over_window):
if average_utilization_rate_over_window is None:
@@ -248,7 +246,10 @@ class EPLBManager:
return list(_chunk_list(all_layer_ids, chunk_size=chunk_size))
def _should_log_expert_location_metadata(self) -> bool:
return self._ps.tp_rank == 0 and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get()
return (
get_parallel().tp_rank == 0
and envs.SGLANG_LOG_EXPERT_LOCATION_METADATA.get()
)
def _log_rebalance_layout_before_update(
self,
@@ -172,7 +172,7 @@ class MlxModelRunnerStub(ModelRunner):
aux_state_size = get_schedule().max_mamba_cache_size
if aux_state_size is None:
return None
return aux_state_size // self.ps.attn_dp_size
return aux_state_size // self.attn_dp_size
def _resolve_max_running_requests(self) -> int:
"""Concurrency cap handed to the scheduler.
@@ -197,7 +197,7 @@ class MlxModelRunnerStub(ModelRunner):
requested_per_worker = None
resolved = min(capacity_cap, 4096)
else:
requested_per_worker = requested // self.ps.attn_dp_size
requested_per_worker = requested // self.attn_dp_size
resolved = min(requested_per_worker, capacity_cap)
aux_state_size = self._explicit_aux_state_size_per_worker()
@@ -209,7 +209,7 @@ class MlxModelRunnerStub(ModelRunner):
resolved = min(resolved, aux_state_size // ratio)
if resolved <= 0:
global_aux_state_size = get_schedule().max_mamba_cache_size
min_global_aux_state_size = ratio * self.ps.attn_dp_size
min_global_aux_state_size = ratio * self.attn_dp_size
raise RuntimeError(
f"MLX auxiliary-state cache is too small to serve any "
f"requests: max_mamba_cache_size={global_aux_state_size} "
@@ -108,7 +108,6 @@ class MlxTpModelWorker(TpModelWorker):
model_config=self.model_config,
mem_fraction_static=get_schedule().mem_fraction_static,
gpu_id=self.gpu_id,
ps=self.ps,
nccl_port=self.nccl_port,
server_args=self.server_args,
is_draft_worker=self.is_draft_worker,
@@ -428,7 +428,7 @@ class AscendAttnBackend(AttentionBackend):
self.is_dllm_model = True
self.dllm_block_size = self.dllm_config.block_size
self.attn_cp_size = model_runner.ps.attn_cp_size
self.attn_cp_size = model_runner.attn_cp_size
def _is_swa_layer(self, layer: RadixAttention) -> bool:
return (
@@ -203,7 +203,7 @@ class FlashAttentionBackend(AttentionBackend):
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
self.kv_index_translator = model_runner.kv_index_translator
self.skip_prefill = skip_prefill
self.attn_cp_size = model_runner.ps.attn_cp_size
self.attn_cp_size = model_runner.attn_cp_size
self._verify_mask = None
# The worker fetches the tree-mask scratch from the target backend
# only; draft-side instances must not allocate it.
@@ -333,10 +333,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.ps.tp_size
// model_runner.tp_size
)
self.num_kv_heads = model_runner.model_config.get_num_kv_heads(
model_runner.ps.tp_size
model_runner.tp_size
)
_softcapping = getattr(
model_runner.model_config.hf_text_config, "attn_logit_softcapping", None
@@ -162,9 +162,8 @@ class HPCOpsAttnBackend(AttentionBackend):
self.use_fp8 = model_runner.kv_cache_dtype == torch.float8_e4m3fn
if self.use_fp8:
heads = (
model_runner.model_config.num_attention_heads
// model_runner.ps.tp_size,
model_runner.model_config.get_num_kv_heads(model_runner.ps.tp_size),
model_runner.model_config.num_attention_heads // model_runner.tp_size,
model_runner.model_config.get_num_kv_heads(model_runner.tp_size),
)
if heads not in FP8_ROPE_SUPPORTED_HEAD_CONFIGS:
raise ValueError(
@@ -177,8 +176,8 @@ class HPCOpsAttnBackend(AttentionBackend):
config = model_runner.model_config
head_dim = config.head_dim
num_q_heads = config.num_attention_heads // model_runner.ps.tp_size
num_kv_heads = config.get_num_kv_heads(model_runner.ps.tp_size)
num_q_heads = config.num_attention_heads // model_runner.tp_size
num_kv_heads = config.get_num_kv_heads(model_runner.tp_size)
gqa_group_size = num_q_heads // num_kv_heads
if head_dim != _SUPPORTED_HEAD_DIM or gqa_group_size not in (
_SUPPORTED_GQA_GROUP_SIZES
@@ -66,7 +66,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.ps.tp_rank}"
new_dir = base_cache_dir / f"worker_{model_runner.tp_rank}"
logger.info(f"Setting Wave cache dir: {new_dir}")
cache.CACHE_BASE_DIR = new_dir
@@ -67,7 +67,7 @@ class XPUAttentionBackend(AttentionBackend):
self.num_attention_heads = (
model_runner.model_config.hf_text_config.num_attention_heads
)
self.tp_size = model_runner.ps.tp_size
self.tp_size = model_runner.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
@@ -510,7 +510,7 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
"PP-parallel DeepGEMM warmup start "
"(pp_rank=%d, tp_rank=%d, batch_sizes=%s, disagg=%s).",
get_parallel().pp_rank,
model_runner.ps.tp_rank,
model_runner.tp_rank,
batch_sizes,
disagg_mode,
)
+18
View File
@@ -47,6 +47,24 @@ if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
def deployment_attn_dp_size() -> int:
"""Attention-DP replicas in the deployment, which no draft scope narrows.
A draft runs on one attention-DP replica and its scope says so, but the
metadata a draft gathers is shaped by the replicas it gathers *with* --
the target's. Those come from the configuration, which the scope leaves
alone, so this answers the same number inside the scope and outside it.
"""
parallel = get_parallel()
attn_dp_size, _ = derive_attention_widths(
tp_size=parallel.tp_size,
attn_cp_size=parallel.attn_cp_size,
dp_size=parallel.dp_size,
enable_dp_attention=parallel.enable_dp_attention,
)
return attn_dp_size
def dp_gather_width() -> int:
"""How many replicas the DP sync gathers over.
+8 -55
View File
@@ -104,12 +104,10 @@ from sglang.srt.disaggregation.utils import (
from sglang.srt.distributed.parallel_state import (
abort_distributed_environment,
)
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.dllm.mixin.scheduler import SchedulerDllmMixin
from sglang.srt.environ import envs, exportable_env_vars
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
from sglang.srt.hardware_backend.mlx.runtime import use_mlx
from sglang.srt.layers.dp_attention import compute_dp_attention_world_info
from sglang.srt.layers.moe import initialize_moe_config
from sglang.srt.layers.quantization.fp4_utils import initialize_fp4_gemm_config
from sglang.srt.layers.quantization.fp8_utils import initialize_fp8_gemm_config
@@ -449,12 +447,8 @@ class Scheduler(
self,
server_args: ServerArgs,
port_args: PortArgs,
gpu_id: int,
tp_rank: int,
moe_ep_rank: int,
pp_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
dp_rank: Optional[int],
):
# NOTE: KEEP THE FOLLOWING CODE STYLE for this function:
@@ -516,38 +510,6 @@ class Scheduler(
self.enable_dp_attention = get_parallel().enable_dp_attention
self.enable_unified_memory = get_memory().enable_unified_memory
# Distributed rank info
attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = (
compute_dp_attention_world_info(
get_parallel().enable_dp_attention,
tp_rank,
get_parallel().tp_size,
get_parallel().dp_size,
get_parallel().attn_cp_size,
)
)
self.ps = ParallelState(
tp_rank=tp_rank,
tp_size=get_parallel().tp_size,
pp_rank=pp_rank,
pp_size=get_parallel().pp_size,
dp_rank=dp_rank,
dp_size=get_parallel().dp_size,
attn_tp_rank=attn_tp_rank,
attn_tp_size=attn_tp_size,
attn_cp_rank=attn_cp_rank,
attn_cp_size=get_parallel().attn_cp_size,
attn_dcp_rank=tp_rank % get_parallel().dcp_size,
attn_dcp_size=get_parallel().dcp_size,
attn_dp_rank=attn_dp_rank,
attn_dp_size=attn_dp_size,
moe_ep_rank=moe_ep_rank,
moe_ep_size=get_parallel().ep_size,
moe_dp_rank=moe_dp_rank,
moe_dp_size=get_parallel().moe_dp_size,
gpu_id=gpu_id,
)
# Init model configs
self.init_model_config()
@@ -779,7 +741,7 @@ class Scheduler(
if get_parallel().pp_size > 1:
logger.error("only zbal mix mode support pp_size > 1!")
init_zbal(
get_parallel().tp_size, self.ps.gpu_id, get_parallel().tp_rank
get_parallel().tp_size, get_device().gpu_id, get_parallel().tp_rank
) # only switch allocator if is mix mode
def init_model_config(self):
@@ -1009,8 +971,7 @@ class Scheduler(
def init_tp_model_worker(self):
worker_kwargs = dict(
server_args=self.server_args,
gpu_id=self.ps.gpu_id,
ps=self.ps,
gpu_id=get_device().gpu_id,
nccl_port=self.nccl_port,
)
@@ -1047,8 +1008,7 @@ class Scheduler(
# — is resolved per runner, not on a config copy.
draft_worker_kwargs = dict(
server_args=self.server_args,
gpu_id=self.ps.gpu_id,
ps=self.ps,
gpu_id=get_device().gpu_id,
nccl_port=self.nccl_port,
target_worker=self.tp_worker,
)
@@ -1246,7 +1206,7 @@ class Scheduler(
# Print debug info
self.startup_available_gpu_memory_gb = get_available_gpu_memory(
self.device, self.ps.gpu_id, empty_cache=False
self.device, get_device().gpu_id, empty_cache=False
)
if get_parallel().tp_rank == 0:
logger.info(
@@ -1582,7 +1542,7 @@ class Scheduler(
tp_rank=get_parallel().tp_rank,
tp_size=get_parallel().tp_size,
dp_size=get_parallel().dp_size,
gpu_id=self.ps.gpu_id,
gpu_id=get_device().gpu_id,
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
max_total_num_tokens=self.max_total_num_tokens,
pp_rank=get_parallel().pp_rank,
@@ -1613,7 +1573,7 @@ class Scheduler(
metadata_buffers=self.disagg_metadata_buffers,
tp_rank=get_parallel().tp_rank,
tp_size=get_parallel().tp_size,
gpu_id=self.ps.gpu_id,
gpu_id=get_device().gpu_id,
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
gloo_group=self.attn_tp_cpu_group,
max_total_num_tokens=self.max_total_num_tokens,
@@ -2248,7 +2208,6 @@ class Scheduler(
def init_profiler(self) -> None:
self.profiler_manager = SchedulerProfilerManager(
ps=self.ps,
dp_tp_cpu_group=self.dp_tp_cpu_group,
get_forward_ct=lambda: self.forward_ct,
)
@@ -2369,7 +2328,6 @@ class Scheduler(
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tree_cache=self.tree_cache,
offload_tags=self.weight_updater.offload_tags,
ps=self.ps,
model_config=self.model_config,
enable_overlap=self.enable_overlap,
spec_algorithm=self.spec_algorithm,
@@ -2457,7 +2415,6 @@ class Scheduler(
self._sched_idled = False
self.load_inquirer = SchedulerLoadInquirer(
disaggregation_mode=self.disaggregation_mode,
ps=self.ps,
server_args=self.server_args,
max_total_num_tokens=self.max_total_num_tokens,
max_running_requests=self.max_running_requests,
@@ -2500,7 +2457,6 @@ class Scheduler(
self.output_streamer = self.get_output_streamer_class()(
send_to_detokenizer=self.ipc_channels.send_to_detokenizer,
tree_cache=self.tree_cache,
ps=self.ps,
server_args=self.server_args,
is_generation=self.is_generation,
spec_algorithm=self.spec_algorithm,
@@ -3120,7 +3076,7 @@ class Scheduler(
self._add_request_to_queue(req)
return
if self.ps.pp_rank == 0 and getattr(
if get_parallel().pp_rank == 0 and getattr(
self.tree_cache.cache_controller, "pp_prefetch_command_group", None
):
recv_req.pp_prefetch_ticketed = bool(self._prefetch_kvcache(req))
@@ -6078,6 +6034,7 @@ def run_scheduler_process(
ranks=SpawnRanks(
world_rank=spawn_world_rank(server_args, tp_rank=tp_rank, pp_rank=pp_rank),
dp_rank=dp_rank,
gpu_id=gpu_id,
),
)
configure_scheduler_process(
@@ -6115,12 +6072,8 @@ def run_scheduler_process(
scheduler = Scheduler(
server_args,
port_args,
gpu_id,
tp_rank,
moe_ep_rank,
pp_rank,
attn_cp_rank,
moe_dp_rank,
dp_rank,
)
@@ -7,7 +7,6 @@ import torch
from sglang.srt.batch_overlap.two_batch_overlap import TboDPAttentionPreparer
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs
from sglang.srt.layers.cp.utils import get_cp_strategy
from sglang.srt.layers.dp_attention import dp_gather_width, world_dp_gather_enabled
@@ -537,7 +536,6 @@ class SchedulerDPAttnAdapter:
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator
tree_cache: BasePrefixCache
offload_tags: set[str]
ps: ParallelState
model_config: ModelConfig
enable_overlap: bool
spec_algorithm: SpeculativeAlgorithm
@@ -17,7 +17,6 @@ from sglang.srt.managers.load_snapshot import (
from sglang.srt.runtime_context import get_lora, get_parallel
if TYPE_CHECKING:
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.scheduler_components.pool_stats_observer import (
SchedulerPoolStatsObserver,
)
@@ -33,7 +32,6 @@ logger = logging.getLogger(__name__)
@dataclass(kw_only=True, slots=True, frozen=True)
class SchedulerLoadInquirer:
disaggregation_mode: DisaggregationMode
ps: ParallelState
server_args: ServerArgs
max_total_num_tokens: int
max_running_requests: int
@@ -308,7 +308,7 @@ class SchedulerMetricsReporter:
self.scheduler.enable_fpm = False
if (
get_observability().enable_forward_pass_metrics
and self.scheduler.ps.attn_tp_rank == 0
and get_parallel().attn_tp_rank == 0
and get_parallel().pp_rank == get_parallel().pp_size - 1
):
from sglang.srt.observability.forward_pass_metrics import (
@@ -316,9 +316,7 @@ class SchedulerMetricsReporter:
)
self.scheduler._fpm_dp_rank = (
self.scheduler.ps.dp_rank
if self.scheduler.ps.dp_rank is not None
else 0
get_parallel().dp_rank if get_parallel().dp_rank is not None else 0
)
self.scheduler._fpm_worker_id = (
get_observability().forward_pass_metrics_worker_id
@@ -483,9 +481,9 @@ class SchedulerMetricsReporter:
num_layers = float(getattr(model_config, "num_attention_layers", 0))
head_dim = float(getattr(model_config, "head_dim", 0))
num_attn_heads = float(
model_config.get_num_attention_heads(self.scheduler.ps.tp_size)
model_config.get_num_attention_heads(get_parallel().tp_size)
)
num_kv_heads = float(model_config.get_num_kv_heads(self.scheduler.ps.tp_size))
num_kv_heads = float(model_config.get_num_kv_heads(get_parallel().tp_size))
intermediate_size = getattr(hf_text_config, "intermediate_size", None)
if intermediate_size is None:
intermediate_size = getattr(hf_text_config, "ffn_hidden_size", 0)
@@ -19,7 +19,6 @@ from sglang.srt.beam_search.output import (
pack_beam_search_output,
)
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import (
BatchEmbeddingOutput,
@@ -53,7 +52,6 @@ class SchedulerOutputStreamer:
send_to_detokenizer: zmq.Socket
tree_cache: BasePrefixCache
ps: ParallelState
server_args: ServerArgs
is_generation: bool
spec_algorithm: SpeculativeAlgorithm
@@ -51,14 +51,12 @@ logger = logging.getLogger(__name__)
@dataclass(kw_only=True)
class SchedulerProfilerManager:
ps: Any
dp_tp_cpu_group: Any
get_forward_ct: Callable[[], int]
def __post_init__(self) -> None:
if envs.SGLANG_PROFILE_V2.get():
self._profile_manager = ProfileManager(
ps=self.ps,
cpu_group=self.dp_tp_cpu_group,
)
return
@@ -274,7 +272,7 @@ class SchedulerProfilerManager:
self.profile_in_progress = True
if "CUDA_PROFILER" in activities:
if self.ps.gpu_id == get_device().base_gpu_id:
if get_device().gpu_id == get_device().base_gpu_id:
torch.cuda.cudart().cudaProfilerStart()
self.profile_in_progress = True
@@ -387,7 +385,7 @@ class SchedulerProfilerManager:
torch.cuda.memory._record_memory_history(enabled=None)
if "CUDA_PROFILER" in self.profiler_activities:
if self.ps.gpu_id == get_device().base_gpu_id:
if get_device().gpu_id == get_device().base_gpu_id:
torch.cuda.cudart().cudaProfilerStop()
merge_message = self._merge_profile_traces()
+7 -11
View File
@@ -22,7 +22,6 @@ from typing import TYPE_CHECKING, List, Optional, Tuple
import torch
from sglang.srt.beam_search.logits_capture import capture_pre_sample_logits
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import (
DestroyWeightsUpdateGroupReqInput,
@@ -216,12 +215,12 @@ class BaseTpWorker(ABC):
return success, message
def _deserialize_own_rank(self, serialized_named_tensors):
"""Each rank deserializes only its own payload (index ps.tp_rank);
"""Each rank deserializes only its own payload (index tp_rank);
deserializing another rank's copy would break producer-side CUDA-IPC
refcounting."""
monkey_patch_torch_reductions()
return MultiprocessingSerializer.deserialize(
serialized_named_tensors[self.ps.tp_rank]
serialized_named_tensors[self.model_runner.tp_rank]
)
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
@@ -290,12 +289,12 @@ class BaseTpWorker(ABC):
extra = [n for n in tensors if n not in exp]
if mismatch or missing or extra:
raise RuntimeError(
f"[LORA-CHECK] rank{self.ps.tp_rank} adapter sync MISMATCH of {len(exp)} expected: "
f"[LORA-CHECK] rank{self.model_runner.tp_rank} adapter sync MISMATCH of {len(exp)} expected: "
f"{len(mismatch)} value-diff {mismatch[:5]}, {len(missing)} missing {missing[:5]}, "
f"{len(extra)} extra {extra[:5]}"
)
logger.info(
f"[LORA-CHECK] rank{self.ps.tp_rank} adapter sync OK: {len(exp)}/{len(exp)} tensors match (sha256)"
f"[LORA-CHECK] rank{self.model_runner.tp_rank} adapter sync OK: {len(exp)}/{len(exp)} tensors match (sha256)"
)
result = self.model_runner.load_lora_adapter_from_tensors(
recv_req.to_ref(),
@@ -322,7 +321,6 @@ class TpModelWorker(BaseTpWorker):
self,
server_args: ServerArgs,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
is_draft_worker: bool = False,
req_to_token_pool: Optional[ReqToTokenPool] = None,
@@ -335,7 +333,6 @@ class TpModelWorker(BaseTpWorker):
):
# Parse args
self.server_args = server_args
self.ps = ps
self.gpu_id = gpu_id
self.nccl_port = nccl_port
self.is_draft_worker = is_draft_worker
@@ -411,14 +408,15 @@ class TpModelWorker(BaseTpWorker):
tp_group = self.model_runner.tp_group
self.random_seed = broadcast_pyobj(
[get_device().random_seed],
tp_group.ranks[self.ps.tp_rank],
tp_group.ranks[self.model_runner.tp_rank],
tp_group.cpu_group,
src=tp_group.ranks[0],
)[0]
else:
self.random_seed = broadcast_pyobj(
[get_device().random_seed],
self.ps.tp_size * get_parallel().pp_rank + self.ps.tp_rank,
self.model_runner.tp_size * get_parallel().pp_rank
+ self.model_runner.tp_rank,
self.world_group.cpu_group,
src=self.world_group.ranks[0],
)[0]
@@ -521,7 +519,6 @@ class TpModelWorker(BaseTpWorker):
model_config=self.model_config,
mem_fraction_static=get_schedule().mem_fraction_static,
gpu_id=self.gpu_id,
ps=self.ps,
nccl_port=self.nccl_port,
server_args=self.server_args,
is_draft_worker=self.is_draft_worker,
@@ -542,7 +539,6 @@ class TpModelWorker(BaseTpWorker):
model_config=self.model_config,
mem_fraction_static=get_schedule().mem_fraction_static,
gpu_id=self.gpu_id,
ps=self.ps,
nccl_port=self.nccl_port,
server_args=self.server_args,
is_draft_worker=self.is_draft_worker,
@@ -1102,13 +1102,12 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
model_runner.lora_manager.prepare_lora_batch(ret)
if (
model_runner.ps.attn_dcp_size > 1
model_runner.attn_dcp_size > 1
and ret.out_cache_loc is not None
and is_hip()
):
ret.dcp_kv_mask = (
ret.positions % model_runner.ps.attn_dcp_size
== model_runner.ps.attn_dcp_rank
ret.positions % model_runner.attn_dcp_size == model_runner.attn_dcp_rank
)
return ret
@@ -37,7 +37,6 @@ from sglang.srt.distributed import bootstrap
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
maybe_init_shared_mooncake_transfer_engine,
)
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.dllm.config import DllmConfig
from sglang.srt.elastic_ep.elastic_ep import (
ElasticEPStateManager,
@@ -322,7 +321,6 @@ class ModelRunner:
model_config: ModelConfig,
mem_fraction_static: float,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
server_args: ServerArgs,
is_draft_worker: bool = False,
@@ -339,7 +337,6 @@ class ModelRunner:
# `server_args._draft_pool_config` mutation hack).
self.memory_pool_config = memory_pool_config
self.gpu_id = gpu_id
self.ps = ps
self.model_config = model_config
self.dist_port = nccl_port
self.server_args = server_args
@@ -416,12 +413,12 @@ class ModelRunner:
# Set device early so that TransferEngine init (e.g. Ascend NPU)
# can access the device context.
try:
torch.get_device_module(self.device).set_device(ps.gpu_id)
torch.get_device_module(self.device).set_device(get_device().gpu_id)
except Exception:
import os
logger.warning(
f"Context: {self.device=} {ps.gpu_id=} {os.environ.get('CUDA_VISIBLE_DEVICES')=} {get_parallel().tp_rank=} {get_parallel().tp_size=}"
f"Context: {self.device=} {get_device().gpu_id=} {os.environ.get('CUDA_VISIBLE_DEVICES')=} {get_parallel().tp_rank=} {get_parallel().tp_size=}"
)
raise
@@ -741,7 +738,6 @@ class ModelRunner:
self.eplb_manager = (
EPLBManager(
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,
@@ -805,8 +801,8 @@ class ModelRunner:
def get_pp_proxy_dspark_hidden_size(self) -> int:
return misc_utils.resolve_pp_proxy_dspark_hidden_size(
model=self.model,
pp_size=self.ps.pp_size,
pp_rank=self.ps.pp_rank,
pp_size=self.pp_size,
pp_rank=self.pp_rank,
)
def get_pp_proxy_topk_size(self) -> Optional[int]:
@@ -1175,7 +1171,6 @@ class ModelRunner:
server_args=self.server_args,
model_config=self.model_config,
device=self.device,
ps=self.ps,
dist_port=self.dist_port,
is_draft_worker=self.is_draft_worker,
local_omp_cpuid=self.local_omp_cpuid if self.device == "cpu" else None,
@@ -1197,6 +1192,10 @@ class ModelRunner:
self.pp_size = parallel.pp_size
self.attn_cp_rank = parallel.attn_cp_rank
self.attn_cp_size = parallel.attn_cp_size
self.attn_dcp_rank = parallel.attn_dcp_rank
self.attn_dcp_size = parallel.attn_dcp_size
self.moe_ep_size = parallel.moe_ep_size
self.dp_rank = parallel.dp_rank
def init_shared_mooncake_transfer_engine(self):
maybe_init_shared_mooncake_transfer_engine(gpu_id=self.gpu_id)
@@ -220,7 +220,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
self._cell_size == 0
and mambaish is not None
and bool(mambaish.full_attention_layer_ids)
and kvc.ps.pp_size > 1
and kvc.pp_size > 1
)
self._zero_kv_max_tokens = (
torch.iinfo(torch.int64).max
@@ -877,7 +877,7 @@ class SWAChunkCapPoolConfigurator(HybridSWAPoolConfigurator):
self._swa_cap = compute_swa_request_cap(
page_size=kvc.page_size,
window=kvc.sliding_window_size,
attn_dp_size=kvc.ps.attn_dp_size,
attn_dp_size=kvc.attn_dp_size,
)
@staticmethod
@@ -1015,7 +1015,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
self.compression_ratios = cfg.compress_ratios[
kvc.layer_info.start_layer : kvc.layer_info.end_layer
]
if kvc.ps.pp_size > 1:
if kvc.pp_size > 1:
logger.info(
f"DSV4 pool PP slice: rank={kvc.pp_group.rank_in_group} "
f"layers=[{kvc.layer_info.start_layer},{kvc.layer_info.end_layer}) "
@@ -1032,9 +1032,9 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
self.page_size = kvc.page_size
self.is_speculative = get_spec().speculative_algorithm is not None
self.online_c128_mtp_max_draft_tokens = max_speculative_num_draft_tokens() or 0
self.attn_dp_size = kvc.ps.attn_dp_size
self.attn_dp_size = kvc.attn_dp_size
self.requested_max_running_requests_per_worker = (
get_schedule().max_running_requests // kvc.ps.attn_dp_size
get_schedule().max_running_requests // kvc.attn_dp_size
if get_schedule().max_running_requests is not None
else None
)
@@ -547,9 +547,9 @@ class BaseRunner(ABC):
if (
capture_forward_mode == ForwardMode.EXTEND
and get_parallel().pp_rank != 0
and mr.ps.attn_cp_size > 1
and mr.attn_cp_size > 1
):
pp_hidden_tokens = num_tokens // mr.ps.attn_cp_size
pp_hidden_tokens = num_tokens // mr.attn_cp_size
pp_proxy_tensors = PPProxyTensors(
{k: v[:pp_hidden_tokens] for k, v in buffers.pp_proxy_tensors.items()}
)
@@ -296,7 +296,7 @@ class EagerRunner(BaseRunner):
or cp_active
or forward_batch.forward_mode.is_target_verify()
):
if model_runner.ps.attn_dcp_size > 1 and hasattr(
if model_runner.attn_dcp_size > 1 and hasattr(
model_runner.model, "prepare_context_parallel_metadata_for_dcp"
):
# prepare kv cache buffer for dcp to gather kv cache
@@ -147,10 +147,10 @@ def flashinfer_autotune_cache_path(model_runner: ModelRunner) -> Path:
str(mr.dtype),
str(get_model().quantization),
str(get_exec().moe.moe_runner_backend),
str(mr.ps.tp_size),
str(mr.tp_size),
str(get_parallel().pp_size),
str(mr.ps.attn_dp_size),
str(mr.ps.moe_ep_size),
str(mr.attn_dp_size),
str(mr.moe_ep_size),
str(mr.model_config.hf_config.__class__.__name__),
]
# A different skip policy must not reuse previously tuned tactics.
@@ -171,7 +171,7 @@ def flashinfer_autotune_cache_path(model_runner: ModelRunner) -> Path:
cache_dir.mkdir(parents=True, exist_ok=True)
return (
cache_dir
/ f"rank_tp{mr.ps.tp_rank}_pp{get_parallel().pp_rank}_dp{mr.ps.dp_rank or 0}.json"
/ f"rank_tp{mr.tp_rank}_pp{get_parallel().pp_rank}_dp{mr.dp_rank or 0}.json"
)
+3 -6
View File
@@ -21,7 +21,7 @@ from typing import Any, Dict, Optional
import ray
from sglang.srt.arg_groups.overrides import declare_resolution
from sglang.srt.runtime_context import SpawnRanks, publish, spawn_world_rank
from sglang.srt.runtime_context import SpawnRanks, get_device, publish, spawn_world_rank
from sglang.srt.server_args import PortArgs, ServerArgs
logger = logging.getLogger(__name__)
@@ -93,6 +93,7 @@ class SchedulerActor:
server_args, tp_rank=tp_rank, pp_rank=pp_rank
),
dp_rank=dp_rank,
gpu_id=actual_gpu_id,
),
)
@@ -124,12 +125,8 @@ class SchedulerActor:
self.scheduler = Scheduler(
server_args=server_args,
port_args=port_args,
gpu_id=actual_gpu_id,
tp_rank=tp_rank,
moe_ep_rank=moe_ep_rank,
pp_rank=pp_rank,
attn_cp_rank=attn_cp_rank,
moe_dp_rank=moe_dp_rank,
dp_rank=dp_rank,
)
@@ -146,7 +143,7 @@ class SchedulerActor:
import torch
# Need to set the GPU id for the event loop for nccl to work
torch.cuda.set_device(self.scheduler.ps.gpu_id)
torch.cuda.set_device(get_device().gpu_id)
self.scheduler.run_event_loop()
except Exception as e:
logger.error(f"Scheduler PP{self._pp_rank} TP{self._tp_rank} crashed: {e}")
+34 -5
View File
@@ -1119,7 +1119,7 @@ def _install_derived_leaves(tops: dict, server_args: Any) -> None:
"""
import importlib
from sglang.srt.arg_groups.arg_utils import Derived
from sglang.srt.arg_groups.arg_utils import _NO_DEFAULT, Derived
from sglang.srt.arg_groups.overrides import resolved_view
namespaces = getattr(type(server_args), "_NAMESPACES", None)
@@ -1131,7 +1131,18 @@ def _install_derived_leaves(tops: dict, server_args: Any) -> None:
if path is None:
continue
for name, decl in vars(source).items():
if not isinstance(decl, Derived) or not decl.fn:
if not isinstance(decl, Derived):
continue
if not decl.fn:
# Nothing computes it. Seed the ones whose absence is itself an
# answer, so a process that never states one still reads it;
# the rest stay unwritten and say so when read.
if decl.default is not _NO_DEFAULT:
bag = tops.get(path.split(".")[0])
for segment in path.split(".")[1:]:
bag = bag and getattr(bag, segment, None)
if bag is not None:
bag._set(name, decl.default)
continue
module, _, attr = decl.fn.rpartition(".")
bag = tops.get(path.split(".")[0])
@@ -1251,7 +1262,20 @@ class RuntimeContext:
# Snapshot resolved config into the namespace bags (the single source of
# truth for config reads). Placed by `namespace_of`; a mock/partial
# config that declares no namespace yields an empty tree (no bags).
# A name the configuration does not carry survives the re-projection.
# `gpu_id` is stated by the spawn, not derived, so rebuilding the bags
# from the record must not unwrite it -- the record has no field for it
# to be rebuilt from.
stated = {}
if self._config_bags is not None:
device = self._config_bags.get("device")
fields = object.__getattribute__(device, "_fields") if device else {}
if "gpu_id" in fields:
stated["gpu_id"] = fields["gpu_id"]
self._config_bags = _build_config_bags(server_args)
device = self._config_bags.get("device")
if stated and device is not None:
device._set("gpu_id", stated["gpu_id"])
spec = self._config_bags.get("spec")
if spec is not None:
from sglang.srt.arg_groups.overrides import (
@@ -1818,8 +1842,13 @@ def publish(
# never built. With DCP on it is a position, and the bundle below states it.
if not _CONTEXT.parallel.dcp_enabled:
_CONTEXT.parallel.override_permanently(attn_dcp_rank=0)
if ranks is not None and ranks.gpu_id is not None:
_CONTEXT.override("spawn", gpu_id=ranks.gpu_id)
# Stated on the bag directly: `gpu_id` is declared but not configured, so
# it is not a leaf `override` can route, and the spawn is the only thing
# that knows it. Written whatever it is, `None` included -- most roles run
# on no device, and that is an answer rather than a name nobody wrote.
_CONTEXT.config_bag("device")._set(
"gpu_id", ranks.gpu_id if ranks is not None else None
)
if ranks is not None:
# The placement, worked out here rather than carried: the widths are on
# the bag a moment ago, and `world_rank` fixes the rest. A read of any
@@ -1882,7 +1911,7 @@ def _attention_ranks(parallel, tp_rank: int) -> dict:
The widths are already on the bag -- `publish` computed them a moment ago --
and the rank comes from the spawn, so the position is known here, before any
process group exists. That is the point: a rank read then works in a process
that never initialises distributed, which is what `ParallelState` provided
that never initialises distributed, which is what the per-runner record used to provide
by being a plain frozen record.
These are stamped rather than written as bag leaves because they are
+6 -6
View File
@@ -119,17 +119,17 @@ class RustServer:
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
# Preserve the DP startup log; ports use node-local offsets.
dp_rank = scheduler.ps.attn_dp_rank if scheduler.ps.dp_size > 1 else None
dp_rank = get_parallel().attn_dp_rank if get_parallel().dp_size > 1 else None
if get_exec().moe.is_ep_scale_joiner:
# The joining TP group is entirely local to this node.
tp_size_per_node = scheduler.ps.tp_size
tp_size_per_node = get_parallel().tp_size
else:
nnodes_per_pp_rank = max(get_parallel().nnodes // get_parallel().pp_size, 1)
tp_size_per_node = scheduler.ps.tp_size // nnodes_per_pp_rank
dp_group_width = scheduler.ps.attn_tp_size * scheduler.ps.attn_cp_size
tp_size_per_node = get_parallel().tp_size // nnodes_per_pp_rank
dp_group_width = get_parallel().attn_tp_size * get_parallel().attn_cp_size
# Count DP leaders within this node's TP range. The first leader must
# use the base port even when a DP group spans multiple nodes.
local_dp_rank = (scheduler.ps.tp_rank % tp_size_per_node) // dp_group_width
local_dp_rank = (get_parallel().tp_rank % tp_size_per_node) // dp_group_width
listen_port = get_serving().port + local_dp_rank
listen_addr = NetworkAddress(get_serving().host, listen_port).to_host_port_str()
@@ -179,7 +179,7 @@ class RustServer:
# Under DP every rank runs its own server on its own port, so the rank is
# what tells two otherwise identical startup lines apart.
dp_note = (
"" if dp_rank is None else f" (DP rank {dp_rank}/{scheduler.ps.dp_size})"
"" if dp_rank is None else f" (DP rank {dp_rank}/{get_parallel().dp_size})"
)
logger.info(
"SGLANG_RUST_SERVER enabled, Rust server listen on %s%s",
@@ -1,7 +1,6 @@
import logging
import math
import os
from dataclasses import replace
from typing import List, Optional, Tuple
import torch
@@ -19,7 +18,6 @@ from sglang.kernels.ops.speculative.dspark.dspark_accept import (
accept_sampling,
)
from sglang.srt.configs.hybrid_arch import mambaish_config
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs
from sglang.srt.layers.logits_processor import should_apply_lm_head_quant_method
from sglang.srt.layers.logprob_processor import compute_spec_logprobs
@@ -364,7 +362,6 @@ class DFlashWorkerV2(BaseSpecWorker):
self,
server_args: ServerArgs,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
target_worker: TpModelWorker,
):
@@ -372,7 +369,6 @@ class DFlashWorkerV2(BaseSpecWorker):
self.server_args = server_args
self.gpu_id = gpu_id
self.ps = ps
self.nccl_port = nccl_port
self._target_worker = target_worker
self.model_runner = target_worker.model_runner
@@ -416,7 +412,6 @@ class DFlashWorkerV2(BaseSpecWorker):
bundle = build_draft_tp_worker(
server_args=server_args,
gpu_id=gpu_id,
ps=replace(ps, pp_rank=0, pp_size=1),
nccl_port=nccl_port,
target_model_config=target_worker.model_runner.model_config,
algo_label="DFLASH",
@@ -453,7 +448,7 @@ class DFlashWorkerV2(BaseSpecWorker):
validate_domino_runtime(
device=torch.device(self.device),
tp_size=int(get_parallel().tp_group.world_size),
tp_rank=int(self.ps.tp_rank),
tp_rank=int(self.model_runner.tp_rank),
target_vocab_size=int(self.model_runner.model_config.vocab_size),
draft_vocab_size=int(self.draft_model_runner.model_config.vocab_size),
hidden_size=int(self.draft_model.config.hidden_size),
@@ -500,7 +495,7 @@ class DFlashWorkerV2(BaseSpecWorker):
)
self._maybe_merge_trained_mask_embedding()
self._cache_full_embed_weight()
if self.ps.tp_rank == 0:
if self.model_runner.tp_rank == 0:
logger.info(
"Initialized DFLASH draft runner. attention_backend=%s, model=%s, block_size=%s, draft_window_size=%s, compact_cache=%s",
bundle.resolved_attention_backend,
@@ -650,7 +645,7 @@ class DFlashWorkerV2(BaseSpecWorker):
# shared graph capture/replay; keep the draft eager under dp
# attention.
capture_decode_cuda_graph = False
if self.ps.tp_rank == 0:
if self.model_runner.tp_rank == 0:
logger.warning(
"Disable DFLASH draft cuda graph because dp attention "
"is enabled (draft runs eager)."
@@ -795,7 +790,7 @@ class DFlashWorkerV2(BaseSpecWorker):
def _maybe_build_draft_sampler(self):
def _eager(reason):
if self.ps.tp_rank == 0:
if self.model_runner.tp_rank == 0:
logger.info("DFLASH draft greedy head kept eager (reason=%s).", reason)
return None
@@ -819,7 +814,7 @@ class DFlashWorkerV2(BaseSpecWorker):
):
return _eager("unsupported quantized lm_head")
self.draft_model.lm_head = lm_head
if self.ps.tp_rank == 0:
if self.model_runner.tp_rank == 0:
logger.info(
"DFLASH selector decode folded into the draft cuda graph "
"(sampling_enabled=%s).",
@@ -843,7 +838,7 @@ class DFlashWorkerV2(BaseSpecWorker):
embed_proj = self.draft_model.embed_proj
if prefix_gru is None or embed_proj is None:
return _eager("Domino projector modules are unavailable")
if self.ps.tp_rank == 0:
if self.model_runner.tp_rank == 0:
logger.info(
"DFLASH Domino rollout folded into the draft cuda graph (tp=%s).",
int(tp_group.world_size),
@@ -882,7 +877,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.ps.tp_rank == 0:
if self.model_runner.tp_rank == 0:
logger.info(
"DFLASH draft greedy head folded into the draft cuda graph (tp=%d).",
tp_group.world_size,
@@ -908,7 +903,7 @@ class DFlashWorkerV2(BaseSpecWorker):
fused_disable_reason = "draft model does not support fused context KV"
if fused_disable_reason is not None:
if self.ps.tp_rank == 0:
if self.model_runner.tp_rank == 0:
logger.info(
"DFLASH fused KV materialization disabled: %s",
fused_disable_reason,
@@ -951,7 +946,7 @@ class DFlashWorkerV2(BaseSpecWorker):
break
if fused_disable_reason is not None:
if self.ps.tp_rank == 0:
if self.model_runner.tp_rank == 0:
logger.info(
"DFLASH fused KV materialization disabled: %s",
fused_disable_reason,
@@ -973,7 +968,7 @@ class DFlashWorkerV2(BaseSpecWorker):
max_position_hint=self.target_worker.model_runner.model_config.context_len
+ int(self.block_size),
)
if self.ps.tp_rank == 0:
if self.model_runner.tp_rank == 0:
logger.info(
"DFLASH fused KV materialization enabled. "
"n_layers=%d, num_kv_heads=%d, head_dim=%d",
@@ -1266,7 +1261,7 @@ class DFlashWorkerV2(BaseSpecWorker):
embedding_tensor.to(embed_module.weight.dtype)
)
if self.ps.tp_rank == 0:
if self.model_runner.tp_rank == 0:
logger.info(
"Merged trained mask embedding into target model "
"(mask_token_id=%s, source=%s)",
@@ -1301,7 +1296,7 @@ class DFlashWorkerV2(BaseSpecWorker):
parts = [torch.empty_like(shard_t) for _ in range(tp_size)]
dist.all_gather(parts, shard_t, group=tp_group.device_group)
self._full_embed_gpu = torch.cat(parts, dim=0)[:vocab_size]
if self.ps.tp_rank == 0:
if self.model_runner.tp_rank == 0:
logger.info(
"DFLASH cached full embed on GPU for dp attention: shape=%s",
list(self._full_embed_gpu.shape),
@@ -1364,7 +1359,7 @@ class DFlashWorkerV2(BaseSpecWorker):
if resolved_id is None:
resolved_id = tokenizer.convert_tokens_to_ids(mask_token)
if added and self.ps.tp_rank == 0:
if added and self.model_runner.tp_rank == 0:
logger.info(
"Added DFLASH mask token to tokenizer. token=%s, mask_token_id=%s, tokenizer_len=%s, model_vocab_size=%s",
mask_token,
@@ -2179,7 +2174,7 @@ class DFlashWorkerV2(BaseSpecWorker):
if self.selector is not None:
if self._selector_sampling_enabled:
return
if not self._warned_sampling_fallback and self.ps.tp_rank == 0:
if not self._warned_sampling_fallback and self.model_runner.tp_rank == 0:
logger.warning(
"DFLASH non-greedy verification is unavailable on this "
"build/device; falling back to greedy argmax verification. "
@@ -2192,7 +2187,7 @@ class DFlashWorkerV2(BaseSpecWorker):
if (
not is_dflash_sampling_verify_available()
and not self._warned_sampling_fallback
and self.ps.tp_rank == 0
and self.model_runner.tp_rank == 0
):
logger.warning(
"DFLASH non-greedy verification is unavailable on this build/device; "
@@ -16,7 +16,6 @@ from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.model_executor.model_runner import ModelRunner
logger = logging.getLogger(__name__)
@@ -63,7 +62,6 @@ def build_draft_tp_worker(
*,
server_args: ServerArgs,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
target_model_config: ModelConfig,
algo_label: str,
@@ -88,7 +86,6 @@ def build_draft_tp_worker(
draft_worker = draft_worker_cls(
server_args=server_args,
gpu_id=gpu_id,
ps=ps,
nccl_port=nccl_port,
is_draft_worker=True,
random_seed=random_seed,
@@ -1,6 +1,5 @@
import logging
from contextlib import nullcontext
from dataclasses import replace
from typing import Callable, Optional, Protocol, runtime_checkable
import torch
@@ -9,8 +8,6 @@ from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import (
is_unified_kv_triton,
)
from sglang.srt.configs.hybrid_arch import mambaish_config
from sglang.srt.distributed import get_pp_group
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs
from sglang.srt.layers.logprob_processor import compute_spec_logprobs
from sglang.srt.lora.layers import unwrap_lora_layer
@@ -134,7 +131,6 @@ class DSparkWorkerV2(BaseSpecWorker):
self,
server_args: ServerArgs,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
target_worker: TpModelWorker,
draft_worker_cls: type[TpModelWorker] = TpModelWorker,
@@ -143,14 +139,13 @@ class DSparkWorkerV2(BaseSpecWorker):
self.server_args = server_args
self.gpu_id = gpu_id
self.ps = ps
self.nccl_port = nccl_port
self._target_worker = target_worker
self.model_runner = target_worker.model_runner
self.page_size = get_schedule().page_size
self.device = target_worker.device
self._draft_worker = None
self._hosts_draft = get_pp_group().is_last_rank
self._hosts_draft = get_parallel().pp_group.is_last_rank
if not self._hosts_draft:
return
@@ -166,7 +161,7 @@ class DSparkWorkerV2(BaseSpecWorker):
if (
get_parallel().enable_dp_attention
and self._draft_is_moe
and ps.attn_tp_size > 1
and get_parallel().attn_tp_size > 1
):
raise ValueError(
"DSpark + dp attention with a DeepSeek-V4 (MoE) draft requires "
@@ -178,7 +173,6 @@ class DSparkWorkerV2(BaseSpecWorker):
bundle = build_draft_tp_worker(
server_args=server_args,
gpu_id=gpu_id,
ps=replace(ps, pp_rank=0, pp_size=1),
nccl_port=nccl_port,
target_model_config=target_worker.model_runner.model_config,
algo_label="DSPARK",
@@ -232,7 +226,7 @@ class DSparkWorkerV2(BaseSpecWorker):
else parallel.tp_group
)
if self.ps.tp_rank == 0:
if self.model_runner.tp_rank == 0:
logger.info(
"Initialized DSpark draft runner. attention_backend=%s, model=%s, "
"gamma=%s, verify_num_draft_tokens=%s, query_token_num=%s, "
@@ -255,7 +249,7 @@ class DSparkWorkerV2(BaseSpecWorker):
)
if getattr(self.draft_model, "uses_own_vocab_modules", False):
if self.ps.tp_rank == 0:
if self.model_runner.tp_rank == 0:
logger.info(
"DSpark draft uses its checkpoint-local embedding and LM head."
)
@@ -279,7 +273,7 @@ class DSparkWorkerV2(BaseSpecWorker):
gamma=self.gamma,
model_runner=self.model_runner,
device=self.device,
tp_rank=self.ps.tp_rank,
tp_rank=self.model_runner.tp_rank,
verify_num_draft_tokens=self.verify_num_draft_tokens,
tp_sync=self._tp_sync,
)
@@ -327,7 +321,7 @@ class DSparkWorkerV2(BaseSpecWorker):
and self._verify_planner.mode_value == "static"
and self._draft_is_moe
and not get_parallel().enable_dp_attention
and self.ps.pp_size == 1
and self.model_runner.pp_size == 1
)
if (
(self._verify_planner.is_compact_mode or static_epilogue_supported)
@@ -391,7 +385,7 @@ class DSparkWorkerV2(BaseSpecWorker):
planner=self._verify_planner,
gamma=self.gamma,
verify_num_draft_tokens=self.verify_num_draft_tokens,
tp_rank=self.ps.tp_rank,
tp_rank=self.model_runner.tp_rank,
device=self.device,
simulate_acc_len=self._simulate_acc_len,
)
@@ -451,7 +445,7 @@ class DSparkWorkerV2(BaseSpecWorker):
draft_model=self.draft_model,
is_deepseek_v4_draft=self._draft_is_moe,
)
if self._target_hidden_projection_enabled and self.ps.tp_rank == 0:
if self._target_hidden_projection_enabled and self.model_runner.tp_rank == 0:
logger.info(
"DSpark prefill target-hidden projection runs before "
"sequence-parallel gather."
@@ -508,7 +502,7 @@ class DSparkWorkerV2(BaseSpecWorker):
gamma=self.gamma,
max_bs=max(get_exec().graph.cuda_graph_config.decode.bs),
device=self.device,
tp_rank=self.ps.tp_rank,
tp_rank=self.model_runner.tp_rank,
tp_sync=self._tp_sync,
available_memory_gb=available_memory_gb,
confidence_fn=(
@@ -10,6 +10,7 @@ from sglang.srt.compilation.torch_compile_decoration import set_torch_compile_co
from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import (
DpPaddingMode,
deployment_attn_dp_size,
set_dp_buffer_len,
set_is_extend_in_batch,
)
@@ -114,8 +115,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.ps.tp_size
self.attn_dp_size = model_runner.ps.attn_dp_size
self.tp_size = model_runner.tp_size
self.attn_dp_size = deployment_attn_dp_size()
self.pp_size = get_parallel().pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = get_exec().graph.disable_cuda_graph_padding
@@ -9,6 +9,7 @@ import torch
from sglang.srt.compilation.torch_compile_decoration import set_torch_compile_config
from sglang.srt.layers.dp_attention import (
DpPaddingMode,
deployment_attn_dp_size,
set_dp_buffer_len,
set_is_extend_in_batch,
)
@@ -111,8 +112,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.ps.tp_size
self.attn_dp_size = model_runner.ps.attn_dp_size
self.tp_size = model_runner.tp_size
self.attn_dp_size = deployment_attn_dp_size()
self.pp_size = get_parallel().pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = get_exec().graph.disable_cuda_graph_padding
@@ -1,14 +1,12 @@
import contextlib
import logging
import time
from dataclasses import replace
from typing import List, Optional
import torch
from sglang.kernels.ops.speculative.topk1 import draft_topk1_postprocess
from sglang.srt.configs.model_config import get_dsa_mtp_topk_width
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs
from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_extend_npu_graph_runner import (
EAGLEDraftExtendNpuGraphRunner,
@@ -238,7 +236,6 @@ class EagleDraftWorker(EagleDraftWorkerBase):
self,
server_args: ServerArgs,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
target_worker: TpModelWorker,
):
@@ -247,7 +244,6 @@ class EagleDraftWorker(EagleDraftWorkerBase):
# copy args
self.server_args = server_args
self.gpu_id = gpu_id
self.ps = ps
self.nccl_port = nccl_port
self.target_worker = target_worker
@@ -288,8 +284,6 @@ class EagleDraftWorker(EagleDraftWorkerBase):
self.draft_worker = TpModelWorker(
server_args=server_args,
gpu_id=gpu_id,
# spec workers don't support pipeline parallelism
ps=replace(ps, pp_rank=0, pp_size=1),
nccl_port=nccl_port,
is_draft_worker=True,
# The draft runs at absolute target positions.
@@ -1293,7 +1287,6 @@ class EAGLEWorkerV2(BaseSpecWorker):
self,
server_args: ServerArgs,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
target_worker: TpModelWorker,
):
@@ -1304,7 +1297,6 @@ class EAGLEWorkerV2(BaseSpecWorker):
self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.ps = ps
self.gpu_id = gpu_id
self.device = get_device().device
self._target_worker = target_worker
@@ -1320,7 +1312,6 @@ class EAGLEWorkerV2(BaseSpecWorker):
EagleDraftWorker(
server_args,
gpu_id,
ps,
nccl_port,
target_worker,
)
@@ -1873,7 +1864,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
monkey_patch_torch_reductions()
named_tensors = MultiprocessingSerializer.deserialize(
recv_req.serialized_named_tensors[self.ps.tp_rank]
recv_req.serialized_named_tensors[self.model_runner.tp_rank]
)
success, message = (
self.draft_worker.draft_runner.weight_updater.update_weights_from_tensor(
@@ -8,6 +8,7 @@ import torch
from sglang.srt.compilation.torch_compile_decoration import set_torch_compile_config
from sglang.srt.layers.dp_attention import (
DpPaddingMode,
deployment_attn_dp_size,
set_dp_buffer_len,
set_is_extend_in_batch,
)
@@ -98,8 +99,8 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather()
self.tp_size = self.model_runner.ps.tp_size
self.attn_dp_size = self.model_runner.ps.attn_dp_size
self.tp_size = self.model_runner.tp_size
self.attn_dp_size = deployment_attn_dp_size()
self.pp_size = get_parallel().pp_size
self.speculative_num_steps = get_spec().speculative_num_steps
self.topk = get_spec().speculative_eagle_topk
@@ -23,12 +23,10 @@ from __future__ import annotations
import logging
import time
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 (
draft_model_build_scope,
speculative_moe_a2a_backend_context,
@@ -102,7 +100,6 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
self,
server_args: ServerArgs,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
target_worker: TpModelWorker,
):
@@ -112,7 +109,6 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.ps = ps
self.gpu_id = gpu_id
self.device = get_device().device
self.target_worker = target_worker
@@ -146,8 +142,6 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
self,
server_args=server_args,
gpu_id=gpu_id,
# spec workers don't support pipeline parallelism
ps=replace(ps, pp_rank=0, pp_size=1),
nccl_port=nccl_port,
is_draft_worker=True,
# The draft runs at absolute target positions.
@@ -702,7 +696,6 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
self,
server_args: ServerArgs,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
target_worker: TpModelWorker,
):
@@ -715,7 +708,6 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.ps = ps
self.gpu_id = gpu_id
self.device = get_device().device
self._target_worker = target_worker
@@ -730,7 +722,6 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
self._draft_worker = FrozenKVMTPDraftWorker(
server_args,
gpu_id,
ps,
nccl_port,
target_worker,
)
@@ -155,7 +155,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.ps.tp_size
self.tp_size = model_runner.tp_size
self.dp_size = get_parallel().dp_size
self.pp_size = get_parallel().pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile
@@ -21,7 +21,6 @@ from typing import TYPE_CHECKING, List
import torch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs
from sglang.srt.hardware_backend.npu.graph_runner.multi_layer_eagle_draft_extend_npu_graph_runner import (
MultiLayerEagleMultiStepDraftExtendNpuGraphRunner,
@@ -121,7 +120,6 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
self,
server_args: ServerArgs,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
target_worker: TpModelWorker,
):
@@ -130,7 +128,6 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
# copy args
self.server_args = server_args
self.gpu_id = gpu_id
self.ps = ps
self.nccl_port = nccl_port
self.target_worker = target_worker
self.draft_extend_attn_backend_list = []
@@ -171,8 +168,6 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
self.draft_worker = TpModelWorker(
server_args=server_args,
gpu_id=gpu_id,
# spec workers don't support pipeline parallelism
ps=replace(ps, pp_rank=0, pp_size=1),
nccl_port=nccl_port,
is_draft_worker=True,
is_multi_layer_eagle=True,
@@ -1013,7 +1008,6 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
self,
server_args: ServerArgs,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
target_worker: TpModelWorker,
):
@@ -1035,7 +1029,6 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
self._draft_worker = MultiLayerEagleDraftWorker(
server_args,
gpu_id,
ps,
nccl_port,
target_worker,
)
@@ -7,7 +7,6 @@ import torch
from sglang.kernels.ops.speculative.cache_locs import (
assign_extend_cache_locs_func as assign_extend_cache_locs_func,
)
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.logprob_processor import compute_spec_logprobs
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.scheduler import GenerationBatchResult
@@ -89,7 +88,6 @@ class NGRAMWorker(BaseSpecWorker):
self,
server_args: ServerArgs,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
target_worker: TpModelWorker,
):
@@ -99,7 +97,7 @@ class NGRAMWorker(BaseSpecWorker):
self.enable_overlap = not get_schedule().disable_overlap_schedule
self._target_worker = target_worker
self.model_runner = target_worker.model_runner
self.tp_rank = ps.tp_rank
self.tp_rank = self.model_runner.tp_rank
self.page_size = get_schedule().page_size
self.draft_token_num: int = get_spec().speculative_num_draft_tokens
self.max_trie_depth: int = get_spec().speculative_ngram_max_trie_depth
@@ -1,10 +1,8 @@
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 (
draft_model_build_scope,
speculative_moe_backend_context,
@@ -45,7 +43,6 @@ class StandaloneDraftWorker(EagleDraftWorker):
self,
server_args: ServerArgs,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
target_worker: TpModelWorker,
):
@@ -54,7 +51,6 @@ class StandaloneDraftWorker(EagleDraftWorker):
# copy args
self.server_args = server_args
self.gpu_id = gpu_id
self.ps = ps
self.nccl_port = nccl_port
self.target_worker = target_worker
@@ -84,8 +80,6 @@ class StandaloneDraftWorker(EagleDraftWorker):
self.draft_worker = TpModelWorker(
server_args=server_args,
gpu_id=gpu_id,
# spec workers don't support pipeline parallelism
ps=replace(ps, pp_rank=0, pp_size=1),
nccl_port=nccl_port,
is_draft_worker=True,
# The draft runs at absolute target positions.
@@ -167,7 +161,6 @@ class StandaloneWorkerV2(EAGLEWorkerV2):
self,
server_args: ServerArgs,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
target_worker: TpModelWorker,
):
@@ -190,7 +183,6 @@ class StandaloneWorkerV2(EAGLEWorkerV2):
self._draft_worker = StandaloneDraftWorker(
server_args,
gpu_id,
ps,
nccl_port,
target_worker,
)
+2 -2
View File
@@ -45,8 +45,8 @@ def init_uno_lora_manager(
dtype=model_runner.dtype,
server_args=model_runner.server_args,
lora_backend="uno_cublas", # fast path
tp_size=model_runner.ps.tp_size,
tp_rank=model_runner.ps.tp_rank,
tp_size=model_runner.tp_size,
tp_rank=model_runner.tp_rank,
# Infer these from the one trained adapter.
max_lora_rank=None,
target_modules=None,
@@ -46,7 +46,6 @@ from sglang.srt.utils.common import (
)
if TYPE_CHECKING:
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.server_args import ServerArgs
@@ -62,7 +61,6 @@ class UnoWorkerV2(BaseSpecWorker):
self,
server_args: ServerArgs,
gpu_id: int,
ps: ParallelState,
nccl_port: int,
target_worker: TpModelWorker,
):
@@ -70,7 +68,6 @@ class UnoWorkerV2(BaseSpecWorker):
self.server_args = server_args
self.gpu_id = gpu_id
self.ps = ps
self.nccl_port = nccl_port
self._target_worker = target_worker
+2 -7
View File
@@ -8,7 +8,6 @@ from typing import Callable, Dict, List, Optional
import torch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import ProfileReqOutput
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
@@ -71,14 +70,13 @@ def export_cuda_graph_capture_trace(prof_context, *, runner_name: str, tp_rank:
class ProfileManager:
def __init__(self, ps: ParallelState, cpu_group):
def __init__(self, cpu_group):
self.stage_based_trigger = _StageBasedTrigger(
on_start=self._do_start,
on_stop=self._do_stop,
)
self.ps = ps
self.cpu_group = cpu_group
self.first_rank_in_node = ps.gpu_id == get_device().base_gpu_id
self.first_rank_in_node = get_device().gpu_id == get_device().base_gpu_id
self.profiler_kwargs = None
self.profiler = None
self.detailed_annotations = False
@@ -154,7 +152,6 @@ class ProfileManager:
set_detailed_annotations_enabled(self.detailed_annotations)
self.profiler = _ProfilerBase.create(
**self.profiler_kwargs,
ps=self.ps,
cpu_group=self.cpu_group,
first_rank_in_node=self.first_rank_in_node,
output_suffix=f"-{stage}" if stage else "",
@@ -303,7 +300,6 @@ class _ProfilerConcreteBase(_ProfilerBase):
output_prefix: str,
output_suffix: str,
profile_id: str,
ps: ParallelState,
cpu_group,
first_rank_in_node: bool,
):
@@ -311,7 +307,6 @@ class _ProfilerConcreteBase(_ProfilerBase):
self.output_prefix = output_prefix
self.output_suffix = output_suffix
self.profile_id = profile_id
self.ps = ps
self.cpu_group = cpu_group
self.first_rank_in_node = first_rank_in_node
+2 -3
View File
@@ -274,15 +274,14 @@ class WeightCacheDaemon:
from sglang.srt.model_loader.loader import get_model_loader
server_args = self.server_args
# The launcher told this daemon where it sits, and it builds the same
# groups a scheduler does, so the same one number places it.
publish(
server_args,
role="weight_cache_daemon",
ranks=SpawnRanks(
world_rank=spawn_world_rank(
server_args, tp_rank=self.tp_rank, pp_rank=self.pp_rank
)
),
gpu_id=self.gpu_id,
),
)
+12 -12
View File
@@ -503,17 +503,17 @@ class IpcModelLoader(BaseModelLoader):
# Build engine's config fingerprint
from sglang.srt.layers.dp_attention import get_moe_cp_size
ps = get_parallel()
tp_size = ps.tp_size
tp_rank = ps.tp_rank
parallel = get_parallel()
tp_size = parallel.tp_size
tp_rank = parallel.tp_rank
pp_size = ps.pp_size
pp_rank = ps.pp_rank
pp_size = parallel.pp_size
pp_rank = parallel.pp_rank
ep_size = ps.moe_ep_size
ep_size = parallel.moe_ep_size
moe_dp_size = get_moe_cp_size()
moe_dp_rank = ps.moe_dp_rank
moe_ep_rank = ps.moe_ep_rank
moe_dp_rank = parallel.moe_dp_rank
moe_ep_rank = parallel.moe_ep_rank
dp_size = get_parallel().dp_size
@@ -535,10 +535,10 @@ class IpcModelLoader(BaseModelLoader):
moe_dp_size=moe_dp_size,
moe_dp_rank=moe_dp_rank,
moe_ep_rank=moe_ep_rank,
enable_dp_attention=ps.enable_dp_attention,
enable_dp_lm_head=ps.enable_dp_lm_head,
attn_cp_size=ps.attn_cp_size,
moe_dense_tp_size=ps.moe_dense_tp_size,
enable_dp_attention=parallel.enable_dp_attention,
enable_dp_lm_head=parallel.enable_dp_lm_head,
attn_cp_size=parallel.attn_cp_size,
moe_dense_tp_size=parallel.moe_dense_tp_size,
moe_a2a_backend=get_exec().moe.moe_a2a_backend,
quant_method=quant_method,
quant_config_hash=hash_quant_config(quant_config),
@@ -7,7 +7,6 @@ import torch.nn.functional as F
from torch import nn
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, ReqToTokenPool
@@ -339,7 +338,6 @@ class MockModelRunner(ModelRunner):
self.tp_size = 1
self.dp_size = 1
self.pp_size = 1
self.ps = ParallelState.trivial()
self.is_draft_worker = False
self.max_running_requests = pool_batch_size
# trtllm_mha __init__ scans model.modules() for ENCODER_ONLY layers;
@@ -5,7 +5,6 @@ from typing import Any
import torch
from torch import nn
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool, ReqToTokenPool
@@ -322,7 +321,6 @@ class DSAMockModelRunner(ModelRunner):
self._kernel_warmed_up = True
self.dp_size = 1
self.pp_size = 1
self.ps = ParallelState.trivial()
self._server_args_override = get_context().override_server_args(
attention_backend=case.backend,
chunked_prefill_size=-1,
@@ -22,7 +22,6 @@ from torch import nn
from sglang.kernels.ops.attention.dsv4.quant_k_cache import (
quant_to_nope_fp8_rope_bf16_pack_triton,
)
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.radix_attention import RadixAttention
@@ -349,7 +348,6 @@ class MockDSV4ModelRunner:
self.tp_size = 1
self.dp_size = 1
self.pp_size = 1
self.ps = ParallelState.trivial()
self._server_args_override = get_context().override_server_args(
attention_backend=case.backend,
chunked_prefill_size=-1,
@@ -10,7 +10,6 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateShape,
)
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
HybridLinearAttnBackend,
@@ -228,7 +227,6 @@ class MockGDNModelRunner(ModelRunner):
self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.spec_algorithm = SpeculativeAlgorithm.NONE
self.canary_manager = None
self.page_size = case.page_size
@@ -10,7 +10,6 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateDType,
)
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
HybridLinearAttnBackend,
@@ -231,7 +230,6 @@ class MockKDAModelRunner(ModelRunner):
self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.spec_algorithm = SpeculativeAlgorithm.NONE
self.canary_manager = None
self.page_size = case.page_size
@@ -10,7 +10,6 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateShape,
)
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.linear.lightning_backend import (
LightningAttentionBackend,
@@ -239,7 +238,6 @@ class MockLightningModelRunner(ModelRunner):
self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.spec_algorithm = SpeculativeAlgorithm.NONE
self.canary_manager = None
self.page_size = case.page_size
@@ -26,7 +26,6 @@ from sglang.srt.configs.mamba_utils import ( # noqa: E402
Mamba2StateShape,
)
from sglang.srt.configs.model_config import AttentionArch # noqa: E402
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ( # noqa: E402
ATTENTION_BACKENDS,
)
@@ -325,7 +324,6 @@ class MockMamba2ModelRunner(ModelRunner):
self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.spec_algorithm = SpeculativeAlgorithm.NONE
self.canary_manager = None
self.page_size = case.page_size
@@ -7,7 +7,6 @@ import torch.nn.functional as F
from torch import nn
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool, ReqToTokenPool
@@ -247,7 +246,6 @@ class MockMLAModelRunner(ModelRunner):
self.tp_size = 1
self.dp_size = 1
self.pp_size = 1
self.ps = ParallelState.trivial()
self.spec_algorithm = SpeculativeAlgorithm.NONE
speculative_num_draft_tokens = (
max(case.input_lens)
@@ -2,6 +2,8 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Dict, Iterator, List, Optional
from sglang.srt.runtime_context import get_parallel
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
from sglang.test.scripted_runtime.context.api import ScriptedContext
@@ -12,7 +14,7 @@ def _get_all_reqs(ctx: ScriptedContext) -> Iterator[Req]:
if s.chunked_req is not None:
yield s.chunked_req
yield from s.waiting_queue
if s.ps.pp_size > 1:
if get_parallel().pp_size > 1:
for mb in (*s.mbs, *s.last_mbs, *s.running_mbs):
if mb is not None:
yield from mb.reqs
@@ -13,6 +13,7 @@ import zmq
from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils.network import get_zmq_socket
from sglang.test.scripted_runtime.background_http_poster import BackgroundHttpPoster
from sglang.test.scripted_runtime.context import ScriptedContext
@@ -125,9 +126,9 @@ class ScriptedSchedulerHook:
) -> None:
self.scheduler = scheduler
self._is_driver = (
scheduler.ps.pp_rank == 0
and scheduler.ps.tp_rank == 0
and scheduler.ps.attn_cp_rank == 0
get_parallel().pp_rank == 0
and get_parallel().tp_rank == 0
and get_parallel().attn_cp_rank == 0
)
self._batch_log: List[ScriptedBatchRecord] = []