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_distributed_environment,
destroy_model_parallel, 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.entrypoints.engine import _set_envs_and_config
from sglang.srt.hardware_backend.mlx.runtime import use_mlx 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.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, 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( runner_kwargs = dict(
model_config=model_config, model_config=model_config,
mem_fraction_static=cfg.mem_fraction_static, mem_fraction_static=cfg.mem_fraction_static,
gpu_id=gpu_id, gpu_id=gpu_id,
ps=ps,
nccl_port=port_args.nccl_port, nccl_port=port_args.nccl_port,
server_args=server_args, server_args=server_args,
) )
@@ -571,7 +548,7 @@ def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
model_runner=model_runner, model_runner=model_runner,
dp_size=get_parallel().dp_size, dp_size=get_parallel().dp_size,
attn_tp_size=get_parallel().attn_tp_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, tp_group=model_runner.tp_group,
get_idle_batch=None, get_idle_batch=None,
disable_cuda_graph=cuda_graph_fully_disabled(), disable_cuda_graph=cuda_graph_fully_disabled(),
@@ -709,13 +686,12 @@ def correctness_test(
gpu_id, gpu_id,
tp_rank, 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( publish(
server_args, server_args,
role="scheduler", role="scheduler",
ranks=SpawnRanks( 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, server_args,
role="scheduler", role="scheduler",
ranks=SpawnRanks( 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() initialize_moe_config()
@@ -92,6 +92,9 @@ class Arg(msgspec.Struct, frozen=True):
fallback: Any = None fallback: Any = None
_NO_DEFAULT = object()
class Derived(msgspec.Struct, frozen=True): class Derived(msgspec.Struct, frozen=True):
"""Metadata for a field the configuration implies, not one anyone types. """Metadata for a field the configuration implies, not one anyone types.
@@ -122,6 +125,12 @@ class Derived(msgspec.Struct, frozen=True):
doc: str = "" doc: str = ""
fn: 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): class NS(msgspec.Struct, frozen=True):
+12 -1
View File
@@ -17,7 +17,7 @@ from typing import (
import msgspec 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): class Device(msgspec.Struct):
@@ -40,6 +40,17 @@ class Device(msgspec.Struct):
int, int,
"The delta between consecutive GPU IDs that are used. For example, setting it to 2 will use GPU 0,2,4,...", "The delta between consecutive GPU IDs that are used. For example, setting it to 2 will use GPU 0,2,4,...",
] = 1 ] = 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 random_seed: A[Optional[int], "The random seed."] = None
mlx_enable_sampling: A[ mlx_enable_sampling: A[
bool, bool,
+4 -3
View File
@@ -111,6 +111,7 @@ from sglang.srt.observability.scheduler_stage_metrics import (
scheduler_stage_method, scheduler_stage_method,
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_device,
get_disagg, get_disagg,
get_memory, get_memory,
get_parallel, get_parallel,
@@ -439,7 +440,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
if get_disagg().disaggregation_enable_kv_checksum: if get_disagg().disaggregation_enable_kv_checksum:
kv_args = self.kv_manager.kv_args kv_args = self.kv_manager.kv_args
self.scheduler.kv_checksum_computer = KvChecksumComputer( 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_data_ptrs=kv_args.kv_data_ptrs,
kv_item_lens=kv_args.kv_item_lens, kv_item_lens=kv_args.kv_item_lens,
state_data_ptrs=kv_args.state_data_ptrs, 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.engine_rank = self.tp_rank % (attn_tp_size)
kv_args.pp_rank = self.pp_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.kv_cache_dtype_str = ( kv_args.kv_cache_dtype_str = (
self.scheduler.tp_worker.model_runner.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.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_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
kv_manager = kv_manager_class( kv_manager = kv_manager_class(
kv_args, 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.cache import media_preprocess_kwargs
from sglang.srt.multimodal.transport import determine_tensor_transport_mode from sglang.srt.multimodal.transport import determine_tensor_transport_mode
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_device,
get_disagg, get_disagg,
get_exec, get_exec,
get_mm, get_mm,
@@ -1921,7 +1922,7 @@ class MMReceiverBase(ABC):
self.scheduler_embedding_port, self.scheduler_embedding_port,
) )
self.scheduler = scheduler 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.wait_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get()
self.embedding_pool = None self.embedding_pool = None
+4 -3
View File
@@ -87,6 +87,7 @@ from sglang.srt.observability.scheduler_stage_metrics import (
scheduler_stage_method, scheduler_stage_method,
) )
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_device,
get_disagg, get_disagg,
get_parallel, get_parallel,
get_schedule, get_schedule,
@@ -212,7 +213,7 @@ class PrefillBootstrapQueue:
if get_disagg().disaggregation_enable_kv_checksum: if get_disagg().disaggregation_enable_kv_checksum:
kv_args = self.kv_manager.kv_args kv_args = self.kv_manager.kv_args
self.scheduler.kv_checksum_computer = KvChecksumComputer( 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_data_ptrs=kv_args.kv_data_ptrs,
kv_item_lens=kv_args.kv_item_lens, kv_item_lens=kv_args.kv_item_lens,
state_data_ptrs=kv_args.state_data_ptrs, state_data_ptrs=kv_args.state_data_ptrs,
@@ -226,7 +227,7 @@ class PrefillBootstrapQueue:
kv_args = kv_args_class() kv_args = kv_args_class()
kv_args.engine_rank = self.tp_rank kv_args.engine_rank = self.tp_rank
kv_args.pp_rank = self.pp_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 = ( kv_args.rust_http_port = (
self.scheduler.rust_server.http_port self.scheduler.rust_server.http_port
if self.scheduler.rust_server is not None if self.scheduler.rust_server is not None
@@ -303,7 +304,7 @@ class PrefillBootstrapQueue:
self.metadata_buffers.get_buf_infos() self.metadata_buffers.get_buf_infos()
) )
kv_args.ib_device = get_disagg().disaggregation_ib_device 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) req_to_token_pool = getattr(self.scheduler, "req_to_token_pool", None)
setup_state_kv_args( 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.common.conn import CommonKVReceiver
from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.managers.io_struct import PdRoleSwitchReqInput, PdRoleSwitchReqOutput 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 from sglang.srt.utils import get_available_gpu_memory
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -95,7 +95,7 @@ def handle_pd_role_switch(
) )
try: try:
available_graph_gb = get_available_gpu_memory( available_graph_gb = get_available_gpu_memory(
scheduler.device, scheduler.ps.gpu_id scheduler.device, get_device().gpu_id
) )
except Exception as e: except Exception as e:
return _fail( 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 ( from sglang.srt.distributed.parallel_state import (
_tag_groups_for_flashinfer_allreduce_only, _tag_groups_for_flashinfer_allreduce_only,
) )
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import initialize_dp_attention from sglang.srt.layers.dp_attention import initialize_dp_attention
from sglang.srt.layers.layernorm_sp import initialize_layernorm_sp from sglang.srt.layers.layernorm_sp import initialize_layernorm_sp
from sglang.srt.platforms import current_platform from sglang.srt.platforms import current_platform
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_device,
get_disagg, get_disagg,
get_exec, get_exec,
get_parallel, get_parallel,
@@ -63,17 +63,17 @@ def init_torch_distributed(
server_args: ServerArgs, server_args: ServerArgs,
model_config: ModelConfig, model_config: ModelConfig,
device: str, device: str,
ps: ParallelState,
dist_port: int, dist_port: int,
is_draft_worker: bool, is_draft_worker: bool,
local_omp_cpuid: Optional[List[int]], local_omp_cpuid: Optional[List[int]],
): ):
tic = time.perf_counter() tic = time.perf_counter()
logger.info("Init torch distributed begin.") logger.info("Init torch distributed begin.")
parallel = get_parallel()
backend = _resolve_backend(device=device) 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: if not get_parallel().enable_p2p_check:
monkey_patch_p2p_access_check() monkey_patch_p2p_access_check()
@@ -83,8 +83,8 @@ def init_torch_distributed(
if not is_draft_worker: if not is_draft_worker:
if device == "cpu": if device == "cpu":
_init_cpu_threads_env( _init_cpu_threads_env(
tp_size=ps.tp_size, tp_size=parallel.tp_size,
tp_rank=ps.tp_rank, tp_rank=parallel.tp_rank,
local_omp_cpuid=local_omp_cpuid, local_omp_cpuid=local_omp_cpuid,
dist_init_method=dist_init_method, dist_init_method=dist_init_method,
) )
@@ -96,16 +96,18 @@ def init_torch_distributed(
dist_init_method=dist_init_method, dist_init_method=dist_init_method,
server_args=server_args, server_args=server_args,
model_config=model_config, model_config=model_config,
gpu_id=ps.gpu_id, gpu_id=get_device().gpu_id,
) )
# Pre-warm NCCL/RCCL/HCCL to eliminate cold-start latency in first request # Pre-warm NCCL/RCCL/HCCL to eliminate cold-start latency in first request
# Controlled by --pre-warm-nccl flag (default: enabled on AMD GPUs) # Controlled by --pre-warm-nccl flag (default: enabled on AMD GPUs)
if get_exec().comm.pre_warm_nccl and ( 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( _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 # CUDA graph capture enables the PyNCCL communicator for TP LM-head
@@ -115,7 +117,7 @@ def init_torch_distributed(
if ( if (
device == "cuda" device == "cuda"
and get_parallel().enable_tp_lm_head_all_to_all 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() _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. # including them in this WORLD reduction would deadlock on absent peers.
pre_model_load_memory = get_available_gpu_memory( pre_model_load_memory = get_available_gpu_memory(
device, device,
ps.gpu_id, get_device().gpu_id,
distributed=get_world_group().world_size > 1 and not is_draft_worker, distributed=get_world_group().world_size > 1 and not is_draft_worker,
cpu_group=get_world_group().cpu_group, cpu_group=get_world_group().cpu_group,
) )
# Check memory for tensor parallelism # Check memory for tensor parallelism
local_gpu_memory = get_available_gpu_memory(device, ps.gpu_id) local_gpu_memory = get_available_gpu_memory(device, get_device().gpu_id)
if ps.tp_size > 1 and not is_draft_worker: if parallel.tp_size > 1 and not is_draft_worker:
_check_tp_memory_balance( _check_tp_memory_balance(
pre_model_load_memory=pre_model_load_memory, pre_model_load_memory=pre_model_load_memory,
local_gpu_memory=local_gpu_memory, local_gpu_memory=local_gpu_memory,
@@ -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, self,
*, *,
model_config: ModelConfig, model_config: ModelConfig,
ps: Any,
get_model: Callable[[], nn.Module], get_model: Callable[[], nn.Module],
get_expert_location_updater: Callable[[], ExpertLocationUpdater], get_expert_location_updater: Callable[[], ExpertLocationUpdater],
get_expert_backup_client: Callable[[], Any], get_expert_backup_client: Callable[[], Any],
@@ -42,7 +41,6 @@ class EPLBManager:
# constructed (model load, expert_backup_client, weight_updater), so # constructed (model load, expert_backup_client, weight_updater), so
# they are read through getters at rebalance time, not captured here. # they are read through getters at rebalance time, not captured here.
self._model_config = model_config self._model_config = model_config
self._ps = ps
self._get_model = get_model self._get_model = get_model
self._get_expert_location_updater = get_expert_location_updater self._get_expert_location_updater = get_expert_location_updater
self._get_expert_backup_client = get_expert_backup_client self._get_expert_backup_client = get_expert_backup_client
@@ -163,7 +161,7 @@ class EPLBManager:
tp_rank=( tp_rank=(
self._elastic_global_rank() self._elastic_global_rank()
if is_post_scale_rebalance if is_post_scale_rebalance
else self._ps.tp_rank else get_parallel().tp_rank
), ),
use_flat_topology=is_post_scale_rebalance, use_flat_topology=is_post_scale_rebalance,
expert_backup_client=self._get_expert_backup_client(), expert_backup_client=self._get_expert_backup_client(),
@@ -223,7 +221,7 @@ class EPLBManager:
) )
def _elastic_global_rank(self) -> int: 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): def _check_rebalance_needed(self, average_utilization_rate_over_window):
if average_utilization_rate_over_window is None: 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)) return list(_chunk_list(all_layer_ids, chunk_size=chunk_size))
def _should_log_expert_location_metadata(self) -> bool: 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( def _log_rebalance_layout_before_update(
self, self,
@@ -172,7 +172,7 @@ class MlxModelRunnerStub(ModelRunner):
aux_state_size = get_schedule().max_mamba_cache_size aux_state_size = get_schedule().max_mamba_cache_size
if aux_state_size is None: if aux_state_size is None:
return 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: def _resolve_max_running_requests(self) -> int:
"""Concurrency cap handed to the scheduler. """Concurrency cap handed to the scheduler.
@@ -197,7 +197,7 @@ class MlxModelRunnerStub(ModelRunner):
requested_per_worker = None requested_per_worker = None
resolved = min(capacity_cap, 4096) resolved = min(capacity_cap, 4096)
else: 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) resolved = min(requested_per_worker, capacity_cap)
aux_state_size = self._explicit_aux_state_size_per_worker() aux_state_size = self._explicit_aux_state_size_per_worker()
@@ -209,7 +209,7 @@ class MlxModelRunnerStub(ModelRunner):
resolved = min(resolved, aux_state_size // ratio) resolved = min(resolved, aux_state_size // ratio)
if resolved <= 0: if resolved <= 0:
global_aux_state_size = get_schedule().max_mamba_cache_size 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( raise RuntimeError(
f"MLX auxiliary-state cache is too small to serve any " f"MLX auxiliary-state cache is too small to serve any "
f"requests: max_mamba_cache_size={global_aux_state_size} " f"requests: max_mamba_cache_size={global_aux_state_size} "
@@ -108,7 +108,6 @@ class MlxTpModelWorker(TpModelWorker):
model_config=self.model_config, model_config=self.model_config,
mem_fraction_static=get_schedule().mem_fraction_static, mem_fraction_static=get_schedule().mem_fraction_static,
gpu_id=self.gpu_id, gpu_id=self.gpu_id,
ps=self.ps,
nccl_port=self.nccl_port, nccl_port=self.nccl_port,
server_args=self.server_args, server_args=self.server_args,
is_draft_worker=self.is_draft_worker, is_draft_worker=self.is_draft_worker,
@@ -428,7 +428,7 @@ class AscendAttnBackend(AttentionBackend):
self.is_dllm_model = True self.is_dllm_model = True
self.dllm_block_size = self.dllm_config.block_size self.dllm_block_size = self.dllm_config.block_size
self.attn_cp_size = model_runner.ps.attn_cp_size self.attn_cp_size = model_runner.attn_cp_size
def _is_swa_layer(self, layer: RadixAttention) -> bool: def _is_swa_layer(self, layer: RadixAttention) -> bool:
return ( return (
@@ -203,7 +203,7 @@ class FlashAttentionBackend(AttentionBackend):
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
self.kv_index_translator = model_runner.kv_index_translator self.kv_index_translator = model_runner.kv_index_translator
self.skip_prefill = skip_prefill 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 self._verify_mask = None
# The worker fetches the tree-mask scratch from the target backend # The worker fetches the tree-mask scratch from the target backend
# only; draft-side instances must not allocate it. # 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.head_dim = model_runner.model_config.head_dim
self.num_attention_heads = ( self.num_attention_heads = (
model_runner.model_config.hf_text_config.num_attention_heads model_runner.model_config.hf_text_config.num_attention_heads
// model_runner.ps.tp_size // model_runner.tp_size
) )
self.num_kv_heads = model_runner.model_config.get_num_kv_heads( self.num_kv_heads = model_runner.model_config.get_num_kv_heads(
model_runner.ps.tp_size model_runner.tp_size
) )
_softcapping = getattr( _softcapping = getattr(
model_runner.model_config.hf_text_config, "attn_logit_softcapping", None model_runner.model_config.hf_text_config, "attn_logit_softcapping", None
@@ -162,9 +162,8 @@ class HPCOpsAttnBackend(AttentionBackend):
self.use_fp8 = model_runner.kv_cache_dtype == torch.float8_e4m3fn self.use_fp8 = model_runner.kv_cache_dtype == torch.float8_e4m3fn
if self.use_fp8: if self.use_fp8:
heads = ( heads = (
model_runner.model_config.num_attention_heads model_runner.model_config.num_attention_heads // model_runner.tp_size,
// model_runner.ps.tp_size, model_runner.model_config.get_num_kv_heads(model_runner.tp_size),
model_runner.model_config.get_num_kv_heads(model_runner.ps.tp_size),
) )
if heads not in FP8_ROPE_SUPPORTED_HEAD_CONFIGS: if heads not in FP8_ROPE_SUPPORTED_HEAD_CONFIGS:
raise ValueError( raise ValueError(
@@ -177,8 +176,8 @@ class HPCOpsAttnBackend(AttentionBackend):
config = model_runner.model_config config = model_runner.model_config
head_dim = config.head_dim head_dim = config.head_dim
num_q_heads = config.num_attention_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.ps.tp_size) num_kv_heads = config.get_num_kv_heads(model_runner.tp_size)
gqa_group_size = num_q_heads // num_kv_heads gqa_group_size = num_q_heads // num_kv_heads
if head_dim != _SUPPORTED_HEAD_DIM or gqa_group_size not in ( if head_dim != _SUPPORTED_HEAD_DIM or gqa_group_size not in (
_SUPPORTED_GQA_GROUP_SIZES _SUPPORTED_GQA_GROUP_SIZES
@@ -66,7 +66,7 @@ class WaveAttnBackend(AttentionBackend):
import wave_lang.kernel.wave.cache as cache import wave_lang.kernel.wave.cache as cache
base_cache_dir = cache.CACHE_BASE_DIR base_cache_dir = cache.CACHE_BASE_DIR
new_dir = base_cache_dir / f"worker_{model_runner.ps.tp_rank}" new_dir = base_cache_dir / f"worker_{model_runner.tp_rank}"
logger.info(f"Setting Wave cache dir: {new_dir}") logger.info(f"Setting Wave cache dir: {new_dir}")
cache.CACHE_BASE_DIR = new_dir cache.CACHE_BASE_DIR = new_dir
@@ -67,7 +67,7 @@ class XPUAttentionBackend(AttentionBackend):
self.num_attention_heads = ( self.num_attention_heads = (
model_runner.model_config.hf_text_config.num_attention_heads model_runner.model_config.hf_text_config.num_attention_heads
) )
self.tp_size = model_runner.ps.tp_size self.tp_size = model_runner.tp_size
assert self.num_attention_heads % self.tp_size == 0 assert self.num_attention_heads % self.tp_size == 0
self.num_local_heads = self.num_attention_heads // self.tp_size self.num_local_heads = self.num_attention_heads // self.tp_size
self.device = model_runner.device self.device = model_runner.device
@@ -510,7 +510,7 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
"PP-parallel DeepGEMM warmup start " "PP-parallel DeepGEMM warmup start "
"(pp_rank=%d, tp_rank=%d, batch_sizes=%s, disagg=%s).", "(pp_rank=%d, tp_rank=%d, batch_sizes=%s, disagg=%s).",
get_parallel().pp_rank, get_parallel().pp_rank,
model_runner.ps.tp_rank, model_runner.tp_rank,
batch_sizes, batch_sizes,
disagg_mode, disagg_mode,
) )
+18
View File
@@ -47,6 +47,24 @@ if TYPE_CHECKING:
from sglang.srt.model_executor.forward_batch_info import ForwardBatch 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: def dp_gather_width() -> int:
"""How many replicas the DP sync gathers over. """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 ( from sglang.srt.distributed.parallel_state import (
abort_distributed_environment, abort_distributed_environment,
) )
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.dllm.mixin.scheduler import SchedulerDllmMixin from sglang.srt.dllm.mixin.scheduler import SchedulerDllmMixin
from sglang.srt.environ import envs, exportable_env_vars from sglang.srt.environ import envs, exportable_env_vars
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder 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.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.moe import initialize_moe_config
from sglang.srt.layers.quantization.fp4_utils import initialize_fp4_gemm_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 from sglang.srt.layers.quantization.fp8_utils import initialize_fp8_gemm_config
@@ -449,12 +447,8 @@ class Scheduler(
self, self,
server_args: ServerArgs, server_args: ServerArgs,
port_args: PortArgs, port_args: PortArgs,
gpu_id: int,
tp_rank: int, tp_rank: int,
moe_ep_rank: int,
pp_rank: int, pp_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
dp_rank: Optional[int], dp_rank: Optional[int],
): ):
# NOTE: KEEP THE FOLLOWING CODE STYLE for this function: # 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_dp_attention = get_parallel().enable_dp_attention
self.enable_unified_memory = get_memory().enable_unified_memory 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 # Init model configs
self.init_model_config() self.init_model_config()
@@ -779,7 +741,7 @@ class Scheduler(
if get_parallel().pp_size > 1: if get_parallel().pp_size > 1:
logger.error("only zbal mix mode support pp_size > 1!") logger.error("only zbal mix mode support pp_size > 1!")
init_zbal( 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 ) # only switch allocator if is mix mode
def init_model_config(self): def init_model_config(self):
@@ -1009,8 +971,7 @@ class Scheduler(
def init_tp_model_worker(self): def init_tp_model_worker(self):
worker_kwargs = dict( worker_kwargs = dict(
server_args=self.server_args, server_args=self.server_args,
gpu_id=self.ps.gpu_id, gpu_id=get_device().gpu_id,
ps=self.ps,
nccl_port=self.nccl_port, nccl_port=self.nccl_port,
) )
@@ -1047,8 +1008,7 @@ class Scheduler(
# — is resolved per runner, not on a config copy. # — is resolved per runner, not on a config copy.
draft_worker_kwargs = dict( draft_worker_kwargs = dict(
server_args=self.server_args, server_args=self.server_args,
gpu_id=self.ps.gpu_id, gpu_id=get_device().gpu_id,
ps=self.ps,
nccl_port=self.nccl_port, nccl_port=self.nccl_port,
target_worker=self.tp_worker, target_worker=self.tp_worker,
) )
@@ -1246,7 +1206,7 @@ class Scheduler(
# Print debug info # Print debug info
self.startup_available_gpu_memory_gb = get_available_gpu_memory( 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: if get_parallel().tp_rank == 0:
logger.info( logger.info(
@@ -1582,7 +1542,7 @@ class Scheduler(
tp_rank=get_parallel().tp_rank, tp_rank=get_parallel().tp_rank,
tp_size=get_parallel().tp_size, tp_size=get_parallel().tp_size,
dp_size=get_parallel().dp_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, bootstrap_port=get_disagg().disaggregation_bootstrap_port,
max_total_num_tokens=self.max_total_num_tokens, max_total_num_tokens=self.max_total_num_tokens,
pp_rank=get_parallel().pp_rank, pp_rank=get_parallel().pp_rank,
@@ -1613,7 +1573,7 @@ class Scheduler(
metadata_buffers=self.disagg_metadata_buffers, metadata_buffers=self.disagg_metadata_buffers,
tp_rank=get_parallel().tp_rank, tp_rank=get_parallel().tp_rank,
tp_size=get_parallel().tp_size, 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, bootstrap_port=get_disagg().disaggregation_bootstrap_port,
gloo_group=self.attn_tp_cpu_group, gloo_group=self.attn_tp_cpu_group,
max_total_num_tokens=self.max_total_num_tokens, max_total_num_tokens=self.max_total_num_tokens,
@@ -2248,7 +2208,6 @@ class Scheduler(
def init_profiler(self) -> None: def init_profiler(self) -> None:
self.profiler_manager = SchedulerProfilerManager( self.profiler_manager = SchedulerProfilerManager(
ps=self.ps,
dp_tp_cpu_group=self.dp_tp_cpu_group, dp_tp_cpu_group=self.dp_tp_cpu_group,
get_forward_ct=lambda: self.forward_ct, get_forward_ct=lambda: self.forward_ct,
) )
@@ -2369,7 +2328,6 @@ class Scheduler(
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tree_cache=self.tree_cache, tree_cache=self.tree_cache,
offload_tags=self.weight_updater.offload_tags, offload_tags=self.weight_updater.offload_tags,
ps=self.ps,
model_config=self.model_config, model_config=self.model_config,
enable_overlap=self.enable_overlap, enable_overlap=self.enable_overlap,
spec_algorithm=self.spec_algorithm, spec_algorithm=self.spec_algorithm,
@@ -2457,7 +2415,6 @@ class Scheduler(
self._sched_idled = False self._sched_idled = False
self.load_inquirer = SchedulerLoadInquirer( self.load_inquirer = SchedulerLoadInquirer(
disaggregation_mode=self.disaggregation_mode, disaggregation_mode=self.disaggregation_mode,
ps=self.ps,
server_args=self.server_args, server_args=self.server_args,
max_total_num_tokens=self.max_total_num_tokens, max_total_num_tokens=self.max_total_num_tokens,
max_running_requests=self.max_running_requests, max_running_requests=self.max_running_requests,
@@ -2500,7 +2457,6 @@ class Scheduler(
self.output_streamer = self.get_output_streamer_class()( self.output_streamer = self.get_output_streamer_class()(
send_to_detokenizer=self.ipc_channels.send_to_detokenizer, send_to_detokenizer=self.ipc_channels.send_to_detokenizer,
tree_cache=self.tree_cache, tree_cache=self.tree_cache,
ps=self.ps,
server_args=self.server_args, server_args=self.server_args,
is_generation=self.is_generation, is_generation=self.is_generation,
spec_algorithm=self.spec_algorithm, spec_algorithm=self.spec_algorithm,
@@ -3120,7 +3076,7 @@ class Scheduler(
self._add_request_to_queue(req) self._add_request_to_queue(req)
return 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 self.tree_cache.cache_controller, "pp_prefetch_command_group", None
): ):
recv_req.pp_prefetch_ticketed = bool(self._prefetch_kvcache(req)) recv_req.pp_prefetch_ticketed = bool(self._prefetch_kvcache(req))
@@ -6078,6 +6034,7 @@ def run_scheduler_process(
ranks=SpawnRanks( ranks=SpawnRanks(
world_rank=spawn_world_rank(server_args, tp_rank=tp_rank, pp_rank=pp_rank), world_rank=spawn_world_rank(server_args, tp_rank=tp_rank, pp_rank=pp_rank),
dp_rank=dp_rank, dp_rank=dp_rank,
gpu_id=gpu_id,
), ),
) )
configure_scheduler_process( configure_scheduler_process(
@@ -6115,12 +6072,8 @@ def run_scheduler_process(
scheduler = Scheduler( scheduler = Scheduler(
server_args, server_args,
port_args, port_args,
gpu_id,
tp_rank, tp_rank,
moe_ep_rank,
pp_rank, pp_rank,
attn_cp_rank,
moe_dp_rank,
dp_rank, dp_rank,
) )
@@ -7,7 +7,6 @@ import torch
from sglang.srt.batch_overlap.two_batch_overlap import TboDPAttentionPreparer from sglang.srt.batch_overlap.two_batch_overlap import TboDPAttentionPreparer
from sglang.srt.configs.model_config import ModelConfig 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.environ import envs
from sglang.srt.layers.cp.utils import get_cp_strategy from sglang.srt.layers.cp.utils import get_cp_strategy
from sglang.srt.layers.dp_attention import dp_gather_width, world_dp_gather_enabled 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 token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator
tree_cache: BasePrefixCache tree_cache: BasePrefixCache
offload_tags: set[str] offload_tags: set[str]
ps: ParallelState
model_config: ModelConfig model_config: ModelConfig
enable_overlap: bool enable_overlap: bool
spec_algorithm: SpeculativeAlgorithm spec_algorithm: SpeculativeAlgorithm
@@ -17,7 +17,6 @@ from sglang.srt.managers.load_snapshot import (
from sglang.srt.runtime_context import get_lora, get_parallel from sglang.srt.runtime_context import get_lora, get_parallel
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.scheduler_components.pool_stats_observer import ( from sglang.srt.managers.scheduler_components.pool_stats_observer import (
SchedulerPoolStatsObserver, SchedulerPoolStatsObserver,
) )
@@ -33,7 +32,6 @@ logger = logging.getLogger(__name__)
@dataclass(kw_only=True, slots=True, frozen=True) @dataclass(kw_only=True, slots=True, frozen=True)
class SchedulerLoadInquirer: class SchedulerLoadInquirer:
disaggregation_mode: DisaggregationMode disaggregation_mode: DisaggregationMode
ps: ParallelState
server_args: ServerArgs server_args: ServerArgs
max_total_num_tokens: int max_total_num_tokens: int
max_running_requests: int max_running_requests: int
@@ -308,7 +308,7 @@ class SchedulerMetricsReporter:
self.scheduler.enable_fpm = False self.scheduler.enable_fpm = False
if ( if (
get_observability().enable_forward_pass_metrics 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 and get_parallel().pp_rank == get_parallel().pp_size - 1
): ):
from sglang.srt.observability.forward_pass_metrics import ( from sglang.srt.observability.forward_pass_metrics import (
@@ -316,9 +316,7 @@ class SchedulerMetricsReporter:
) )
self.scheduler._fpm_dp_rank = ( self.scheduler._fpm_dp_rank = (
self.scheduler.ps.dp_rank get_parallel().dp_rank if get_parallel().dp_rank is not None else 0
if self.scheduler.ps.dp_rank is not None
else 0
) )
self.scheduler._fpm_worker_id = ( self.scheduler._fpm_worker_id = (
get_observability().forward_pass_metrics_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)) num_layers = float(getattr(model_config, "num_attention_layers", 0))
head_dim = float(getattr(model_config, "head_dim", 0)) head_dim = float(getattr(model_config, "head_dim", 0))
num_attn_heads = float( 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) intermediate_size = getattr(hf_text_config, "intermediate_size", None)
if intermediate_size is None: if intermediate_size is None:
intermediate_size = getattr(hf_text_config, "ffn_hidden_size", 0) 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, pack_beam_search_output,
) )
from sglang.srt.disaggregation.utils import DisaggregationMode 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.environ import envs
from sglang.srt.managers.io_struct import ( from sglang.srt.managers.io_struct import (
BatchEmbeddingOutput, BatchEmbeddingOutput,
@@ -53,7 +52,6 @@ class SchedulerOutputStreamer:
send_to_detokenizer: zmq.Socket send_to_detokenizer: zmq.Socket
tree_cache: BasePrefixCache tree_cache: BasePrefixCache
ps: ParallelState
server_args: ServerArgs server_args: ServerArgs
is_generation: bool is_generation: bool
spec_algorithm: SpeculativeAlgorithm spec_algorithm: SpeculativeAlgorithm
@@ -51,14 +51,12 @@ logger = logging.getLogger(__name__)
@dataclass(kw_only=True) @dataclass(kw_only=True)
class SchedulerProfilerManager: class SchedulerProfilerManager:
ps: Any
dp_tp_cpu_group: Any dp_tp_cpu_group: Any
get_forward_ct: Callable[[], int] get_forward_ct: Callable[[], int]
def __post_init__(self) -> None: def __post_init__(self) -> None:
if envs.SGLANG_PROFILE_V2.get(): if envs.SGLANG_PROFILE_V2.get():
self._profile_manager = ProfileManager( self._profile_manager = ProfileManager(
ps=self.ps,
cpu_group=self.dp_tp_cpu_group, cpu_group=self.dp_tp_cpu_group,
) )
return return
@@ -274,7 +272,7 @@ class SchedulerProfilerManager:
self.profile_in_progress = True self.profile_in_progress = True
if "CUDA_PROFILER" in activities: 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() torch.cuda.cudart().cudaProfilerStart()
self.profile_in_progress = True self.profile_in_progress = True
@@ -387,7 +385,7 @@ class SchedulerProfilerManager:
torch.cuda.memory._record_memory_history(enabled=None) torch.cuda.memory._record_memory_history(enabled=None)
if "CUDA_PROFILER" in self.profiler_activities: 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() torch.cuda.cudart().cudaProfilerStop()
merge_message = self._merge_profile_traces() merge_message = self._merge_profile_traces()
+7 -11
View File
@@ -22,7 +22,6 @@ from typing import TYPE_CHECKING, List, Optional, Tuple
import torch import torch
from sglang.srt.beam_search.logits_capture import capture_pre_sample_logits 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.environ import envs
from sglang.srt.managers.io_struct import ( from sglang.srt.managers.io_struct import (
DestroyWeightsUpdateGroupReqInput, DestroyWeightsUpdateGroupReqInput,
@@ -216,12 +215,12 @@ class BaseTpWorker(ABC):
return success, message return success, message
def _deserialize_own_rank(self, serialized_named_tensors): 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 deserializing another rank's copy would break producer-side CUDA-IPC
refcounting.""" refcounting."""
monkey_patch_torch_reductions() monkey_patch_torch_reductions()
return MultiprocessingSerializer.deserialize( 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): 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] extra = [n for n in tensors if n not in exp]
if mismatch or missing or extra: if mismatch or missing or extra:
raise RuntimeError( 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(mismatch)} value-diff {mismatch[:5]}, {len(missing)} missing {missing[:5]}, "
f"{len(extra)} extra {extra[:5]}" f"{len(extra)} extra {extra[:5]}"
) )
logger.info( 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( result = self.model_runner.load_lora_adapter_from_tensors(
recv_req.to_ref(), recv_req.to_ref(),
@@ -322,7 +321,6 @@ class TpModelWorker(BaseTpWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
ps: ParallelState,
nccl_port: int, nccl_port: int,
is_draft_worker: bool = False, is_draft_worker: bool = False,
req_to_token_pool: Optional[ReqToTokenPool] = None, req_to_token_pool: Optional[ReqToTokenPool] = None,
@@ -335,7 +333,6 @@ class TpModelWorker(BaseTpWorker):
): ):
# Parse args # Parse args
self.server_args = server_args self.server_args = server_args
self.ps = ps
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.nccl_port = nccl_port self.nccl_port = nccl_port
self.is_draft_worker = is_draft_worker self.is_draft_worker = is_draft_worker
@@ -411,14 +408,15 @@ class TpModelWorker(BaseTpWorker):
tp_group = self.model_runner.tp_group tp_group = self.model_runner.tp_group
self.random_seed = broadcast_pyobj( self.random_seed = broadcast_pyobj(
[get_device().random_seed], [get_device().random_seed],
tp_group.ranks[self.ps.tp_rank], tp_group.ranks[self.model_runner.tp_rank],
tp_group.cpu_group, tp_group.cpu_group,
src=tp_group.ranks[0], src=tp_group.ranks[0],
)[0] )[0]
else: else:
self.random_seed = broadcast_pyobj( self.random_seed = broadcast_pyobj(
[get_device().random_seed], [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, self.world_group.cpu_group,
src=self.world_group.ranks[0], src=self.world_group.ranks[0],
)[0] )[0]
@@ -521,7 +519,6 @@ class TpModelWorker(BaseTpWorker):
model_config=self.model_config, model_config=self.model_config,
mem_fraction_static=get_schedule().mem_fraction_static, mem_fraction_static=get_schedule().mem_fraction_static,
gpu_id=self.gpu_id, gpu_id=self.gpu_id,
ps=self.ps,
nccl_port=self.nccl_port, nccl_port=self.nccl_port,
server_args=self.server_args, server_args=self.server_args,
is_draft_worker=self.is_draft_worker, is_draft_worker=self.is_draft_worker,
@@ -542,7 +539,6 @@ class TpModelWorker(BaseTpWorker):
model_config=self.model_config, model_config=self.model_config,
mem_fraction_static=get_schedule().mem_fraction_static, mem_fraction_static=get_schedule().mem_fraction_static,
gpu_id=self.gpu_id, gpu_id=self.gpu_id,
ps=self.ps,
nccl_port=self.nccl_port, nccl_port=self.nccl_port,
server_args=self.server_args, server_args=self.server_args,
is_draft_worker=self.is_draft_worker, is_draft_worker=self.is_draft_worker,
@@ -1102,13 +1102,12 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
model_runner.lora_manager.prepare_lora_batch(ret) model_runner.lora_manager.prepare_lora_batch(ret)
if ( if (
model_runner.ps.attn_dcp_size > 1 model_runner.attn_dcp_size > 1
and ret.out_cache_loc is not None and ret.out_cache_loc is not None
and is_hip() and is_hip()
): ):
ret.dcp_kv_mask = ( ret.dcp_kv_mask = (
ret.positions % model_runner.ps.attn_dcp_size ret.positions % model_runner.attn_dcp_size == model_runner.attn_dcp_rank
== model_runner.ps.attn_dcp_rank
) )
return ret return ret
@@ -37,7 +37,6 @@ from sglang.srt.distributed import bootstrap
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import ( from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
maybe_init_shared_mooncake_transfer_engine, maybe_init_shared_mooncake_transfer_engine,
) )
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.dllm.config import DllmConfig from sglang.srt.dllm.config import DllmConfig
from sglang.srt.elastic_ep.elastic_ep import ( from sglang.srt.elastic_ep.elastic_ep import (
ElasticEPStateManager, ElasticEPStateManager,
@@ -322,7 +321,6 @@ class ModelRunner:
model_config: ModelConfig, model_config: ModelConfig,
mem_fraction_static: float, mem_fraction_static: float,
gpu_id: int, gpu_id: int,
ps: ParallelState,
nccl_port: int, nccl_port: int,
server_args: ServerArgs, server_args: ServerArgs,
is_draft_worker: bool = False, is_draft_worker: bool = False,
@@ -339,7 +337,6 @@ class ModelRunner:
# `server_args._draft_pool_config` mutation hack). # `server_args._draft_pool_config` mutation hack).
self.memory_pool_config = memory_pool_config self.memory_pool_config = memory_pool_config
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.ps = ps
self.model_config = model_config self.model_config = model_config
self.dist_port = nccl_port self.dist_port = nccl_port
self.server_args = server_args self.server_args = server_args
@@ -416,12 +413,12 @@ class ModelRunner:
# Set device early so that TransferEngine init (e.g. Ascend NPU) # Set device early so that TransferEngine init (e.g. Ascend NPU)
# can access the device context. # can access the device context.
try: 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: except Exception:
import os import os
logger.warning( 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 raise
@@ -741,7 +738,6 @@ class ModelRunner:
self.eplb_manager = ( self.eplb_manager = (
EPLBManager( EPLBManager(
model_config=self.model_config, model_config=self.model_config,
ps=self.ps,
get_model=lambda: self.model, get_model=lambda: self.model,
get_expert_location_updater=lambda: self.expert_location_updater, get_expert_location_updater=lambda: self.expert_location_updater,
get_expert_backup_client=lambda: self.expert_backup_client, get_expert_backup_client=lambda: self.expert_backup_client,
@@ -805,8 +801,8 @@ class ModelRunner:
def get_pp_proxy_dspark_hidden_size(self) -> int: def get_pp_proxy_dspark_hidden_size(self) -> int:
return misc_utils.resolve_pp_proxy_dspark_hidden_size( return misc_utils.resolve_pp_proxy_dspark_hidden_size(
model=self.model, model=self.model,
pp_size=self.ps.pp_size, pp_size=self.pp_size,
pp_rank=self.ps.pp_rank, pp_rank=self.pp_rank,
) )
def get_pp_proxy_topk_size(self) -> Optional[int]: def get_pp_proxy_topk_size(self) -> Optional[int]:
@@ -1175,7 +1171,6 @@ class ModelRunner:
server_args=self.server_args, server_args=self.server_args,
model_config=self.model_config, model_config=self.model_config,
device=self.device, device=self.device,
ps=self.ps,
dist_port=self.dist_port, dist_port=self.dist_port,
is_draft_worker=self.is_draft_worker, is_draft_worker=self.is_draft_worker,
local_omp_cpuid=self.local_omp_cpuid if self.device == "cpu" else None, 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.pp_size = parallel.pp_size
self.attn_cp_rank = parallel.attn_cp_rank self.attn_cp_rank = parallel.attn_cp_rank
self.attn_cp_size = parallel.attn_cp_size 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): def init_shared_mooncake_transfer_engine(self):
maybe_init_shared_mooncake_transfer_engine(gpu_id=self.gpu_id) maybe_init_shared_mooncake_transfer_engine(gpu_id=self.gpu_id)
@@ -220,7 +220,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
self._cell_size == 0 self._cell_size == 0
and mambaish is not None and mambaish is not None
and bool(mambaish.full_attention_layer_ids) and bool(mambaish.full_attention_layer_ids)
and kvc.ps.pp_size > 1 and kvc.pp_size > 1
) )
self._zero_kv_max_tokens = ( self._zero_kv_max_tokens = (
torch.iinfo(torch.int64).max torch.iinfo(torch.int64).max
@@ -877,7 +877,7 @@ class SWAChunkCapPoolConfigurator(HybridSWAPoolConfigurator):
self._swa_cap = compute_swa_request_cap( self._swa_cap = compute_swa_request_cap(
page_size=kvc.page_size, page_size=kvc.page_size,
window=kvc.sliding_window_size, window=kvc.sliding_window_size,
attn_dp_size=kvc.ps.attn_dp_size, attn_dp_size=kvc.attn_dp_size,
) )
@staticmethod @staticmethod
@@ -1015,7 +1015,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
self.compression_ratios = cfg.compress_ratios[ self.compression_ratios = cfg.compress_ratios[
kvc.layer_info.start_layer : kvc.layer_info.end_layer kvc.layer_info.start_layer : kvc.layer_info.end_layer
] ]
if kvc.ps.pp_size > 1: if kvc.pp_size > 1:
logger.info( logger.info(
f"DSV4 pool PP slice: rank={kvc.pp_group.rank_in_group} " f"DSV4 pool PP slice: rank={kvc.pp_group.rank_in_group} "
f"layers=[{kvc.layer_info.start_layer},{kvc.layer_info.end_layer}) " 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.page_size = kvc.page_size
self.is_speculative = get_spec().speculative_algorithm is not None 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.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 = ( 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 if get_schedule().max_running_requests is not None
else None else None
) )
@@ -547,9 +547,9 @@ class BaseRunner(ABC):
if ( if (
capture_forward_mode == ForwardMode.EXTEND capture_forward_mode == ForwardMode.EXTEND
and get_parallel().pp_rank != 0 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( pp_proxy_tensors = PPProxyTensors(
{k: v[:pp_hidden_tokens] for k, v in buffers.pp_proxy_tensors.items()} {k: v[:pp_hidden_tokens] for k, v in buffers.pp_proxy_tensors.items()}
) )
@@ -296,7 +296,7 @@ class EagerRunner(BaseRunner):
or cp_active or cp_active
or forward_batch.forward_mode.is_target_verify() 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" model_runner.model, "prepare_context_parallel_metadata_for_dcp"
): ):
# prepare kv cache buffer for dcp to gather kv cache # 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(mr.dtype),
str(get_model().quantization), str(get_model().quantization),
str(get_exec().moe.moe_runner_backend), str(get_exec().moe.moe_runner_backend),
str(mr.ps.tp_size), str(mr.tp_size),
str(get_parallel().pp_size), str(get_parallel().pp_size),
str(mr.ps.attn_dp_size), str(mr.attn_dp_size),
str(mr.ps.moe_ep_size), str(mr.moe_ep_size),
str(mr.model_config.hf_config.__class__.__name__), str(mr.model_config.hf_config.__class__.__name__),
] ]
# A different skip policy must not reuse previously tuned tactics. # 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) cache_dir.mkdir(parents=True, exist_ok=True)
return ( return (
cache_dir 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 import ray
from sglang.srt.arg_groups.overrides import declare_resolution 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 from sglang.srt.server_args import PortArgs, ServerArgs
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -93,6 +93,7 @@ class SchedulerActor:
server_args, tp_rank=tp_rank, pp_rank=pp_rank server_args, tp_rank=tp_rank, pp_rank=pp_rank
), ),
dp_rank=dp_rank, dp_rank=dp_rank,
gpu_id=actual_gpu_id,
), ),
) )
@@ -124,12 +125,8 @@ class SchedulerActor:
self.scheduler = Scheduler( self.scheduler = Scheduler(
server_args=server_args, server_args=server_args,
port_args=port_args, port_args=port_args,
gpu_id=actual_gpu_id,
tp_rank=tp_rank, tp_rank=tp_rank,
moe_ep_rank=moe_ep_rank,
pp_rank=pp_rank, pp_rank=pp_rank,
attn_cp_rank=attn_cp_rank,
moe_dp_rank=moe_dp_rank,
dp_rank=dp_rank, dp_rank=dp_rank,
) )
@@ -146,7 +143,7 @@ class SchedulerActor:
import torch import torch
# Need to set the GPU id for the event loop for nccl to work # 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() self.scheduler.run_event_loop()
except Exception as e: except Exception as e:
logger.error(f"Scheduler PP{self._pp_rank} TP{self._tp_rank} crashed: {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 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 from sglang.srt.arg_groups.overrides import resolved_view
namespaces = getattr(type(server_args), "_NAMESPACES", None) 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: if path is None:
continue continue
for name, decl in vars(source).items(): 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 continue
module, _, attr = decl.fn.rpartition(".") module, _, attr = decl.fn.rpartition(".")
bag = tops.get(path.split(".")[0]) bag = tops.get(path.split(".")[0])
@@ -1251,7 +1262,20 @@ class RuntimeContext:
# Snapshot resolved config into the namespace bags (the single source of # Snapshot resolved config into the namespace bags (the single source of
# truth for config reads). Placed by `namespace_of`; a mock/partial # truth for config reads). Placed by `namespace_of`; a mock/partial
# config that declares no namespace yields an empty tree (no bags). # 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) 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") spec = self._config_bags.get("spec")
if spec is not None: if spec is not None:
from sglang.srt.arg_groups.overrides import ( 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. # never built. With DCP on it is a position, and the bundle below states it.
if not _CONTEXT.parallel.dcp_enabled: if not _CONTEXT.parallel.dcp_enabled:
_CONTEXT.parallel.override_permanently(attn_dcp_rank=0) _CONTEXT.parallel.override_permanently(attn_dcp_rank=0)
if ranks is not None and ranks.gpu_id is not None: # Stated on the bag directly: `gpu_id` is declared but not configured, so
_CONTEXT.override("spawn", gpu_id=ranks.gpu_id) # 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: if ranks is not None:
# The placement, worked out here rather than carried: the widths are on # 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 # 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 -- 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 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 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. by being a plain frozen record.
These are stamped rather than written as bag leaves because they are 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") os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
# Preserve the DP startup log; ports use node-local offsets. # 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: if get_exec().moe.is_ep_scale_joiner:
# The joining TP group is entirely local to this node. # 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: else:
nnodes_per_pp_rank = max(get_parallel().nnodes // get_parallel().pp_size, 1) 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 tp_size_per_node = get_parallel().tp_size // nnodes_per_pp_rank
dp_group_width = scheduler.ps.attn_tp_size * scheduler.ps.attn_cp_size 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 # 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. # 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_port = get_serving().port + local_dp_rank
listen_addr = NetworkAddress(get_serving().host, listen_port).to_host_port_str() 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 # Under DP every rank runs its own server on its own port, so the rank is
# what tells two otherwise identical startup lines apart. # what tells two otherwise identical startup lines apart.
dp_note = ( 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( logger.info(
"SGLANG_RUST_SERVER enabled, Rust server listen on %s%s", "SGLANG_RUST_SERVER enabled, Rust server listen on %s%s",
@@ -1,7 +1,6 @@
import logging import logging
import math import math
import os import os
from dataclasses import replace
from typing import List, Optional, Tuple from typing import List, Optional, Tuple
import torch import torch
@@ -19,7 +18,6 @@ from sglang.kernels.ops.speculative.dspark.dspark_accept import (
accept_sampling, accept_sampling,
) )
from sglang.srt.configs.hybrid_arch import mambaish_config 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.environ import envs
from sglang.srt.layers.logits_processor import should_apply_lm_head_quant_method from sglang.srt.layers.logits_processor import should_apply_lm_head_quant_method
from sglang.srt.layers.logprob_processor import compute_spec_logprobs from sglang.srt.layers.logprob_processor import compute_spec_logprobs
@@ -364,7 +362,6 @@ class DFlashWorkerV2(BaseSpecWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
ps: ParallelState,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -372,7 +369,6 @@ class DFlashWorkerV2(BaseSpecWorker):
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.ps = ps
self.nccl_port = nccl_port self.nccl_port = nccl_port
self._target_worker = target_worker self._target_worker = target_worker
self.model_runner = target_worker.model_runner self.model_runner = target_worker.model_runner
@@ -416,7 +412,6 @@ class DFlashWorkerV2(BaseSpecWorker):
bundle = build_draft_tp_worker( bundle = build_draft_tp_worker(
server_args=server_args, server_args=server_args,
gpu_id=gpu_id, gpu_id=gpu_id,
ps=replace(ps, pp_rank=0, pp_size=1),
nccl_port=nccl_port, nccl_port=nccl_port,
target_model_config=target_worker.model_runner.model_config, target_model_config=target_worker.model_runner.model_config,
algo_label="DFLASH", algo_label="DFLASH",
@@ -453,7 +448,7 @@ class DFlashWorkerV2(BaseSpecWorker):
validate_domino_runtime( validate_domino_runtime(
device=torch.device(self.device), device=torch.device(self.device),
tp_size=int(get_parallel().tp_group.world_size), 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), target_vocab_size=int(self.model_runner.model_config.vocab_size),
draft_vocab_size=int(self.draft_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), hidden_size=int(self.draft_model.config.hidden_size),
@@ -500,7 +495,7 @@ class DFlashWorkerV2(BaseSpecWorker):
) )
self._maybe_merge_trained_mask_embedding() self._maybe_merge_trained_mask_embedding()
self._cache_full_embed_weight() self._cache_full_embed_weight()
if self.ps.tp_rank == 0: if self.model_runner.tp_rank == 0:
logger.info( logger.info(
"Initialized DFLASH draft runner. attention_backend=%s, model=%s, block_size=%s, draft_window_size=%s, compact_cache=%s", "Initialized DFLASH draft runner. attention_backend=%s, model=%s, block_size=%s, draft_window_size=%s, compact_cache=%s",
bundle.resolved_attention_backend, bundle.resolved_attention_backend,
@@ -650,7 +645,7 @@ class DFlashWorkerV2(BaseSpecWorker):
# shared graph capture/replay; keep the draft eager under dp # shared graph capture/replay; keep the draft eager under dp
# attention. # attention.
capture_decode_cuda_graph = False capture_decode_cuda_graph = False
if self.ps.tp_rank == 0: if self.model_runner.tp_rank == 0:
logger.warning( logger.warning(
"Disable DFLASH draft cuda graph because dp attention " "Disable DFLASH draft cuda graph because dp attention "
"is enabled (draft runs eager)." "is enabled (draft runs eager)."
@@ -795,7 +790,7 @@ class DFlashWorkerV2(BaseSpecWorker):
def _maybe_build_draft_sampler(self): def _maybe_build_draft_sampler(self):
def _eager(reason): 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) logger.info("DFLASH draft greedy head kept eager (reason=%s).", reason)
return None return None
@@ -819,7 +814,7 @@ class DFlashWorkerV2(BaseSpecWorker):
): ):
return _eager("unsupported quantized lm_head") return _eager("unsupported quantized lm_head")
self.draft_model.lm_head = lm_head self.draft_model.lm_head = lm_head
if self.ps.tp_rank == 0: if self.model_runner.tp_rank == 0:
logger.info( logger.info(
"DFLASH selector decode folded into the draft cuda graph " "DFLASH selector decode folded into the draft cuda graph "
"(sampling_enabled=%s).", "(sampling_enabled=%s).",
@@ -843,7 +838,7 @@ class DFlashWorkerV2(BaseSpecWorker):
embed_proj = self.draft_model.embed_proj embed_proj = self.draft_model.embed_proj
if prefix_gru is None or embed_proj is None: if prefix_gru is None or embed_proj is None:
return _eager("Domino projector modules are unavailable") return _eager("Domino projector modules are unavailable")
if self.ps.tp_rank == 0: if self.model_runner.tp_rank == 0:
logger.info( logger.info(
"DFLASH Domino rollout folded into the draft cuda graph (tp=%s).", "DFLASH Domino rollout folded into the draft cuda graph (tp=%s).",
int(tp_group.world_size), int(tp_group.world_size),
@@ -882,7 +877,7 @@ class DFlashWorkerV2(BaseSpecWorker):
return _eager("added vocab") return _eager("added vocab")
num_org = int(shard.num_org_elements) num_org = int(shard.num_org_elements)
org_vocab_start = int(shard.org_vocab_start_index) org_vocab_start = int(shard.org_vocab_start_index)
if self.ps.tp_rank == 0: if self.model_runner.tp_rank == 0:
logger.info( logger.info(
"DFLASH draft greedy head folded into the draft cuda graph (tp=%d).", "DFLASH draft greedy head folded into the draft cuda graph (tp=%d).",
tp_group.world_size, tp_group.world_size,
@@ -908,7 +903,7 @@ class DFlashWorkerV2(BaseSpecWorker):
fused_disable_reason = "draft model does not support fused context KV" fused_disable_reason = "draft model does not support fused context KV"
if fused_disable_reason is not None: if fused_disable_reason is not None:
if self.ps.tp_rank == 0: if self.model_runner.tp_rank == 0:
logger.info( logger.info(
"DFLASH fused KV materialization disabled: %s", "DFLASH fused KV materialization disabled: %s",
fused_disable_reason, fused_disable_reason,
@@ -951,7 +946,7 @@ class DFlashWorkerV2(BaseSpecWorker):
break break
if fused_disable_reason is not None: if fused_disable_reason is not None:
if self.ps.tp_rank == 0: if self.model_runner.tp_rank == 0:
logger.info( logger.info(
"DFLASH fused KV materialization disabled: %s", "DFLASH fused KV materialization disabled: %s",
fused_disable_reason, fused_disable_reason,
@@ -973,7 +968,7 @@ class DFlashWorkerV2(BaseSpecWorker):
max_position_hint=self.target_worker.model_runner.model_config.context_len max_position_hint=self.target_worker.model_runner.model_config.context_len
+ int(self.block_size), + int(self.block_size),
) )
if self.ps.tp_rank == 0: if self.model_runner.tp_rank == 0:
logger.info( logger.info(
"DFLASH fused KV materialization enabled. " "DFLASH fused KV materialization enabled. "
"n_layers=%d, num_kv_heads=%d, head_dim=%d", "n_layers=%d, num_kv_heads=%d, head_dim=%d",
@@ -1266,7 +1261,7 @@ class DFlashWorkerV2(BaseSpecWorker):
embedding_tensor.to(embed_module.weight.dtype) embedding_tensor.to(embed_module.weight.dtype)
) )
if self.ps.tp_rank == 0: if self.model_runner.tp_rank == 0:
logger.info( logger.info(
"Merged trained mask embedding into target model " "Merged trained mask embedding into target model "
"(mask_token_id=%s, source=%s)", "(mask_token_id=%s, source=%s)",
@@ -1301,7 +1296,7 @@ class DFlashWorkerV2(BaseSpecWorker):
parts = [torch.empty_like(shard_t) for _ in range(tp_size)] parts = [torch.empty_like(shard_t) for _ in range(tp_size)]
dist.all_gather(parts, shard_t, group=tp_group.device_group) dist.all_gather(parts, shard_t, group=tp_group.device_group)
self._full_embed_gpu = torch.cat(parts, dim=0)[:vocab_size] 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( logger.info(
"DFLASH cached full embed on GPU for dp attention: shape=%s", "DFLASH cached full embed on GPU for dp attention: shape=%s",
list(self._full_embed_gpu.shape), list(self._full_embed_gpu.shape),
@@ -1364,7 +1359,7 @@ class DFlashWorkerV2(BaseSpecWorker):
if resolved_id is None: if resolved_id is None:
resolved_id = tokenizer.convert_tokens_to_ids(mask_token) 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( logger.info(
"Added DFLASH mask token to tokenizer. token=%s, mask_token_id=%s, tokenizer_len=%s, model_vocab_size=%s", "Added DFLASH mask token to tokenizer. token=%s, mask_token_id=%s, tokenizer_len=%s, model_vocab_size=%s",
mask_token, mask_token,
@@ -2179,7 +2174,7 @@ class DFlashWorkerV2(BaseSpecWorker):
if self.selector is not None: if self.selector is not None:
if self._selector_sampling_enabled: if self._selector_sampling_enabled:
return 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( logger.warning(
"DFLASH non-greedy verification is unavailable on this " "DFLASH non-greedy verification is unavailable on this "
"build/device; falling back to greedy argmax verification. " "build/device; falling back to greedy argmax verification. "
@@ -2192,7 +2187,7 @@ class DFlashWorkerV2(BaseSpecWorker):
if ( if (
not is_dflash_sampling_verify_available() not is_dflash_sampling_verify_available()
and not self._warned_sampling_fallback and not self._warned_sampling_fallback
and self.ps.tp_rank == 0 and self.model_runner.tp_rank == 0
): ):
logger.warning( logger.warning(
"DFLASH non-greedy verification is unavailable on this build/device; " "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: if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig 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 from sglang.srt.model_executor.model_runner import ModelRunner
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -63,7 +62,6 @@ def build_draft_tp_worker(
*, *,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
ps: ParallelState,
nccl_port: int, nccl_port: int,
target_model_config: ModelConfig, target_model_config: ModelConfig,
algo_label: str, algo_label: str,
@@ -88,7 +86,6 @@ def build_draft_tp_worker(
draft_worker = draft_worker_cls( draft_worker = draft_worker_cls(
server_args=server_args, server_args=server_args,
gpu_id=gpu_id, gpu_id=gpu_id,
ps=ps,
nccl_port=nccl_port, nccl_port=nccl_port,
is_draft_worker=True, is_draft_worker=True,
random_seed=random_seed, random_seed=random_seed,
@@ -1,6 +1,5 @@
import logging import logging
from contextlib import nullcontext from contextlib import nullcontext
from dataclasses import replace
from typing import Callable, Optional, Protocol, runtime_checkable from typing import Callable, Optional, Protocol, runtime_checkable
import torch import torch
@@ -9,8 +8,6 @@ from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import (
is_unified_kv_triton, is_unified_kv_triton,
) )
from sglang.srt.configs.hybrid_arch import mambaish_config 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.environ import envs
from sglang.srt.layers.logprob_processor import compute_spec_logprobs from sglang.srt.layers.logprob_processor import compute_spec_logprobs
from sglang.srt.lora.layers import unwrap_lora_layer from sglang.srt.lora.layers import unwrap_lora_layer
@@ -134,7 +131,6 @@ class DSparkWorkerV2(BaseSpecWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
ps: ParallelState,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
draft_worker_cls: type[TpModelWorker] = TpModelWorker, draft_worker_cls: type[TpModelWorker] = TpModelWorker,
@@ -143,14 +139,13 @@ class DSparkWorkerV2(BaseSpecWorker):
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.ps = ps
self.nccl_port = nccl_port self.nccl_port = nccl_port
self._target_worker = target_worker self._target_worker = target_worker
self.model_runner = target_worker.model_runner self.model_runner = target_worker.model_runner
self.page_size = get_schedule().page_size self.page_size = get_schedule().page_size
self.device = target_worker.device self.device = target_worker.device
self._draft_worker = None 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: if not self._hosts_draft:
return return
@@ -166,7 +161,7 @@ class DSparkWorkerV2(BaseSpecWorker):
if ( if (
get_parallel().enable_dp_attention get_parallel().enable_dp_attention
and self._draft_is_moe and self._draft_is_moe
and ps.attn_tp_size > 1 and get_parallel().attn_tp_size > 1
): ):
raise ValueError( raise ValueError(
"DSpark + dp attention with a DeepSeek-V4 (MoE) draft requires " "DSpark + dp attention with a DeepSeek-V4 (MoE) draft requires "
@@ -178,7 +173,6 @@ class DSparkWorkerV2(BaseSpecWorker):
bundle = build_draft_tp_worker( bundle = build_draft_tp_worker(
server_args=server_args, server_args=server_args,
gpu_id=gpu_id, gpu_id=gpu_id,
ps=replace(ps, pp_rank=0, pp_size=1),
nccl_port=nccl_port, nccl_port=nccl_port,
target_model_config=target_worker.model_runner.model_config, target_model_config=target_worker.model_runner.model_config,
algo_label="DSPARK", algo_label="DSPARK",
@@ -232,7 +226,7 @@ class DSparkWorkerV2(BaseSpecWorker):
else parallel.tp_group else parallel.tp_group
) )
if self.ps.tp_rank == 0: if self.model_runner.tp_rank == 0:
logger.info( logger.info(
"Initialized DSpark draft runner. attention_backend=%s, model=%s, " "Initialized DSpark draft runner. attention_backend=%s, model=%s, "
"gamma=%s, verify_num_draft_tokens=%s, query_token_num=%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 getattr(self.draft_model, "uses_own_vocab_modules", False):
if self.ps.tp_rank == 0: if self.model_runner.tp_rank == 0:
logger.info( logger.info(
"DSpark draft uses its checkpoint-local embedding and LM head." "DSpark draft uses its checkpoint-local embedding and LM head."
) )
@@ -279,7 +273,7 @@ class DSparkWorkerV2(BaseSpecWorker):
gamma=self.gamma, gamma=self.gamma,
model_runner=self.model_runner, model_runner=self.model_runner,
device=self.device, 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, verify_num_draft_tokens=self.verify_num_draft_tokens,
tp_sync=self._tp_sync, tp_sync=self._tp_sync,
) )
@@ -327,7 +321,7 @@ class DSparkWorkerV2(BaseSpecWorker):
and self._verify_planner.mode_value == "static" and self._verify_planner.mode_value == "static"
and self._draft_is_moe and self._draft_is_moe
and not get_parallel().enable_dp_attention and not get_parallel().enable_dp_attention
and self.ps.pp_size == 1 and self.model_runner.pp_size == 1
) )
if ( if (
(self._verify_planner.is_compact_mode or static_epilogue_supported) (self._verify_planner.is_compact_mode or static_epilogue_supported)
@@ -391,7 +385,7 @@ class DSparkWorkerV2(BaseSpecWorker):
planner=self._verify_planner, planner=self._verify_planner,
gamma=self.gamma, gamma=self.gamma,
verify_num_draft_tokens=self.verify_num_draft_tokens, 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, device=self.device,
simulate_acc_len=self._simulate_acc_len, simulate_acc_len=self._simulate_acc_len,
) )
@@ -451,7 +445,7 @@ class DSparkWorkerV2(BaseSpecWorker):
draft_model=self.draft_model, draft_model=self.draft_model,
is_deepseek_v4_draft=self._draft_is_moe, 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( logger.info(
"DSpark prefill target-hidden projection runs before " "DSpark prefill target-hidden projection runs before "
"sequence-parallel gather." "sequence-parallel gather."
@@ -508,7 +502,7 @@ class DSparkWorkerV2(BaseSpecWorker):
gamma=self.gamma, gamma=self.gamma,
max_bs=max(get_exec().graph.cuda_graph_config.decode.bs), max_bs=max(get_exec().graph.cuda_graph_config.decode.bs),
device=self.device, device=self.device,
tp_rank=self.ps.tp_rank, tp_rank=self.model_runner.tp_rank,
tp_sync=self._tp_sync, tp_sync=self._tp_sync,
available_memory_gb=available_memory_gb, available_memory_gb=available_memory_gb,
confidence_fn=( 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.environ import envs
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
DpPaddingMode, DpPaddingMode,
deployment_attn_dp_size,
set_dp_buffer_len, set_dp_buffer_len,
set_is_extend_in_batch, set_is_extend_in_batch,
) )
@@ -114,8 +115,8 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
# Fields the parent's capture() reads: # Fields the parent's capture() reads:
self.device = model_runner.device self.device = model_runner.device
self.device_module = torch.get_device_module(self.device) self.device_module = torch.get_device_module(self.device)
self.tp_size = model_runner.ps.tp_size self.tp_size = model_runner.tp_size
self.attn_dp_size = model_runner.ps.attn_dp_size self.attn_dp_size = deployment_attn_dp_size()
self.pp_size = get_parallel().pp_size self.pp_size = get_parallel().pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = get_exec().graph.disable_cuda_graph_padding 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.compilation.torch_compile_decoration import set_torch_compile_config
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
DpPaddingMode, DpPaddingMode,
deployment_attn_dp_size,
set_dp_buffer_len, set_dp_buffer_len,
set_is_extend_in_batch, set_is_extend_in_batch,
) )
@@ -111,8 +112,8 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
# Fields the parent's capture() reads: # Fields the parent's capture() reads:
self.device = model_runner.device self.device = model_runner.device
self.device_module = torch.get_device_module(self.device) self.device_module = torch.get_device_module(self.device)
self.tp_size = model_runner.ps.tp_size self.tp_size = model_runner.tp_size
self.attn_dp_size = model_runner.ps.attn_dp_size self.attn_dp_size = deployment_attn_dp_size()
self.pp_size = get_parallel().pp_size self.pp_size = get_parallel().pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = get_exec().graph.disable_cuda_graph_padding self.disable_padding = get_exec().graph.disable_cuda_graph_padding
@@ -1,14 +1,12 @@
import contextlib import contextlib
import logging import logging
import time import time
from dataclasses import replace
from typing import List, Optional from typing import List, Optional
import torch import torch
from sglang.kernels.ops.speculative.topk1 import draft_topk1_postprocess 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.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.environ import envs
from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_extend_npu_graph_runner import ( from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_extend_npu_graph_runner import (
EAGLEDraftExtendNpuGraphRunner, EAGLEDraftExtendNpuGraphRunner,
@@ -238,7 +236,6 @@ class EagleDraftWorker(EagleDraftWorkerBase):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
ps: ParallelState,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -247,7 +244,6 @@ class EagleDraftWorker(EagleDraftWorkerBase):
# copy args # copy args
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.ps = ps
self.nccl_port = nccl_port self.nccl_port = nccl_port
self.target_worker = target_worker self.target_worker = target_worker
@@ -288,8 +284,6 @@ class EagleDraftWorker(EagleDraftWorkerBase):
self.draft_worker = TpModelWorker( self.draft_worker = TpModelWorker(
server_args=server_args, server_args=server_args,
gpu_id=gpu_id, gpu_id=gpu_id,
# spec workers don't support pipeline parallelism
ps=replace(ps, pp_rank=0, pp_size=1),
nccl_port=nccl_port, nccl_port=nccl_port,
is_draft_worker=True, is_draft_worker=True,
# The draft runs at absolute target positions. # The draft runs at absolute target positions.
@@ -1293,7 +1287,6 @@ class EAGLEWorkerV2(BaseSpecWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
ps: ParallelState,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -1304,7 +1297,6 @@ class EAGLEWorkerV2(BaseSpecWorker):
self.topk = get_spec().speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = get_spec().speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.ps = ps
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.device = get_device().device self.device = get_device().device
self._target_worker = target_worker self._target_worker = target_worker
@@ -1320,7 +1312,6 @@ class EAGLEWorkerV2(BaseSpecWorker):
EagleDraftWorker( EagleDraftWorker(
server_args, server_args,
gpu_id, gpu_id,
ps,
nccl_port, nccl_port,
target_worker, target_worker,
) )
@@ -1873,7 +1864,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput): def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
monkey_patch_torch_reductions() monkey_patch_torch_reductions()
named_tensors = MultiprocessingSerializer.deserialize( 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 = ( success, message = (
self.draft_worker.draft_runner.weight_updater.update_weights_from_tensor( 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.compilation.torch_compile_decoration import set_torch_compile_config
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
DpPaddingMode, DpPaddingMode,
deployment_attn_dp_size,
set_dp_buffer_len, set_dp_buffer_len,
set_is_extend_in_batch, set_is_extend_in_batch,
) )
@@ -98,8 +99,8 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
self.require_mlp_tp_gather = require_mlp_tp_gather() self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync() self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather() self.require_attn_tp_gather = require_attn_tp_gather()
self.tp_size = self.model_runner.ps.tp_size self.tp_size = self.model_runner.tp_size
self.attn_dp_size = self.model_runner.ps.attn_dp_size self.attn_dp_size = deployment_attn_dp_size()
self.pp_size = get_parallel().pp_size self.pp_size = get_parallel().pp_size
self.speculative_num_steps = get_spec().speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.topk = get_spec().speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
@@ -23,12 +23,10 @@ from __future__ import annotations
import logging import logging
import time import time
from dataclasses import replace
from typing import Optional from typing import Optional
import torch import torch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.moe.utils import ( from sglang.srt.layers.moe.utils import (
draft_model_build_scope, draft_model_build_scope,
speculative_moe_a2a_backend_context, speculative_moe_a2a_backend_context,
@@ -102,7 +100,6 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
ps: ParallelState,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -112,7 +109,6 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
self.topk = get_spec().speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = get_spec().speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.ps = ps
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.device = get_device().device self.device = get_device().device
self.target_worker = target_worker self.target_worker = target_worker
@@ -146,8 +142,6 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
self, self,
server_args=server_args, server_args=server_args,
gpu_id=gpu_id, gpu_id=gpu_id,
# spec workers don't support pipeline parallelism
ps=replace(ps, pp_rank=0, pp_size=1),
nccl_port=nccl_port, nccl_port=nccl_port,
is_draft_worker=True, is_draft_worker=True,
# The draft runs at absolute target positions. # The draft runs at absolute target positions.
@@ -702,7 +696,6 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
ps: ParallelState,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -715,7 +708,6 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
self.topk = get_spec().speculative_eagle_topk self.topk = get_spec().speculative_eagle_topk
self.speculative_num_steps = get_spec().speculative_num_steps self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
self.ps = ps
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.device = get_device().device self.device = get_device().device
self._target_worker = target_worker self._target_worker = target_worker
@@ -730,7 +722,6 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
self._draft_worker = FrozenKVMTPDraftWorker( self._draft_worker = FrozenKVMTPDraftWorker(
server_args, server_args,
gpu_id, gpu_id,
ps,
nccl_port, nccl_port,
target_worker, target_worker,
) )
@@ -155,7 +155,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
# Fields the parent's capture() reads: # Fields the parent's capture() reads:
self.device = model_runner.device self.device = model_runner.device
self.device_module = torch.get_device_module(self.device) self.device_module = torch.get_device_module(self.device)
self.tp_size = model_runner.ps.tp_size self.tp_size = model_runner.tp_size
self.dp_size = get_parallel().dp_size self.dp_size = get_parallel().dp_size
self.pp_size = get_parallel().pp_size self.pp_size = get_parallel().pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile self.enable_torch_compile = get_flags().capture.enable_torch_compile
@@ -21,7 +21,6 @@ from typing import TYPE_CHECKING, List
import torch import torch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.hardware_backend.npu.graph_runner.multi_layer_eagle_draft_extend_npu_graph_runner import ( from sglang.srt.hardware_backend.npu.graph_runner.multi_layer_eagle_draft_extend_npu_graph_runner import (
MultiLayerEagleMultiStepDraftExtendNpuGraphRunner, MultiLayerEagleMultiStepDraftExtendNpuGraphRunner,
@@ -121,7 +120,6 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
ps: ParallelState,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -130,7 +128,6 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
# copy args # copy args
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.ps = ps
self.nccl_port = nccl_port self.nccl_port = nccl_port
self.target_worker = target_worker self.target_worker = target_worker
self.draft_extend_attn_backend_list = [] self.draft_extend_attn_backend_list = []
@@ -171,8 +168,6 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
self.draft_worker = TpModelWorker( self.draft_worker = TpModelWorker(
server_args=server_args, server_args=server_args,
gpu_id=gpu_id, gpu_id=gpu_id,
# spec workers don't support pipeline parallelism
ps=replace(ps, pp_rank=0, pp_size=1),
nccl_port=nccl_port, nccl_port=nccl_port,
is_draft_worker=True, is_draft_worker=True,
is_multi_layer_eagle=True, is_multi_layer_eagle=True,
@@ -1013,7 +1008,6 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
ps: ParallelState,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -1035,7 +1029,6 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
self._draft_worker = MultiLayerEagleDraftWorker( self._draft_worker = MultiLayerEagleDraftWorker(
server_args, server_args,
gpu_id, gpu_id,
ps,
nccl_port, nccl_port,
target_worker, target_worker,
) )
@@ -7,7 +7,6 @@ import torch
from sglang.kernels.ops.speculative.cache_locs import ( from sglang.kernels.ops.speculative.cache_locs import (
assign_extend_cache_locs_func as assign_extend_cache_locs_func, 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.layers.logprob_processor import compute_spec_logprobs
from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.scheduler import GenerationBatchResult
@@ -89,7 +88,6 @@ class NGRAMWorker(BaseSpecWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
ps: ParallelState,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -99,7 +97,7 @@ class NGRAMWorker(BaseSpecWorker):
self.enable_overlap = not get_schedule().disable_overlap_schedule self.enable_overlap = not get_schedule().disable_overlap_schedule
self._target_worker = target_worker self._target_worker = target_worker
self.model_runner = target_worker.model_runner 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.page_size = get_schedule().page_size
self.draft_token_num: int = get_spec().speculative_num_draft_tokens self.draft_token_num: int = get_spec().speculative_num_draft_tokens
self.max_trie_depth: int = get_spec().speculative_ngram_max_trie_depth self.max_trie_depth: int = get_spec().speculative_ngram_max_trie_depth
@@ -1,10 +1,8 @@
import logging import logging
from dataclasses import replace
from typing import Optional from typing import Optional
import torch import torch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.moe.utils import ( from sglang.srt.layers.moe.utils import (
draft_model_build_scope, draft_model_build_scope,
speculative_moe_backend_context, speculative_moe_backend_context,
@@ -45,7 +43,6 @@ class StandaloneDraftWorker(EagleDraftWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
ps: ParallelState,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -54,7 +51,6 @@ class StandaloneDraftWorker(EagleDraftWorker):
# copy args # copy args
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.ps = ps
self.nccl_port = nccl_port self.nccl_port = nccl_port
self.target_worker = target_worker self.target_worker = target_worker
@@ -84,8 +80,6 @@ class StandaloneDraftWorker(EagleDraftWorker):
self.draft_worker = TpModelWorker( self.draft_worker = TpModelWorker(
server_args=server_args, server_args=server_args,
gpu_id=gpu_id, gpu_id=gpu_id,
# spec workers don't support pipeline parallelism
ps=replace(ps, pp_rank=0, pp_size=1),
nccl_port=nccl_port, nccl_port=nccl_port,
is_draft_worker=True, is_draft_worker=True,
# The draft runs at absolute target positions. # The draft runs at absolute target positions.
@@ -167,7 +161,6 @@ class StandaloneWorkerV2(EAGLEWorkerV2):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
ps: ParallelState,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -190,7 +183,6 @@ class StandaloneWorkerV2(EAGLEWorkerV2):
self._draft_worker = StandaloneDraftWorker( self._draft_worker = StandaloneDraftWorker(
server_args, server_args,
gpu_id, gpu_id,
ps,
nccl_port, nccl_port,
target_worker, target_worker,
) )
+2 -2
View File
@@ -45,8 +45,8 @@ def init_uno_lora_manager(
dtype=model_runner.dtype, dtype=model_runner.dtype,
server_args=model_runner.server_args, server_args=model_runner.server_args,
lora_backend="uno_cublas", # fast path lora_backend="uno_cublas", # fast path
tp_size=model_runner.ps.tp_size, tp_size=model_runner.tp_size,
tp_rank=model_runner.ps.tp_rank, tp_rank=model_runner.tp_rank,
# Infer these from the one trained adapter. # Infer these from the one trained adapter.
max_lora_rank=None, max_lora_rank=None,
target_modules=None, target_modules=None,
@@ -46,7 +46,6 @@ from sglang.srt.utils.common import (
) )
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
@@ -62,7 +61,6 @@ class UnoWorkerV2(BaseSpecWorker):
self, self,
server_args: ServerArgs, server_args: ServerArgs,
gpu_id: int, gpu_id: int,
ps: ParallelState,
nccl_port: int, nccl_port: int,
target_worker: TpModelWorker, target_worker: TpModelWorker,
): ):
@@ -70,7 +68,6 @@ class UnoWorkerV2(BaseSpecWorker):
self.server_args = server_args self.server_args = server_args
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.ps = ps
self.nccl_port = nccl_port self.nccl_port = nccl_port
self._target_worker = target_worker self._target_worker = target_worker
+2 -7
View File
@@ -8,7 +8,6 @@ from typing import Callable, Dict, List, Optional
import torch import torch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import ProfileReqOutput from sglang.srt.managers.io_struct import ProfileReqOutput
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode 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: class ProfileManager:
def __init__(self, ps: ParallelState, cpu_group): def __init__(self, cpu_group):
self.stage_based_trigger = _StageBasedTrigger( self.stage_based_trigger = _StageBasedTrigger(
on_start=self._do_start, on_start=self._do_start,
on_stop=self._do_stop, on_stop=self._do_stop,
) )
self.ps = ps
self.cpu_group = cpu_group 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_kwargs = None
self.profiler = None self.profiler = None
self.detailed_annotations = False self.detailed_annotations = False
@@ -154,7 +152,6 @@ class ProfileManager:
set_detailed_annotations_enabled(self.detailed_annotations) set_detailed_annotations_enabled(self.detailed_annotations)
self.profiler = _ProfilerBase.create( self.profiler = _ProfilerBase.create(
**self.profiler_kwargs, **self.profiler_kwargs,
ps=self.ps,
cpu_group=self.cpu_group, cpu_group=self.cpu_group,
first_rank_in_node=self.first_rank_in_node, first_rank_in_node=self.first_rank_in_node,
output_suffix=f"-{stage}" if stage else "", output_suffix=f"-{stage}" if stage else "",
@@ -303,7 +300,6 @@ class _ProfilerConcreteBase(_ProfilerBase):
output_prefix: str, output_prefix: str,
output_suffix: str, output_suffix: str,
profile_id: str, profile_id: str,
ps: ParallelState,
cpu_group, cpu_group,
first_rank_in_node: bool, first_rank_in_node: bool,
): ):
@@ -311,7 +307,6 @@ class _ProfilerConcreteBase(_ProfilerBase):
self.output_prefix = output_prefix self.output_prefix = output_prefix
self.output_suffix = output_suffix self.output_suffix = output_suffix
self.profile_id = profile_id self.profile_id = profile_id
self.ps = ps
self.cpu_group = cpu_group self.cpu_group = cpu_group
self.first_rank_in_node = first_rank_in_node 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 from sglang.srt.model_loader.loader import get_model_loader
server_args = self.server_args 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( publish(
server_args, server_args,
role="weight_cache_daemon", role="weight_cache_daemon",
ranks=SpawnRanks( ranks=SpawnRanks(
world_rank=spawn_world_rank( world_rank=spawn_world_rank(
server_args, tp_rank=self.tp_rank, pp_rank=self.pp_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 # Build engine's config fingerprint
from sglang.srt.layers.dp_attention import get_moe_cp_size from sglang.srt.layers.dp_attention import get_moe_cp_size
ps = get_parallel() parallel = get_parallel()
tp_size = ps.tp_size tp_size = parallel.tp_size
tp_rank = ps.tp_rank tp_rank = parallel.tp_rank
pp_size = ps.pp_size pp_size = parallel.pp_size
pp_rank = ps.pp_rank 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_size = get_moe_cp_size()
moe_dp_rank = ps.moe_dp_rank moe_dp_rank = parallel.moe_dp_rank
moe_ep_rank = ps.moe_ep_rank moe_ep_rank = parallel.moe_ep_rank
dp_size = get_parallel().dp_size dp_size = get_parallel().dp_size
@@ -535,10 +535,10 @@ class IpcModelLoader(BaseModelLoader):
moe_dp_size=moe_dp_size, moe_dp_size=moe_dp_size,
moe_dp_rank=moe_dp_rank, moe_dp_rank=moe_dp_rank,
moe_ep_rank=moe_ep_rank, moe_ep_rank=moe_ep_rank,
enable_dp_attention=ps.enable_dp_attention, enable_dp_attention=parallel.enable_dp_attention,
enable_dp_lm_head=ps.enable_dp_lm_head, enable_dp_lm_head=parallel.enable_dp_lm_head,
attn_cp_size=ps.attn_cp_size, attn_cp_size=parallel.attn_cp_size,
moe_dense_tp_size=ps.moe_dense_tp_size, moe_dense_tp_size=parallel.moe_dense_tp_size,
moe_a2a_backend=get_exec().moe.moe_a2a_backend, moe_a2a_backend=get_exec().moe.moe_a2a_backend,
quant_method=quant_method, quant_method=quant_method,
quant_config_hash=hash_quant_config(quant_config), quant_config_hash=hash_quant_config(quant_config),
@@ -7,7 +7,6 @@ import torch.nn.functional as F
from torch import nn from torch import nn
from sglang.srt.configs.model_config import AttentionArch 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.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, ReqToTokenPool
@@ -339,7 +338,6 @@ class MockModelRunner(ModelRunner):
self.tp_size = 1 self.tp_size = 1
self.dp_size = 1 self.dp_size = 1
self.pp_size = 1 self.pp_size = 1
self.ps = ParallelState.trivial()
self.is_draft_worker = False self.is_draft_worker = False
self.max_running_requests = pool_batch_size self.max_running_requests = pool_batch_size
# trtllm_mha __init__ scans model.modules() for ENCODER_ONLY layers; # trtllm_mha __init__ scans model.modules() for ENCODER_ONLY layers;
@@ -5,7 +5,6 @@ from typing import Any
import torch import torch
from torch import nn 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.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool, ReqToTokenPool
@@ -322,7 +321,6 @@ class DSAMockModelRunner(ModelRunner):
self._kernel_warmed_up = True self._kernel_warmed_up = True
self.dp_size = 1 self.dp_size = 1
self.pp_size = 1 self.pp_size = 1
self.ps = ParallelState.trivial()
self._server_args_override = get_context().override_server_args( self._server_args_override = get_context().override_server_args(
attention_backend=case.backend, attention_backend=case.backend,
chunked_prefill_size=-1, chunked_prefill_size=-1,
@@ -22,7 +22,6 @@ from torch import nn
from sglang.kernels.ops.attention.dsv4.quant_k_cache import ( from sglang.kernels.ops.attention.dsv4.quant_k_cache import (
quant_to_nope_fp8_rope_bf16_pack_triton, 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.environ import envs
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
@@ -349,7 +348,6 @@ class MockDSV4ModelRunner:
self.tp_size = 1 self.tp_size = 1
self.dp_size = 1 self.dp_size = 1
self.pp_size = 1 self.pp_size = 1
self.ps = ParallelState.trivial()
self._server_args_override = get_context().override_server_args( self._server_args_override = get_context().override_server_args(
attention_backend=case.backend, attention_backend=case.backend,
chunked_prefill_size=-1, chunked_prefill_size=-1,
@@ -10,7 +10,6 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateShape, Mamba2StateShape,
) )
from sglang.srt.configs.model_config import AttentionArch 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.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
HybridLinearAttnBackend, HybridLinearAttnBackend,
@@ -228,7 +227,6 @@ class MockGDNModelRunner(ModelRunner):
self.decode_attention_backend_str = case.backend self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None self.draft_attention_backend = None
self.gpu_id = 0 self.gpu_id = 0
self.ps = ParallelState.trivial()
self.spec_algorithm = SpeculativeAlgorithm.NONE self.spec_algorithm = SpeculativeAlgorithm.NONE
self.canary_manager = None self.canary_manager = None
self.page_size = case.page_size self.page_size = case.page_size
@@ -10,7 +10,6 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateDType, Mamba2StateDType,
) )
from sglang.srt.configs.model_config import AttentionArch 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.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( from sglang.srt.layers.attention.hybrid_linear_attn_backend import (
HybridLinearAttnBackend, HybridLinearAttnBackend,
@@ -231,7 +230,6 @@ class MockKDAModelRunner(ModelRunner):
self.decode_attention_backend_str = case.backend self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None self.draft_attention_backend = None
self.gpu_id = 0 self.gpu_id = 0
self.ps = ParallelState.trivial()
self.spec_algorithm = SpeculativeAlgorithm.NONE self.spec_algorithm = SpeculativeAlgorithm.NONE
self.canary_manager = None self.canary_manager = None
self.page_size = case.page_size self.page_size = case.page_size
@@ -10,7 +10,6 @@ from sglang.srt.configs.mamba_utils import (
Mamba2StateShape, Mamba2StateShape,
) )
from sglang.srt.configs.model_config import AttentionArch 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.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.linear.lightning_backend import ( from sglang.srt.layers.attention.linear.lightning_backend import (
LightningAttentionBackend, LightningAttentionBackend,
@@ -239,7 +238,6 @@ class MockLightningModelRunner(ModelRunner):
self.decode_attention_backend_str = case.backend self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None self.draft_attention_backend = None
self.gpu_id = 0 self.gpu_id = 0
self.ps = ParallelState.trivial()
self.spec_algorithm = SpeculativeAlgorithm.NONE self.spec_algorithm = SpeculativeAlgorithm.NONE
self.canary_manager = None self.canary_manager = None
self.page_size = case.page_size self.page_size = case.page_size
@@ -26,7 +26,6 @@ from sglang.srt.configs.mamba_utils import ( # noqa: E402
Mamba2StateShape, Mamba2StateShape,
) )
from sglang.srt.configs.model_config import AttentionArch # noqa: E402 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 from sglang.srt.layers.attention.attention_registry import ( # noqa: E402
ATTENTION_BACKENDS, ATTENTION_BACKENDS,
) )
@@ -325,7 +324,6 @@ class MockMamba2ModelRunner(ModelRunner):
self.decode_attention_backend_str = case.backend self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None self.draft_attention_backend = None
self.gpu_id = 0 self.gpu_id = 0
self.ps = ParallelState.trivial()
self.spec_algorithm = SpeculativeAlgorithm.NONE self.spec_algorithm = SpeculativeAlgorithm.NONE
self.canary_manager = None self.canary_manager = None
self.page_size = case.page_size self.page_size = case.page_size
@@ -7,7 +7,6 @@ import torch.nn.functional as F
from torch import nn from torch import nn
from sglang.srt.configs.model_config import AttentionArch 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.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool, ReqToTokenPool
@@ -247,7 +246,6 @@ class MockMLAModelRunner(ModelRunner):
self.tp_size = 1 self.tp_size = 1
self.dp_size = 1 self.dp_size = 1
self.pp_size = 1 self.pp_size = 1
self.ps = ParallelState.trivial()
self.spec_algorithm = SpeculativeAlgorithm.NONE self.spec_algorithm = SpeculativeAlgorithm.NONE
speculative_num_draft_tokens = ( speculative_num_draft_tokens = (
max(case.input_lens) max(case.input_lens)
@@ -2,6 +2,8 @@ from __future__ import annotations
from typing import TYPE_CHECKING, Dict, Iterator, List, Optional from typing import TYPE_CHECKING, Dict, Iterator, List, Optional
from sglang.srt.runtime_context import get_parallel
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
from sglang.test.scripted_runtime.context.api import ScriptedContext 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: if s.chunked_req is not None:
yield s.chunked_req yield s.chunked_req
yield from s.waiting_queue 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): for mb in (*s.mbs, *s.last_mbs, *s.running_mbs):
if mb is not None: if mb is not None:
yield from mb.reqs yield from mb.reqs
@@ -13,6 +13,7 @@ import zmq
from sglang.srt.arg_groups.overrides import resolving_view from sglang.srt.arg_groups.overrides import resolving_view
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle 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.srt.utils.network import get_zmq_socket
from sglang.test.scripted_runtime.background_http_poster import BackgroundHttpPoster from sglang.test.scripted_runtime.background_http_poster import BackgroundHttpPoster
from sglang.test.scripted_runtime.context import ScriptedContext from sglang.test.scripted_runtime.context import ScriptedContext
@@ -125,9 +126,9 @@ class ScriptedSchedulerHook:
) -> None: ) -> None:
self.scheduler = scheduler self.scheduler = scheduler
self._is_driver = ( self._is_driver = (
scheduler.ps.pp_rank == 0 get_parallel().pp_rank == 0
and scheduler.ps.tp_rank == 0 and get_parallel().tp_rank == 0
and scheduler.ps.attn_cp_rank == 0 and get_parallel().attn_cp_rank == 0
) )
self._batch_log: List[ScriptedBatchRecord] = [] self._batch_log: List[ScriptedBatchRecord] = []
@@ -15,7 +15,6 @@ import torch
from sglang.benchmark.one_batch import TreeCacheNamespace from sglang.benchmark.one_batch import TreeCacheNamespace
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.model_executor.forward_context import ( from sglang.srt.model_executor.forward_context import (
@@ -65,7 +64,6 @@ class TestForwardSplitPrefill(CustomTestCase):
model_config=cls.model_config, model_config=cls.model_config,
mem_fraction_static=cls.server_args.mem_fraction_static, mem_fraction_static=cls.server_args.mem_fraction_static,
gpu_id=0, gpu_id=0,
ps=ParallelState.trivial(tp_size=cls.tp_size),
nccl_port=cls.port_args.nccl_port, nccl_port=cls.port_args.nccl_port,
server_args=cls.server_args, server_args=cls.server_args,
) )
-2
View File
@@ -9,7 +9,6 @@ import torch.nn.functional as F
from transformers import AutoModel, AutoProcessor, AutoTokenizer from transformers import AutoModel, AutoProcessor, AutoTokenizer
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.entrypoints.openai.protocol import ChatCompletionRequest from sglang.srt.entrypoints.openai.protocol import ChatCompletionRequest
from sglang.srt.managers.mm_utils import embed_mm_inputs, init_mm_embedding_cache from sglang.srt.managers.mm_utils import embed_mm_inputs, init_mm_embedding_cache
from sglang.srt.managers.schedule_batch import ( from sglang.srt.managers.schedule_batch import (
@@ -151,7 +150,6 @@ class VisionLLMLogitsBase(unittest.IsolatedAsyncioTestCase):
model_config=ModelConfig(self.model_path, model_override_args="{}"), model_config=ModelConfig(self.model_path, model_override_args="{}"),
mem_fraction_static=0.8, mem_fraction_static=0.8,
gpu_id=0, gpu_id=0,
ps=ParallelState.trivial(),
nccl_port=12435, nccl_port=12435,
server_args=server_args, server_args=server_args,
) )
@@ -11,7 +11,6 @@ from sglang.srt.disaggregation.decode import (
) )
from sglang.srt.disaggregation.fake.conn import FakeKVManager, FakeKVReceiver from sglang.srt.disaggregation.fake.conn import FakeKVManager, FakeKVReceiver
from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.schedule_batch import FINISH_ABORT from sglang.srt.managers.schedule_batch import FINISH_ABORT
from sglang.srt.managers.scheduler import Scheduler from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.runtime_context import get_context, publish, reset_context from sglang.srt.runtime_context import get_context, publish, reset_context
@@ -440,7 +439,6 @@ class TestDecodeQueueCleanup(CustomTestCase):
scheduler.last_batch = None scheduler.last_batch = None
scheduler.cur_batch_for_debug = None scheduler.cur_batch_for_debug = None
scheduler.enable_overlap = False scheduler.enable_overlap = False
scheduler.ps = ParallelState.trivial()
scheduler.running_mbs = [] scheduler.running_mbs = []
scheduler.waiting_queue = [] scheduler.waiting_queue = []
scheduler.grammar_manager = SimpleNamespace(grammar_queue=[]) scheduler.grammar_manager = SimpleNamespace(grammar_queue=[])
@@ -131,7 +131,6 @@ class TestHandlePdRoleSwitch(unittest.TestCase):
def test_rejected_when_decode_graph_headroom_is_insufficient(self): def test_rejected_when_decode_graph_headroom_is_insufficient(self):
s = self._scheduler(DisaggregationMode.PREFILL) s = self._scheduler(DisaggregationMode.PREFILL)
s.device = "cuda" s.device = "cuda"
s.ps = SimpleNamespace(gpu_id=0)
s.tp_worker.get_decode_cuda_graph_bs.return_value = [] s.tp_worker.get_decode_cuda_graph_bs.return_value = []
with patch.object(role_switch, "get_available_gpu_memory", return_value=0.5): with patch.object(role_switch, "get_available_gpu_memory", return_value=0.5):
out = Scheduler.handle_pd_role_switch( out = Scheduler.handle_pd_role_switch(
@@ -151,7 +150,6 @@ class TestHandlePdRoleSwitch(unittest.TestCase):
def test_decode_graph_headroom_allows_flip(self): def test_decode_graph_headroom_allows_flip(self):
s = self._scheduler(DisaggregationMode.PREFILL) s = self._scheduler(DisaggregationMode.PREFILL)
s.device = "cuda" s.device = "cuda"
s.ps = SimpleNamespace(gpu_id=0)
s.tp_worker.get_decode_cuda_graph_bs.return_value = [] s.tp_worker.get_decode_cuda_graph_bs.return_value = []
with patch.object(role_switch, "get_available_gpu_memory", return_value=1.0): with patch.object(role_switch, "get_available_gpu_memory", return_value=1.0):
out = Scheduler.handle_pd_role_switch( out = Scheduler.handle_pd_role_switch(
@@ -43,17 +43,17 @@ def test_dp_leaders_reuse_node_local_ports(
for dp_rank, tp_rank in enumerate(ranks): for dp_rank, tp_rank in enumerate(ranks):
scheduler = SimpleNamespace( scheduler = SimpleNamespace(
server_args=SimpleNamespace(), server_args=SimpleNamespace(),
ps=SimpleNamespace(
tp_rank=tp_rank,
tp_size=parallel.tp_size,
pp_size=parallel.pp_size,
attn_tp_size=parallel.attn_tp_size,
attn_cp_size=parallel.attn_cp_size,
attn_dp_rank=dp_rank,
dp_size=dp_size,
),
model_config=SimpleNamespace(is_multimodal=False), model_config=SimpleNamespace(is_multimodal=False),
) )
# Where this rank sits, stated whole: the attention rank follows
# from the TP rank and the attention-TP width, and the identities
# refuse the combination if it describes no real layout.
with parallel.override(
tp_rank=tp_rank,
attn_dp_rank=dp_rank,
attn_tp_rank=tp_rank % parallel.attn_tp_size,
attn_cp_rank=0,
):
ports.append(rust_server.RustServer.launch(scheduler).http_port) ports.append(rust_server.RustServer.launch(scheduler).http_port)
calls = extension.return_value.Server.call_args_list calls = extension.return_value.Server.call_args_list
@@ -23,7 +23,6 @@ _HAS_MLX = importlib.util.find_spec("mlx") is not None
_SKIP_REASON = "requires mlx" _SKIP_REASON = "requires mlx"
if _HAS_MLX: if _HAS_MLX:
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.hardware_backend.mlx.model_runner_stub import ( from sglang.srt.hardware_backend.mlx.model_runner_stub import (
MLX_AUX_STATE_SIZE_MAX_RUNNING_REQUESTS_RATIO as RATIO, MLX_AUX_STATE_SIZE_MAX_RUNNING_REQUESTS_RATIO as RATIO,
) )
@@ -62,7 +61,7 @@ def _stub_for_initialize(
stub = MlxModelRunnerStub.__new__(MlxModelRunnerStub) stub = MlxModelRunnerStub.__new__(MlxModelRunnerStub)
stub._mlx_pool_size = pool_size stub._mlx_pool_size = pool_size
stub.device = "cpu" stub.device = "cpu"
stub.ps = ParallelState.trivial(dp_size=dp_size, attn_dp_size=attn_dp_size) stub.attn_dp_size = attn_dp_size
stub.server_args = server_args stub.server_args = server_args
stub.model_config = SimpleNamespace( stub.model_config = SimpleNamespace(
is_hybrid_swa=False, is_hybrid_swa=False,
@@ -208,14 +208,7 @@ class TestSchedulerProfilerManagerMPS(unittest.TestCase):
SchedulerProfilerManager, SchedulerProfilerManager,
) )
class FakePS: mgr = SchedulerProfilerManager(dp_tp_cpu_group=None, get_forward_ct=lambda: 0)
tp_rank = dp_rank = pp_rank = moe_ep_rank = 0
dp_size = pp_size = moe_ep_size = 1
gpu_id = 0
mgr = SchedulerProfilerManager(
ps=FakePS(), dp_tp_cpu_group=None, get_forward_ct=lambda: 0
)
mgr._init_profile(output_dir, None, None, None, None, None, False, "test") mgr._init_profile(output_dir, None, None, None, None, None, False, "test")
return mgr return mgr
@@ -228,7 +228,7 @@ class TestMambaPrefillTrackMetadata(unittest.TestCase):
prefill_attention_backend_str="torch_native", prefill_attention_backend_str="torch_native",
ngram_embedding_manager=SimpleNamespace(enabled=False), ngram_embedding_manager=SimpleNamespace(enabled=False),
lora_manager=None, lora_manager=None,
ps=SimpleNamespace(attn_dcp_size=1), attn_dcp_size=1,
attn_backend=SimpleNamespace( attn_backend=SimpleNamespace(
get_cpu_graph_seq_len_fill_value=lambda: 1, get_cpu_graph_seq_len_fill_value=lambda: 1,
get_cuda_graph_seq_len_fill_value=lambda: 1, get_cuda_graph_seq_len_fill_value=lambda: 1,
@@ -20,7 +20,11 @@ from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import (
CustomTestCase,
enter_scope,
published_topology,
)
register_cpu_ci(est_time=1, suite="base-a-test-cpu") register_cpu_ci(est_time=1, suite="base-a-test-cpu")
@@ -58,6 +62,12 @@ def load_mlx_scheduler_module():
class TestSchedulerIdleStepCounters(CustomTestCase): class TestSchedulerIdleStepCounters(CustomTestCase):
def setUp(self):
super().setUp()
# The loop asks the context where this process sits; nothing here
# builds a process group, so the placement arrives by publishing one.
enter_scope(self, published_topology(role="scheduler"))
@parameterized.expand( @parameterized.expand(
[ [
( (
@@ -406,7 +416,6 @@ class TestSchedulerIdleStepCounters(CustomTestCase):
scheduler.forward_ct = 0 scheduler.forward_ct = 0
scheduler.processed_tokens_counter = 0 scheduler.processed_tokens_counter = 0
scheduler.spec_algorithm = SpeculativeAlgorithm.NONE scheduler.spec_algorithm = SpeculativeAlgorithm.NONE
scheduler.ps = SimpleNamespace(pp_rank=0, attn_tp_rank=0, attn_cp_rank=0)
scheduler._poll_timeout_aborts = Mock(return_value=[]) scheduler._poll_timeout_aborts = Mock(return_value=[])
scheduler.scheduler_stage_metrics = None scheduler.scheduler_stage_metrics = None
scheduler.metrics_reporter = SimpleNamespace(record_scheduler_active=Mock()) scheduler.metrics_reporter = SimpleNamespace(record_scheduler_active=Mock())
@@ -455,7 +464,7 @@ class TestSchedulerIdleStepCounters(CustomTestCase):
return scheduler return scheduler
def prepare_pp_scheduler(self, scheduler): def prepare_pp_scheduler(self, scheduler):
scheduler.ps.pp_size = 2 enter_scope(self, get_parallel().override(pp_size=2, pp_rank=0))
scheduler.pp_group = SimpleNamespace(is_last_rank=True) scheduler.pp_group = SimpleNamespace(is_last_rank=True)
scheduler.forward_stream_ctx = nullcontext() scheduler.forward_stream_ctx = nullcontext()
scheduler.forward_stream = Mock() scheduler.forward_stream = Mock()
@@ -148,7 +148,6 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
streamer = Streamer( streamer = Streamer(
send_to_detokenizer=SimpleNamespace(send_output=outputs.append), send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
tree_cache=None, tree_cache=None,
ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0),
server_args=SimpleNamespace( server_args=SimpleNamespace(
stream_interval=1, stream_interval=1,
enable_request_time_stats_logging=False, enable_request_time_stats_logging=False,
@@ -181,7 +180,6 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
streamer = Streamer( streamer = Streamer(
send_to_detokenizer=SimpleNamespace(send_output=outputs.append), send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
tree_cache=None, tree_cache=None,
ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0),
server_args=SimpleNamespace( server_args=SimpleNamespace(
stream_interval=1, stream_interval=1,
enable_request_time_stats_logging=False, enable_request_time_stats_logging=False,
@@ -218,7 +216,6 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
streamer = Streamer( streamer = Streamer(
send_to_detokenizer=SimpleNamespace(send_output=outputs.append), send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
tree_cache=None, tree_cache=None,
ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0),
server_args=SimpleNamespace( server_args=SimpleNamespace(
stream_interval=1, stream_interval=1,
enable_request_time_stats_logging=False, enable_request_time_stats_logging=False,
@@ -258,7 +255,6 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
streamer = Streamer( streamer = Streamer(
send_to_detokenizer=SimpleNamespace(send_output=outputs.append), send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
tree_cache=None, tree_cache=None,
ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0),
server_args=SimpleNamespace( server_args=SimpleNamespace(
stream_interval=1, stream_interval=1,
enable_request_time_stats_logging=False, enable_request_time_stats_logging=False,
@@ -286,7 +282,6 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
streamer = Streamer( streamer = Streamer(
send_to_detokenizer=SimpleNamespace(send_output=outputs.append), send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
tree_cache=None, tree_cache=None,
ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0),
server_args=SimpleNamespace(), server_args=SimpleNamespace(),
is_generation=True, is_generation=True,
spec_algorithm=SpeculativeAlgorithm.NONE, spec_algorithm=SpeculativeAlgorithm.NONE,
@@ -312,7 +307,6 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
streamer = Streamer( streamer = Streamer(
send_to_detokenizer=SimpleNamespace(send_output=outputs.append), send_to_detokenizer=SimpleNamespace(send_output=outputs.append),
tree_cache=None, tree_cache=None,
ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0),
server_args=SimpleNamespace(), server_args=SimpleNamespace(),
is_generation=True, is_generation=True,
spec_algorithm=SpeculativeAlgorithm.NONE, spec_algorithm=SpeculativeAlgorithm.NONE,
@@ -334,7 +328,6 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase):
Streamer( Streamer(
send_to_detokenizer=SimpleNamespace(), send_to_detokenizer=SimpleNamespace(),
tree_cache=None, tree_cache=None,
ps=SimpleNamespace(),
server_args=SimpleNamespace(), server_args=SimpleNamespace(),
is_generation=True, is_generation=True,
spec_algorithm=SpeculativeAlgorithm.NONE, spec_algorithm=SpeculativeAlgorithm.NONE,
@@ -11,7 +11,6 @@ from types import SimpleNamespace
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
from sglang.srt.configs.model_config import AttentionArch from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_memory, get_memory,
get_parallel, get_parallel,
@@ -36,10 +35,17 @@ def mock_cpu_env(kv_size=2, tp_size=1, swa_eviction_interval=4):
with ( with (
patch("torch._utils._element_size", return_value=kv_size), patch("torch._utils._element_size", return_value=kv_size),
# A width is a whole topology: state the TP siblings the identities # The whole attention triple, not just one leaf: a width that does not
# relate it to, not the attention share alone. # factor describes no layout.
get_parallel().override( get_parallel().override(
tp_size=tp_size, attn_tp_size=tp_size, moe_tp_size=tp_size tp_size=tp_size,
attn_tp_size=tp_size,
attn_dp_size=1,
attn_cp_size=1,
moe_ep_size=1,
moe_ep_group=None,
moe_dp_size=1,
moe_tp_size=tp_size,
), ),
envs.SGLANG_SWA_EVICTION_INTERVAL.override(swa_eviction_interval), envs.SGLANG_SWA_EVICTION_INTERVAL.override(swa_eviction_interval),
): ):
@@ -174,7 +180,8 @@ def _make_model_runner(
mr.layer_info = SimpleNamespace( mr.layer_info = SimpleNamespace(
start_layer=0, end_layer=num_layers, num_effective_layers=num_layers start_layer=0, end_layer=num_layers, num_effective_layers=num_layers
) )
mr.ps = ParallelState.trivial() mr.attn_dp_size = 1
mr.pp_size = 1
mr.pp_group = SimpleNamespace(rank_in_group=0) mr.pp_group = SimpleNamespace(rank_in_group=0)
mr.spec_aux_config = SimpleNamespace( mr.spec_aux_config = SimpleNamespace(
eagle_draft_num_layers=None, eagle_draft_num_layers=None,
@@ -1271,7 +1278,8 @@ class TestSWAPoolFloor(CustomTestCase):
kv_cache_dtype_str="fp8_e4m3", kv_cache_dtype_str="fp8_e4m3",
model_config=cfg, model_config=cfg,
layer_info=SimpleNamespace(start_layer=0, end_layer=40), layer_info=SimpleNamespace(start_layer=0, end_layer=40),
ps=SimpleNamespace(pp_size=1, attn_dp_size=1), pp_size=1,
attn_dp_size=1,
sliding_window_size=128, sliding_window_size=128,
page_size=256, page_size=256,
spec_algorithm=spec, spec_algorithm=spec,
@@ -29,14 +29,6 @@ class TestRustServerExtension(CustomTestCase):
self.server.start_mm_workers(sentinel.spec, 8) self.server.start_mm_workers(sentinel.spec, 8)
scheduler = SimpleNamespace( scheduler = SimpleNamespace(
ps=SimpleNamespace(
dp_size=2,
attn_dp_rank=1,
tp_size=2,
tp_rank=1,
attn_tp_size=1,
attn_cp_size=1,
),
model_config=SimpleNamespace(is_multimodal=True), model_config=SimpleNamespace(is_multimodal=True),
) )
with ( with (
@@ -50,7 +42,16 @@ class TestRustServerExtension(CustomTestCase):
patch.object( patch.object(
server_module, server_module,
"get_parallel", "get_parallel",
return_value=SimpleNamespace(nnodes=1, pp_size=1), return_value=SimpleNamespace(
nnodes=1,
pp_size=1,
dp_size=2,
attn_dp_rank=1,
tp_size=2,
tp_rank=1,
attn_tp_size=1,
attn_cp_size=1,
),
), ),
patch.object(ModelServer, "_partition_cores", return_value=(None, None)), patch.object(ModelServer, "_partition_cores", return_value=(None, None)),
patch.object( patch.object(
@@ -9,7 +9,6 @@ import unittest
from unittest.mock import patch from unittest.mock import patch
from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.scheduler_components.metrics_reporter import ( from sglang.srt.managers.scheduler_components.metrics_reporter import (
PrefillStats, PrefillStats,
SchedulerMetricsReporter, SchedulerMetricsReporter,
@@ -18,16 +17,6 @@ from sglang.srt.managers.scheduler_components.metrics_reporter import (
from sglang.test.test_utils import CustomTestCase, enter_scope from sglang.test.test_utils import CustomTestCase, enter_scope
def _make_ps(**overrides) -> ParallelState:
"""Build a ParallelState with reasonable defaults for tests; override fields via kwargs."""
defaults = dict(
dp_rank=None,
moe_dp_rank=None,
)
defaults.update(overrides)
return ParallelState.trivial(**defaults)
class _FakeReq: class _FakeReq:
def __init__( def __init__(
self, self,
@@ -75,11 +64,30 @@ class _DummyPublisherThread:
def _publish_server_args(test, **fields): def _publish_server_args(test, **fields):
"""Publish a config for the reporter under test and return the instance.""" """Publish a config for the reporter under test and return the instance.
The collector asks the context where this process sits, so the ranks are
stated too: without them a rank read falls through to a process group that
a unit test has not built.
"""
fields.setdefault("decode_log_interval", 40) fields.setdefault("decode_log_interval", 40)
override = get_context().override_server_args(**fields) override = get_context().override_server_args(**fields)
server_args = override.install() server_args = override.install()
test.addCleanup(override.restore) test.addCleanup(override.restore)
enter_scope(
test,
get_parallel().override(
tp_rank=0,
attn_tp_rank=0,
attn_cp_rank=0,
moe_ep_rank=0,
attn_dp_rank=0,
dp_rank=0,
moe_ep_size=1,
moe_dp_size=1,
moe_tp_size=1,
),
)
return server_args return server_args
@@ -93,8 +101,6 @@ def _make_reporter(test, scheduler) -> SchedulerMetricsReporter:
enable_mfu_metrics=False, enable_mfu_metrics=False,
enable_forward_pass_metrics=False, enable_forward_pass_metrics=False,
) )
if not hasattr(scheduler, "ps"):
scheduler.ps = ParallelState.trivial()
if not hasattr(scheduler, "kv_events_publisher"): if not hasattr(scheduler, "kv_events_publisher"):
scheduler.kv_events_publisher = types.SimpleNamespace( scheduler.kv_events_publisher = types.SimpleNamespace(
init_kv_events=lambda *a, **kw: None, init_kv_events=lambda *a, **kw: None,
@@ -291,9 +297,9 @@ class TestForwardPassMetrics(unittest.TestCase):
forward_pass_metrics_ipc_name=None, forward_pass_metrics_ipc_name=None,
kv_events_config=None, kv_events_config=None,
) )
scheduler.ps = _make_ps(attn_tp_rank=0, dp_rank=2, pp_rank=0, pp_size=1) # The reporter asks the context whether this is the last stage, and
# The reporter asks the context whether this is the last stage. # which replica it is reporting for.
enter_scope(self, get_parallel().override(pp_rank=0, pp_size=1)) enter_scope(self, get_parallel().override(pp_rank=0, pp_size=1, dp_rank=2))
scheduler.enable_kv_cache_events = False scheduler.enable_kv_cache_events = False
with patch( with patch(
@@ -330,7 +336,6 @@ class TestForwardPassMetrics(unittest.TestCase):
forward_pass_metrics_ipc_name=None, forward_pass_metrics_ipc_name=None,
kv_events_config=None, kv_events_config=None,
) )
scheduler.ps = _make_ps(attn_tp_rank=0, dp_rank=0, pp_rank=0, pp_size=2)
# The reporter asks the context whether this is the last stage. # The reporter asks the context whether this is the last stage.
enter_scope(self, get_parallel().override(pp_rank=0, pp_size=2)) enter_scope(self, get_parallel().override(pp_rank=0, pp_size=2))
scheduler.enable_kv_cache_events = False scheduler.enable_kv_cache_events = False
@@ -175,7 +175,6 @@ class TestDraftPerRunnerConfig(CustomTestCase):
scheduler.tp_worker = SimpleNamespace( scheduler.tp_worker = SimpleNamespace(
model_runner=SimpleNamespace(model_config=SimpleNamespace(context_len=4096)) model_runner=SimpleNamespace(model_config=SimpleNamespace(context_len=4096))
) )
scheduler.ps = SimpleNamespace(gpu_id=0)
scheduler.nccl_port = 0 scheduler.nccl_port = 0
scheduler.spec_algorithm = SimpleNamespace( scheduler.spec_algorithm = SimpleNamespace(
is_none=lambda: False, is_none=lambda: False,
@@ -6,6 +6,7 @@ import torch
from sglang.srt.layers.aux_hidden_states import pack_aux_hidden_states from sglang.srt.layers.aux_hidden_states import pack_aux_hidden_states
from sglang.srt.models.dspark import DSparkDraftMixin from sglang.srt.models.dspark import DSparkDraftMixin
from sglang.srt.runtime_context import get_parallel
from sglang.srt.speculative.dspark_components.dspark_kv_inject import ( from sglang.srt.speculative.dspark_components.dspark_kv_inject import (
TargetHiddenKvInjector, TargetHiddenKvInjector,
) )
@@ -48,16 +49,13 @@ class DSparkTargetHiddenProjectionTest(CustomTestCase):
), ),
) )
with ( with (
mock.patch( get_parallel().override(pp_group=SimpleNamespace(is_last_rank=False)),
"sglang.srt.speculative.dspark_components.dspark_worker_v2.get_pp_group",
return_value=SimpleNamespace(is_last_rank=False),
),
mock.patch( mock.patch(
"sglang.srt.speculative.dspark_components.dspark_worker_v2.get_schedule", "sglang.srt.speculative.dspark_components.dspark_worker_v2.get_schedule",
return_value=SimpleNamespace(page_size=1), return_value=SimpleNamespace(page_size=1),
), ),
): ):
worker = DSparkWorkerV2(None, 0, None, 0, target) worker = DSparkWorkerV2(None, 0, 0, target)
worker.alloc_memory_pool() worker.alloc_memory_pool()
worker.init_attention_backends() worker.init_attention_backends()
worker.init_cuda_graphs() worker.init_cuda_graphs()
+83 -1
View File
@@ -46,6 +46,7 @@ from sglang.srt.runtime_context import (
assert_published, assert_published,
derive_parallel_widths, derive_parallel_widths,
get_context, get_context,
get_device,
get_exec, get_exec,
get_flags, get_flags,
get_parallel, get_parallel,
@@ -366,6 +367,35 @@ class TestSpawnIdentities(_IsolatedOverrides):
self.assertEqual(parallel.pp_rank, 1) self.assertEqual(parallel.pp_rank, 1)
self.assertEqual(parallel.dp_rank, 2) self.assertEqual(parallel.dp_rank, 2)
def test_the_spawn_states_the_device_and_the_record_stays_clean(self):
"""The parent picks the device, so it arrives with the rest of the
placement. It is stamped onto the bag: the record is the startup
input and stays as the caller handed it over."""
server_args = ServerArgs(model_path="dummy")
publish(
server_args,
role="test",
ranks=SpawnRanks(world_rank=0, gpu_id=3),
)
self.assertEqual(get_device().gpu_id, 3)
# Not on the record at all. An `Arg` is the operator's input and is
# collected into `ServerArgs`; nobody types this one, so it is
# declared rather than carried, and the startup input has no field
# for the spawn to have to leave alone.
self.assertNotIn(
"gpu_id", {f.name for f in msgspec.structs.fields(type(server_args))}
)
def test_a_process_on_no_device_is_told_nothing(self):
"""Most roles run on no device at all, so the bundle leaves it out and
the bag keeps the declared default rather than inventing a zero."""
publish(
ServerArgs(model_path="dummy"),
role="test",
ranks=SpawnRanks(world_rank=0),
)
self.assertIsNone(get_device().gpu_id)
def test_no_controller_is_an_answer_not_a_failure(self): def test_no_controller_is_an_answer_not_a_failure(self):
"""`dp_rank=None` means "not under a data parallel controller", which """`dp_rank=None` means "not under a data parallel controller", which
is a fact about the deployment, unlike never having been told. The is a fact about the deployment, unlike never having been told. The
@@ -394,7 +424,7 @@ class TestSpawnIdentities(_IsolatedOverrides):
class TestAttentionRanksComeFromPublish(_IsolatedOverrides): class TestAttentionRanksComeFromPublish(_IsolatedOverrides):
"""With a spawn bundle, a rank read works before any group exists. """With a spawn bundle, a rank read works before any group exists.
This is what `ParallelState` provided by being a plain frozen record, and This is what the per-runner record provided by being a plain frozen object, and
what the topology init could not: it needs the groups. Deriving at publish what the topology init could not: it needs the groups. Deriving at publish
is what lets a reader ask the context in a process that never initialises is what lets a reader ask the context in a process that never initialises
distributed -- every unit test that builds a scheduler component, for one. distributed -- every unit test that builds a scheduler component, for one.
@@ -3068,6 +3098,58 @@ class TestWhoAnswersDuringADraftScope(CustomTestCase):
self.assertEqual((info.pp_rank, info.pp_size), (0, 1)) self.assertEqual((info.pp_rank, info.pp_size), (0, 1))
class TestTheRecordIsNeverWrittenTo(CustomTestCase):
"""`server_args` is the startup record; the bags are the truth afterwards.
Writing a field onto it after `resolve_once()` has sealed it puts a second
answer where there is supposed to be one, and it is invisible to anything
reading the bag. The sanctioned writer is `RuntimeContext.override`, which
writes the bag and says so in its own contract. `arg_groups/` is exempt: it
is the resolution pipeline, so building the record is its job.
"""
#: Assignments here are the record being built, not mutated behind a reader.
EXEMPT = ("srt/arg_groups/",)
def test_nothing_assigns_a_field_of_the_record(self):
import ast as _ast
from sglang.srt.arg_groups.arg_utils import namespace_of
from sglang.srt.server_args import ServerArgs
fields = set(namespace_of(ServerArgs))
offenders = []
for path in _sources():
rel = path.as_posix()
if "sglang/srt/" not in rel and "sglang/benchmark/" not in rel:
continue
if any(part in rel for part in self.EXEMPT):
continue
for node in _ast.walk(_ast.parse(path.read_text(encoding="utf-8-sig"))):
targets = (
node.targets
if isinstance(node, _ast.Assign)
else [node.target]
if isinstance(node, (_ast.AugAssign, _ast.AnnAssign))
else []
)
for target in targets:
if not isinstance(target, _ast.Attribute):
continue
base = target.value
name = getattr(base, "id", getattr(base, "attr", None))
if target.attr.startswith("_"):
continue
if name == "server_args" and target.attr in fields:
offenders.append(f"{rel}:{target.lineno} .{target.attr}")
self.assertEqual(
offenders,
[],
"write the bag through get_context().override(source, ...) instead "
"-- the record is not a channel:\n " + "\n ".join(offenders),
)
class TestNothingReadsThePlacementBeforeItIsFrozen(CustomTestCase): class TestNothingReadsThePlacementBeforeItIsFrozen(CustomTestCase):
"""`ModelRunner.__init__` freezes its placement partway through. """`ModelRunner.__init__` freezes its placement partway through.
@@ -20,7 +20,6 @@ from unittest.mock import patch
import torch import torch
from torch import nn from torch import nn
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.quantization.fp8_utils import ( from sglang.srt.layers.quantization.fp8_utils import (
quant_weight_ue8m0, quant_weight_ue8m0,
transform_scale_ue8m0, transform_scale_ue8m0,
@@ -181,15 +180,6 @@ class _FakeModelRunner:
attn_dp_size: int | None = None, attn_dp_size: int | None = None,
): ):
self.model = model self.model = model
self.ps = ParallelState.trivial(
tp_rank=tp_rank,
tp_size=tp_size,
dp_rank=dp_rank,
dp_size=dp_size,
attn_dp_size=attn_dp_size if attn_dp_size is not None else dp_size,
pp_rank=pp_rank,
pp_size=pp_size,
)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------