From 43797cc804a335b587be1563e600008fee91d0f3 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Sat, 16 May 2026 09:23:51 +0800 Subject: [PATCH] Bundle Scheduler rank/size fields into a frozen ParallelState (#25444) --- python/sglang/srt/disaggregation/decode.py | 4 +- python/sglang/srt/disaggregation/prefill.py | 8 +- .../srt/distributed/parallel_state_wrapper.py | 23 +++ python/sglang/srt/layers/dp_attention.py | 4 +- .../srt/managers/data_parallel_controller.py | 2 +- python/sglang/srt/managers/scheduler.py | 167 ++++++++++-------- .../srt/managers/scheduler_dp_attn_mixin.py | 4 +- .../scheduler_output_processor_mixin.py | 4 +- .../sglang/srt/managers/scheduler_pp_mixin.py | 42 ++--- .../srt/managers/scheduler_profiler_mixin.py | 38 ++-- .../observability/scheduler_metrics_mixin.py | 24 +-- .../srt/ray/data_parallel_controller.py | 2 +- python/sglang/srt/ray/scheduler_actor.py | 2 +- .../test_forward_pass_metrics.py | 36 +++- 14 files changed, 208 insertions(+), 152 deletions(-) create mode 100644 python/sglang/srt/distributed/parallel_state_wrapper.py diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 3cdf2af17..9ef9a2096 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -378,7 +378,7 @@ class DecodePreallocQueue: kv_args.engine_rank = self.tp_rank % (attn_tp_size) kv_args.pp_rank = self.pp_rank - kv_args.system_dp_rank = self.scheduler.dp_rank + kv_args.system_dp_rank = self.scheduler.ps.dp_rank if self.scheduler.enable_hisparse: # Direct-to-host: register host pool pointers so P writes to D's host memory host_pool = self.scheduler.hisparse_coordinator.mem_pool_host @@ -420,7 +420,7 @@ class DecodePreallocQueue: ) kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device - kv_args.gpu_id = self.scheduler.gpu_id + kv_args.gpu_id = self.scheduler.ps.gpu_id kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER) kv_manager = kv_manager_class( kv_args, diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 7ddcbe169..e2cc107de 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -144,7 +144,7 @@ class PrefillBootstrapQueue: kv_args = kv_args_class() kv_args.engine_rank = self.tp_rank kv_args.pp_rank = self.pp_rank - kv_args.system_dp_rank = self.scheduler.dp_rank + kv_args.system_dp_rank = self.scheduler.ps.dp_rank kv_args.prefill_start_layer = self.token_to_kv_pool.start_layer kv_data_ptrs, kv_data_lens, kv_item_lens = ( self.token_to_kv_pool.get_contiguous_buf_infos() @@ -174,7 +174,7 @@ class PrefillBootstrapQueue: self.metadata_buffers.get_buf_infos() ) kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device - kv_args.gpu_id = self.scheduler.gpu_id + kv_args.gpu_id = self.scheduler.ps.gpu_id req_to_token_pool = getattr(self.scheduler, "req_to_token_pool", None) setup_state_kv_args( @@ -620,7 +620,7 @@ class SchedulerDisaggregationPrefillMixin: KVPoll.Failed, ): logger.warning_once( - f"PP rank {self.pp_rank}: unexpected poll state {poll} for rid {req.rid} " + f"PP rank {self.ps.pp_rank}: unexpected poll state {poll} for rid {req.rid} " f"from consensus; treating as undone", ) undone_reqs.append(req) @@ -637,7 +637,7 @@ class SchedulerDisaggregationPrefillMixin: done_reqs.append(req) req.time_stats.set_prefill_kv_transfer_finish_time() elif poll == KVPoll.Failed: - error_message = f"Prefill transfer failed for request rank={self.tp_rank} {req.rid=} {req.bootstrap_room=}" + error_message = f"Prefill transfer failed for request rank={self.ps.tp_rank} {req.rid=} {req.bootstrap_room=}" try: req.disagg_kv_sender.failure_exception() except Exception as e: diff --git a/python/sglang/srt/distributed/parallel_state_wrapper.py b/python/sglang/srt/distributed/parallel_state_wrapper.py new file mode 100644 index 000000000..bf692f5e9 --- /dev/null +++ b/python/sglang/srt/distributed/parallel_state_wrapper.py @@ -0,0 +1,23 @@ +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_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 diff --git a/python/sglang/srt/layers/dp_attention.py b/python/sglang/srt/layers/dp_attention.py index 89d238928..fc4182b96 100644 --- a/python/sglang/srt/layers/dp_attention.py +++ b/python/sglang/srt/layers/dp_attention.py @@ -251,7 +251,7 @@ def compute_dp_attention_world_info( # tp_rank = (attn_dp_rank * attn_cp_size + attn_cp_rank) * attn_tp_size + attn_tp_rank attn_dp_rank = tp_rank // (attn_tp_size * attn_cp_size) - return attn_tp_rank, attn_tp_size, attn_dp_rank + return attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size def compute_dp_attention_local_info( @@ -287,7 +287,7 @@ def initialize_dp_attention( tp_rank = get_tensor_model_parallel_rank() tp_size = get_tensor_model_parallel_world_size() - _, _, _ATTN_DP_RANK = compute_dp_attention_world_info( + _, _, _ATTN_DP_RANK, _ = compute_dp_attention_world_info( enable_dp_attention, tp_rank, tp_size, dp_size, attn_cp_size ) _, _, _LOCAL_ATTN_DP_RANK = compute_dp_attention_local_info( diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index e86c3140e..bf158695f 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -480,7 +480,7 @@ class DataParallelController: if server_args.enable_dp_attention: # dp attention has different sharding logic - _, _, dp_rank = compute_dp_attention_world_info( + _, _, dp_rank, _ = compute_dp_attention_world_info( server_args.enable_dp_attention, tp_rank, server_args.tp_size, diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 43290075a..4091d884a 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -64,6 +64,7 @@ from sglang.srt.disaggregation.utils import ( ) from sglang.srt.distributed import get_pp_group, get_world_group from sglang.srt.distributed.parallel_state import get_tp_group +from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.dllm.mixin.scheduler import SchedulerDllmMixin from sglang.srt.environ import envs from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder @@ -356,18 +357,6 @@ class Scheduler( # Parse args self.server_args = server_args - self.tp_rank = tp_rank - self.moe_ep_rank = moe_ep_rank - self.pp_rank = pp_rank - self.attn_cp_rank = attn_cp_rank - self.attn_cp_size = server_args.attn_cp_size - self.moe_dp_rank = moe_dp_rank - self.moe_dp_size = server_args.moe_dp_size - self.dp_rank = dp_rank - self.tp_size = server_args.tp_size - self.moe_ep_size = server_args.ep_size - self.pp_size = server_args.pp_size - self.dp_size = server_args.dp_size self.nccl_port = port_args.nccl_port self.schedule_policy = server_args.schedule_policy self.enable_priority_scheduling = server_args.enable_priority_scheduling @@ -391,7 +380,6 @@ class Scheduler( self.spec_algorithm = SpeculativeAlgorithm.from_string( server_args.speculative_algorithm ) - self.gpu_id = gpu_id self.page_size = server_args.page_size self.enable_hierarchical_cache = server_args.enable_hierarchical_cache self.enable_hicache_storage = server_args.hicache_storage_backend is not None @@ -400,15 +388,34 @@ class Scheduler( self.hisparse_coordinator: Optional[HiSparseCoordinator] = None # Distributed rank info - self.attn_tp_rank, self.attn_tp_size, self.attn_dp_rank = ( + attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = ( compute_dp_attention_world_info( server_args.enable_dp_attention, - self.tp_rank, - self.tp_size, - self.dp_size, - self.attn_cp_size, + tp_rank, + server_args.tp_size, + server_args.dp_size, + server_args.attn_cp_size, ) ) + self.ps = ParallelState( + tp_rank=tp_rank, + tp_size=server_args.tp_size, + pp_rank=pp_rank, + pp_size=server_args.pp_size, + dp_rank=dp_rank, + dp_size=server_args.dp_size, + attn_tp_rank=attn_tp_rank, + attn_tp_size=attn_tp_size, + attn_cp_rank=attn_cp_rank, + attn_cp_size=server_args.attn_cp_size, + attn_dp_rank=attn_dp_rank, + attn_dp_size=attn_dp_size, + moe_ep_rank=moe_ep_rank, + moe_ep_size=server_args.ep_size, + moe_dp_rank=moe_dp_rank, + moe_dp_size=server_args.moe_dp_size, + gpu_id=gpu_id, + ) # Init model configs self.init_model_config() @@ -505,10 +512,10 @@ class Scheduler( if _is_npu: from sglang.srt.hardware_backend.npu.utils import init_zbal - if self.pp_size > 1: + if self.ps.pp_size > 1: logger.error(f"only zbal mix mode support pp_size > 1!") init_zbal( - self.tp_size, self.gpu_id, self.tp_rank + self.ps.tp_size, self.ps.gpu_id, self.ps.tp_rank ) # only switch allocator if is mix mode def init_model_config(self): @@ -536,7 +543,11 @@ class Scheduler( context = zmq.Context(2) self.idle_sleeper = None - if self.pp_rank == 0 and self.attn_tp_rank == 0 and self.attn_cp_rank == 0: + if ( + self.ps.pp_rank == 0 + and self.ps.attn_tp_rank == 0 + and self.ps.attn_cp_rank == 0 + ): self.recv_from_tokenizer = get_zmq_socket( context, zmq.PULL, port_args.scheduler_input_ipc_name, False ) @@ -654,13 +665,13 @@ class Scheduler( def init_tp_model_worker(self): worker_kwargs = dict( server_args=self.server_args, - gpu_id=self.gpu_id, - tp_rank=self.tp_rank, - moe_ep_rank=self.moe_ep_rank, - pp_rank=self.pp_rank, - attn_cp_rank=self.attn_cp_rank, - moe_dp_rank=self.moe_dp_rank, - dp_rank=self.dp_rank, + gpu_id=self.ps.gpu_id, + tp_rank=self.ps.tp_rank, + moe_ep_rank=self.ps.moe_ep_rank, + pp_rank=self.ps.pp_rank, + attn_cp_rank=self.ps.attn_cp_rank, + moe_dp_rank=self.ps.moe_dp_rank, + dp_rank=self.ps.dp_rank, nccl_port=self.nccl_port, ) @@ -683,14 +694,14 @@ class Scheduler( # Launch a draft worker for speculative decoding draft_worker_kwargs = dict( server_args=self.server_args, - gpu_id=self.gpu_id, - tp_rank=self.tp_rank, - moe_ep_rank=self.moe_ep_rank, + gpu_id=self.ps.gpu_id, + tp_rank=self.ps.tp_rank, + moe_ep_rank=self.ps.moe_ep_rank, nccl_port=self.nccl_port, target_worker=self.tp_worker, - dp_rank=self.dp_rank, - attn_cp_rank=self.attn_cp_rank, - moe_dp_rank=self.moe_dp_rank, + dp_rank=self.ps.dp_rank, + attn_cp_rank=self.ps.attn_cp_rank, + moe_dp_rank=self.ps.moe_dp_rank, ) if self.server_args.speculative_draft_load_format is not None: @@ -743,7 +754,7 @@ class Scheduler( ) = self.tp_worker.get_worker_info() if not get_global_server_args().pp_max_micro_batch_size: get_global_server_args().pp_max_micro_batch_size = max( - self.max_running_requests // self.pp_size, 1 + self.max_running_requests // self.ps.pp_size, 1 ) self.tp_group = get_tp_group() @@ -772,9 +783,9 @@ class Scheduler( # Print debug info avail_mem = get_available_gpu_memory( - self.device, self.gpu_id, empty_cache=False + self.device, self.ps.gpu_id, empty_cache=False ) - if self.tp_rank == 0: + if self.ps.tp_rank == 0: logger.info( f"max_total_num_tokens={self.max_total_num_tokens}, " f"chunked_prefill_size={self.server_args.chunked_prefill_size}, " @@ -876,8 +887,8 @@ class Scheduler( enable_metrics=self.enable_metrics, enable_kv_cache_events=self.enable_kv_cache_events, enable_mamba_extra_buffer=server_args.enable_mamba_extra_buffer(), - pp_rank=self.pp_rank, - pp_size=self.pp_size, + pp_rank=self.ps.pp_rank, + pp_size=self.ps.pp_size, chunked_prefill_size=effective_chunked_prefill_size, sliding_window_size=self.sliding_window_size, ) @@ -952,8 +963,8 @@ class Scheduler( self.tree_cache = LMCRadixCache( params=params, model_config=self.model_config, - tp_size=self.tp_size, - rank=self.tp_rank, + tp_size=self.ps.tp_size, + rank=self.ps.tp_rank, tp_group=self.tp_group, ) else: @@ -1102,7 +1113,7 @@ class Scheduler( # Init the dynamic chunking predictor for PP self.enable_dynamic_chunking = ( - self.server_args.enable_dynamic_chunking and self.pp_size > 1 + self.server_args.enable_dynamic_chunking and self.ps.pp_size > 1 ) if self.enable_dynamic_chunking: try: @@ -1133,8 +1144,8 @@ class Scheduler( ) else: self.prefill_delayer = PrefillDelayer( - dp_size=self.dp_size, - attn_tp_size=self.attn_tp_size, + dp_size=self.ps.dp_size, + attn_tp_size=self.ps.attn_tp_size, cpu_group=self.tp_cpu_group, device_group=self.tp_group.device_group, server_args=self.server_args, @@ -1187,7 +1198,7 @@ class Scheduler( # Init recv skipper and input blocker self.recv_skipper = SchedulerRecvSkipper.maybe_create(self.server_args) self.input_blocker = ( - SchedulerInputBlocker(noop=self.attn_tp_rank != 0) + SchedulerInputBlocker(noop=self.ps.attn_tp_rank != 0) if get_bool_env_var("SGLANG_ENABLE_COLOCATED_BATCH_GEN") else None ) @@ -1238,7 +1249,7 @@ class Scheduler( self.disagg_decode_transfer_queue = DecodeTransferQueue( gloo_group=self.attn_tp_cpu_group, req_to_metadata_buffer_idx_allocator=self.req_to_metadata_buffer_idx_allocator, - tp_rank=self.tp_rank, + tp_rank=self.ps.tp_rank, metadata_buffers=self.disagg_metadata_buffers, scheduler=self, tree_cache=self.tree_cache, @@ -1255,13 +1266,13 @@ class Scheduler( transfer_queue=self.disagg_decode_transfer_queue, tree_cache=self.tree_cache, gloo_group=self.attn_tp_cpu_group, - tp_rank=self.tp_rank, - tp_size=self.tp_size, + tp_rank=self.ps.tp_rank, + tp_size=self.ps.tp_size, dp_size=self.server_args.dp_size, - gpu_id=self.gpu_id, + gpu_id=self.ps.gpu_id, bootstrap_port=self.server_args.disaggregation_bootstrap_port, max_total_num_tokens=self.max_total_num_tokens, - pp_rank=self.pp_rank, + pp_rank=self.ps.pp_rank, num_reserved_decode_tokens=self.server_args.num_reserved_decode_tokens, transfer_backend=self.transfer_backend, ) @@ -1294,15 +1305,15 @@ class Scheduler( draft_token_to_kv_pool=draft_token_to_kv_pool, req_to_metadata_buffer_idx_allocator=self.req_to_metadata_buffer_idx_allocator, metadata_buffers=self.disagg_metadata_buffers, - tp_rank=self.tp_rank, - tp_size=self.tp_size, - gpu_id=self.gpu_id, + tp_rank=self.ps.tp_rank, + tp_size=self.ps.tp_size, + gpu_id=self.ps.gpu_id, bootstrap_port=self.server_args.disaggregation_bootstrap_port, gloo_group=self.attn_tp_cpu_group, max_total_num_tokens=self.max_total_num_tokens, scheduler=self, - pp_rank=self.pp_rank, - pp_size=self.pp_size, + pp_rank=self.ps.pp_rank, + pp_size=self.ps.pp_size, transfer_backend=self.transfer_backend, ) # The prefill requests that are in the middle of kv sending @@ -1316,8 +1327,8 @@ class Scheduler( self.mm_receiver = create_mm_receiver( self.server_args, hf_config=self.model_config.hf_config, - pp_rank=self.pp_rank, - tp_rank=self.tp_rank, + pp_rank=self.ps.pp_rank, + tp_rank=self.ps.tp_rank, tp_group=self.tp_group, scheduler=self, ) @@ -1667,8 +1678,8 @@ class Scheduler( if not self.recv_skipper.handle(last_forward_mode): return [] - if self.pp_rank == 0: - if self.attn_tp_rank == 0 and self.attn_cp_rank == 0: + if self.ps.pp_rank == 0: + if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0: recv_reqs = [] while True: @@ -1691,14 +1702,14 @@ class Scheduler( else: recv_reqs = None else: - if self.attn_tp_rank == 0 and self.attn_cp_rank == 0: - dp_offset = self.attn_dp_rank * self.attn_tp_size + if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0: + dp_offset = self.ps.attn_dp_rank * self.ps.attn_tp_size recv_reqs = point_to_point_pyobj( [], - self.pp_rank * self.tp_size + dp_offset, + self.ps.pp_rank * self.ps.tp_size + dp_offset, self.world_group.cpu_group, - (self.pp_rank - 1) * self.tp_size + dp_offset, - self.pp_rank * self.tp_size + dp_offset, + (self.ps.pp_rank - 1) * self.ps.tp_size + dp_offset, + self.ps.pp_rank * self.ps.tp_size + dp_offset, ) else: recv_reqs = None @@ -1707,13 +1718,13 @@ class Scheduler( recv_reqs = self.input_blocker.handle(recv_reqs) if self.server_args.enable_dp_attention: - if self.attn_tp_rank == 0 and self.attn_cp_rank == 0: + if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0: work_reqs, control_reqs = self._split_work_and_control_reqs(recv_reqs) else: work_reqs = None control_reqs = None - if self.attn_tp_size != 1: + if self.ps.attn_tp_size != 1: work_reqs = broadcast_pyobj( work_reqs, self.attn_tp_group.rank, @@ -1721,7 +1732,7 @@ class Scheduler( src=self.attn_tp_group.ranks[0], ) - if self.attn_cp_size != 1: + if self.ps.attn_cp_size != 1: work_reqs = broadcast_pyobj( work_reqs, self.attn_cp_group.rank, @@ -1736,21 +1747,21 @@ class Scheduler( # all-ranks gloo sync. _local_ctrl = self.server_args.enable_dp_attention_local_control_broadcast if _local_ctrl: - if self.attn_tp_size != 1: + if self.ps.attn_tp_size != 1: control_reqs = broadcast_pyobj( control_reqs, self.attn_tp_group.rank, self.attn_tp_cpu_group, src=self.attn_tp_group.ranks[0], ) - if self.attn_cp_size != 1: + if self.ps.attn_cp_size != 1: control_reqs = broadcast_pyobj( control_reqs, self.attn_cp_group.rank, self.attn_cp_cpu_group, src=self.attn_cp_group.ranks[0], ) - elif self.tp_size != 1: + elif self.ps.tp_size != 1: control_reqs = broadcast_pyobj( control_reqs, self.tp_group.rank, @@ -1758,7 +1769,7 @@ class Scheduler( src=self.tp_group.ranks[0], ) recv_reqs = work_reqs + control_reqs - elif self.tp_size != 1: + elif self.ps.tp_size != 1: recv_reqs = broadcast_pyobj( recv_reqs, self.tp_group.rank, @@ -1768,7 +1779,7 @@ class Scheduler( # Process MM requests under EPD-disaggregation mode if ( - self.pp_rank == 0 + self.ps.pp_rank == 0 and self.server_args.language_only and self.server_args.encoder_transfer_backend == "zmq_to_scheduler" ): @@ -1802,7 +1813,7 @@ class Scheduler( # removes the name; already-open handles stay valid. if ( not self.server_args.enable_dp_attention - and self.tp_size > 1 + and self.ps.tp_size > 1 and self.model_config.is_multimodal and has_shm_features(recv_reqs) ): @@ -3129,7 +3140,7 @@ class Scheduler( tp_active_ranks = self.tp_group.active_ranks.detach().cpu().numpy() tp_active_ranks_cpu = self.tp_group.active_ranks_cpu.detach().numpy() tp_active_ranks &= tp_active_ranks_cpu - dp_active_ranks = tp_active_ranks.reshape(self.dp_size, -1).prod(axis=1) + dp_active_ranks = tp_active_ranks.reshape(self.ps.dp_size, -1).prod(axis=1) self.send_to_tokenizer.send_output( ActiveRanksOutput(status=dp_active_ranks.tolist()) ) @@ -3303,7 +3314,7 @@ class Scheduler( and (self.last_batch is None or self.last_batch.is_empty()) and (self.cur_batch is None or self.cur_batch.is_empty()) and (not self.enable_overlap or len(self.result_queue) == 0) - and (self.pp_size == 1 or all(x.is_empty() for x in self.running_mbs)) + and (self.ps.pp_size == 1 or all(x.is_empty() for x in self.running_mbs)) ) # Waiting queues: waiting + bootstrapping + preallocation + kv transfer (decode) @@ -3506,10 +3517,10 @@ class Scheduler( if_success = False break elif k == "pp_max_micro_batch_size" and ( - v > self.max_running_requests // self.pp_size or v < 1 + v > self.max_running_requests // self.ps.pp_size or v < 1 ): logging.warning( - f"Updating {k} to {v} is rejected because it is out of the valid range [1, {self.max_running_requests // self.pp_size}]." + f"Updating {k} to {v} is rejected because it is out of the valid range [1, {self.max_running_requests // self.ps.pp_size}]." ) if_success = False break @@ -3776,7 +3787,7 @@ class Scheduler( def open_session(self, recv_req: OpenSessionReqInput): output = self.session_controller.open(recv_req) - if self.pp_rank == 0 and self.tp_rank == 0 and self.attn_cp_rank == 0: + if self.ps.pp_rank == 0 and self.ps.tp_rank == 0 and self.ps.attn_cp_rank == 0: return output return None diff --git a/python/sglang/srt/managers/scheduler_dp_attn_mixin.py b/python/sglang/srt/managers/scheduler_dp_attn_mixin.py index 5331fc033..b135f2d2d 100644 --- a/python/sglang/srt/managers/scheduler_dp_attn_mixin.py +++ b/python/sglang/srt/managers/scheduler_dp_attn_mixin.py @@ -230,8 +230,8 @@ class SchedulerDPAttnMixin: return prepare_mlp_sync_batch_raw( local_batch, dp_size=self.server_args.dp_size, - attn_tp_size=self.attn_tp_size, - attn_cp_size=self.attn_cp_size, + attn_tp_size=self.ps.attn_tp_size, + attn_cp_size=self.ps.attn_cp_size, tp_group=self.tp_group, get_idle_batch=self.get_idle_batch, disable_cuda_graph=self.server_args.disable_cuda_graph, diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index ea2c38c0d..85ee86d2a 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -1253,12 +1253,12 @@ class SchedulerOutputProcessorMixin: if ( req.finished() - and self.attn_tp_rank == 0 + and self.ps.attn_tp_rank == 0 and self.server_args.enable_request_time_stats_logging ): req.log_time_stats() - dp_ranks = [self.dp_rank] * len(rids) if rids else None + dp_ranks = [self.ps.dp_rank] * len(rids) if rids else None # Send to detokenizer if reqs or is_idle_batch: diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index 0b2270ce9..ab1cc6d2a 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -77,7 +77,7 @@ class SchedulerPPMixin: for mb_id in range(self.pp_loop_size): self.running_batch = self.running_mbs[mb_id] self.last_batch = self.last_mbs[mb_id] - next_first_rank_mb_id = (mb_id + self.pp_size) % self.pp_loop_size + next_first_rank_mb_id = (mb_id + self.ps.pp_size) % self.pp_loop_size next_mb_id = (mb_id + 1) % self.pp_loop_size with torch.profiler.record_function("recv_requests"): recv_reqs = self.recv_requests() @@ -205,7 +205,7 @@ class SchedulerPPMixin: for mb_id in range(self.pp_loop_size): self.running_batch = self.running_mbs[mb_id] self.last_batch = self.last_mbs[mb_id] - next_first_rank_mb_id = (mb_id + self.pp_size) % self.pp_loop_size + next_first_rank_mb_id = (mb_id + self.ps.pp_size) % self.pp_loop_size next_mb_id = (mb_id + 1) % self.pp_loop_size next_pp_outputs = None @@ -350,7 +350,7 @@ class SchedulerPPMixin: for mb_id in range(self.pp_loop_size): self.running_batch = self.running_mbs[mb_id] self.last_batch = self.last_mbs[mb_id] - next_first_rank_mb_id = (mb_id + self.pp_size) % self.pp_loop_size + next_first_rank_mb_id = (mb_id + self.ps.pp_size) % self.pp_loop_size next_mb_id = (mb_id + 1) % self.pp_loop_size next_pp_outputs = None @@ -520,7 +520,7 @@ class SchedulerPPMixin: self.on_idle() def init_pp_loop_state(self: Scheduler): - self.pp_loop_size: int = self.pp_size + self.server_args.pp_async_batch_depth + self.pp_loop_size: int = self.ps.pp_size + self.server_args.pp_async_batch_depth # In CP mode, attention weights are duplicated, eliminating the need for the attention TP all-gather operation. self.require_attn_tp_allgather = ( not self.server_args.enable_nsa_prefill_context_parallel @@ -665,7 +665,7 @@ class SchedulerPPMixin: f"seq_lens={seq_lens}, latencies_ms={latencies}" ) - if self.attn_tp_size > 1: + if self.ps.attn_tp_size > 1: data_to_sync_tp = [seq_lens, latencies] data_to_sync_tp = broadcast_pyobj( data_to_sync_tp, @@ -675,7 +675,7 @@ class SchedulerPPMixin: ) seq_lens, latencies = data_to_sync_tp - if self.attn_cp_size > 1: + if self.ps.attn_cp_size > 1: data_to_sync_tp = [seq_lens, latencies] data_to_sync_tp = broadcast_pyobj( data_to_sync_tp, @@ -696,7 +696,7 @@ class SchedulerPPMixin: self.length_predictor.set_target_latency(self.chunked_prefill_size) self.length_predictor.is_ready = True logger.info( - f"[PP Dynamic Chunk] [PP{self.pp_rank}] Predictor ready (quadratic). " + f"[PP Dynamic Chunk] [PP{self.ps.pp_rank}] Predictor ready (quadratic). " f"Target latency: {self.length_predictor.target_latency:.2f}ms" ) @@ -728,7 +728,7 @@ class SchedulerPPMixin: if predicted_size is not None: logger.debug( - f"[PP Dynamic Chunk] [PP{self.pp_rank}] Predicted chunk size: " + f"[PP Dynamic Chunk] [PP{self.ps.pp_rank}] Predicted chunk size: " f"{predicted_size} (history_len={history_len})" ) @@ -886,32 +886,32 @@ class SchedulerPPMixin: def _pp_send_pyobj_to_next_stage(self: Scheduler, data, async_send: bool = False): p2p_work = [] - if self.attn_tp_rank == 0 and self.attn_cp_rank == 0: - dp_offset = self.attn_dp_rank * self.attn_tp_size + if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0: + dp_offset = self.ps.attn_dp_rank * self.ps.attn_tp_size p2p_work = point_to_point_pyobj( data, - self.pp_rank * self.tp_size + dp_offset, + self.ps.pp_rank * self.ps.tp_size + dp_offset, self.world_group.cpu_group, - self.pp_rank * self.tp_size + dp_offset, - ((self.pp_rank + 1) % self.pp_size) * self.tp_size + dp_offset, + self.ps.pp_rank * self.ps.tp_size + dp_offset, + ((self.ps.pp_rank + 1) % self.ps.pp_size) * self.ps.tp_size + dp_offset, async_send=async_send, ) return p2p_work def _pp_recv_pyobj_from_prev_stage(self: Scheduler): - if self.attn_tp_rank == 0 and self.attn_cp_rank == 0: - dp_offset = self.attn_dp_rank * self.attn_tp_size + if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0: + dp_offset = self.ps.attn_dp_rank * self.ps.attn_tp_size data = point_to_point_pyobj( [], - self.pp_rank * self.tp_size + dp_offset, + self.ps.pp_rank * self.ps.tp_size + dp_offset, self.world_group.cpu_group, - ((self.pp_rank - 1) % self.pp_size) * self.tp_size + dp_offset, - self.pp_rank * self.tp_size + dp_offset, + ((self.ps.pp_rank - 1) % self.ps.pp_size) * self.ps.tp_size + dp_offset, + self.ps.pp_rank * self.ps.tp_size + dp_offset, ) else: data = None - if self.attn_tp_size > 1: + if self.ps.attn_tp_size > 1: data = broadcast_pyobj( data, self.attn_tp_group.rank, @@ -919,7 +919,7 @@ class SchedulerPPMixin: src=self.attn_tp_group.ranks[0], ) - if self.attn_cp_size > 1: + if self.ps.attn_cp_size > 1: data = broadcast_pyobj( data, self.attn_cp_group.rank, @@ -1120,7 +1120,7 @@ class SchedulerPPMixin: # CUDA: send first # XPU: even ranks send first, odd ranks recv first. - send_first = (not is_xpu()) or ((self.pp_rank % 2) == 0) + send_first = (not is_xpu()) or ((self.ps.pp_rank % 2) == 0) def _do_send(): return self._pp_send_output_to_next_stage( diff --git a/python/sglang/srt/managers/scheduler_profiler_mixin.py b/python/sglang/srt/managers/scheduler_profiler_mixin.py index c02ed7997..77585988f 100644 --- a/python/sglang/srt/managers/scheduler_profiler_mixin.py +++ b/python/sglang/srt/managers/scheduler_profiler_mixin.py @@ -38,9 +38,9 @@ class SchedulerProfilerMixin: def init_profiler(self: Scheduler): if envs.SGLANG_PROFILE_V2.get(): self._profile_manager = ProfileManager( - tp_rank=self.tp_rank, + tp_rank=self.ps.tp_rank, cpu_group=self.dp_tp_cpu_group, - gpu_id=self.gpu_id, + gpu_id=self.ps.gpu_id, ) return @@ -167,10 +167,10 @@ class SchedulerProfilerMixin: self.rpd_profile_path = os.path.join( self.torch_profiler_output_dir, - "rpd-" + str(time.time()) + f"-TP-{self.tp_rank}" + ".trace.json.gz", + "rpd-" + str(time.time()) + f"-TP-{self.ps.tp_rank}" + ".trace.json.gz", ) - if self.tp_rank == 0: + if self.ps.tp_rank == 0: import sqlite3 from rocpd.schema import RocpdSchema @@ -210,7 +210,7 @@ class SchedulerProfilerMixin: self.profile_in_progress = True if "CUDA_PROFILER" in activities: - if self.gpu_id == get_global_server_args().base_gpu_id: + if self.ps.gpu_id == get_global_server_args().base_gpu_id: torch.cuda.cudart().cudaProfilerStart() self.profile_in_progress = True @@ -220,13 +220,13 @@ class SchedulerProfilerMixin: if not self.merge_profiles: return "" - if self.tp_rank != 0: + if self.ps.tp_rank != 0: return "" - if getattr(self, "dp_size", 1) > 1 and getattr(self, "dp_rank", 0) != 0: + if self.ps.dp_size > 1 and self.ps.dp_rank != 0: return "" - if getattr(self, "pp_size", 1) > 1 and getattr(self, "pp_rank", 0) != 0: + if self.ps.pp_size > 1 and self.ps.pp_rank != 0: return "" - if getattr(self, "moe_ep_size", 1) > 1 and getattr(self, "moe_ep_rank", 0) != 0: + if self.ps.moe_ep_size > 1 and self.ps.moe_ep_rank != 0: return "" try: @@ -273,15 +273,15 @@ class SchedulerProfilerMixin: self.torch_profiler.stop() if not _is_npu: # Build filename with only non-zero ranks to maintain backward compatibility - filename_parts = [self.profile_id, f"TP-{self.tp_rank}"] + filename_parts = [self.profile_id, f"TP-{self.ps.tp_rank}"] # Only add other ranks if parallelism is enabled (size > 1) - if getattr(self, "dp_size", 1) > 1: - filename_parts.append(f"DP-{getattr(self, 'dp_rank', 0)}") - if getattr(self, "pp_size", 1) > 1: - filename_parts.append(f"PP-{getattr(self, 'pp_rank', 0)}") - if getattr(self, "moe_ep_size", 1) > 1: - filename_parts.append(f"EP-{getattr(self, 'moe_ep_rank', 0)}") + if self.ps.dp_size > 1: + filename_parts.append(f"DP-{self.ps.dp_rank}") + if self.ps.pp_size > 1: + filename_parts.append(f"PP-{self.ps.pp_rank}") + if self.ps.moe_ep_size > 1: + filename_parts.append(f"EP-{self.ps.moe_ep_rank}") filename = ( stage_prefix @@ -301,7 +301,7 @@ class SchedulerProfilerMixin: self.rpd_profiler.flush() torch.distributed.barrier(self.dp_tp_cpu_group) - if self.tp_rank == 0: + if self.ps.tp_rank == 0: from sglang.srt.utils.rpd_utils import rpd_to_chrome_trace rpd_to_chrome_trace("trace.rpd", self.rpd_profile_path) @@ -312,7 +312,7 @@ class SchedulerProfilerMixin: memory_profile_path = os.path.join( self.torch_profiler_output_dir, str(time.time()) - + f"-TP-{self.tp_rank}-memory" + + f"-TP-{self.ps.tp_rank}-memory" + stage_suffix + ".pickle", ) @@ -320,7 +320,7 @@ class SchedulerProfilerMixin: torch.cuda.memory._record_memory_history(enabled=None) if "CUDA_PROFILER" in self.profiler_activities: - if self.gpu_id == get_global_server_args().base_gpu_id: + if self.ps.gpu_id == get_global_server_args().base_gpu_id: torch.cuda.cudart().cudaProfilerStop() merge_message = self._merge_profile_traces() diff --git a/python/sglang/srt/observability/scheduler_metrics_mixin.py b/python/sglang/srt/observability/scheduler_metrics_mixin.py index a252474b4..9a2ab6f64 100644 --- a/python/sglang/srt/observability/scheduler_metrics_mixin.py +++ b/python/sglang/srt/observability/scheduler_metrics_mixin.py @@ -121,7 +121,7 @@ class SchedulerMetricsMixin: # Metrics self.enable_metrics = self.server_args.enable_metrics - self.is_stats_logging_rank = self.attn_tp_rank == 0 + self.is_stats_logging_rank = self.ps.attn_tp_rank == 0 self.current_scheduler_metrics_enabled = self.enable_metrics and ( self.is_stats_logging_rank or self.server_args.enable_metrics_for_all_schedulers @@ -138,7 +138,7 @@ class SchedulerMetricsMixin: "engine_type": engine_type, "tp_rank": tp_rank, "pp_rank": pp_rank, - "moe_ep_rank": self.moe_ep_rank, + "moe_ep_rank": self.ps.moe_ep_rank, } if self.enable_priority_scheduling: labels["priority"] = "" @@ -201,12 +201,12 @@ class SchedulerMetricsMixin: def init_kv_events(self: Scheduler, kv_events_config: Optional[str]): self.enable_kv_cache_events = bool( - kv_events_config and self.attn_tp_rank == 0 and self.attn_cp_rank == 0 + kv_events_config and self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0 ) if self.enable_kv_cache_events: self.kv_event_publisher = EventPublisherFactory.create( - kv_events_config, self.attn_dp_rank + kv_events_config, self.ps.attn_dp_rank ) def _init_fpm(self: Scheduler): @@ -214,14 +214,14 @@ class SchedulerMetricsMixin: self.enable_fpm = False if ( self.server_args.enable_forward_pass_metrics - and self.attn_tp_rank == 0 - and self.pp_rank == self.pp_size - 1 + and self.ps.attn_tp_rank == 0 + and self.ps.pp_rank == self.ps.pp_size - 1 ): from sglang.srt.observability.forward_pass_metrics import ( _FpmPublisherThread, ) - self._fpm_dp_rank = self.dp_rank if self.dp_rank is not None else 0 + self._fpm_dp_rank = self.ps.dp_rank if self.ps.dp_rank is not None else 0 self._fpm_worker_id = self.server_args.forward_pass_metrics_worker_id base_endpoint = self.server_args.forward_pass_metrics_ipc_name if base_endpoint is None: @@ -345,8 +345,8 @@ class SchedulerMetricsMixin: hidden_size = float(model_config.hidden_size) num_layers = float(getattr(model_config, "num_attention_layers", 0)) head_dim = float(getattr(model_config, "head_dim", 0)) - num_attn_heads = float(model_config.get_num_attention_heads(self.tp_size)) - num_kv_heads = float(model_config.get_num_kv_heads(self.tp_size)) + num_attn_heads = float(model_config.get_num_attention_heads(self.ps.tp_size)) + num_kv_heads = float(model_config.get_num_kv_heads(self.ps.tp_size)) intermediate_size = getattr(hf_text_config, "intermediate_size", None) if intermediate_size is None: intermediate_size = getattr(hf_text_config, "ffn_hidden_size", 0) @@ -834,7 +834,9 @@ class SchedulerMetricsMixin: kv_metrics.num_requests_waiting = self.stats.num_queue_reqs.total kv_metrics.gpu_cache_usage_perc = self.stats.token_usage kv_metrics.gpu_prefix_cache_hit_rate = self.stats.cache_hit_rate - kv_metrics.data_parallel_rank = self.dp_rank if self.dp_rank is not None else 0 + kv_metrics.data_parallel_rank = ( + self.ps.dp_rank if self.ps.dp_rank is not None else 0 + ) if not self.send_metrics_from_scheduler.closed: self.send_metrics_from_scheduler.send_pyobj(kv_metrics) @@ -1092,7 +1094,7 @@ class SchedulerMetricsMixin: ) return GetLoadsReqOutput( - dp_rank=self.dp_rank, + dp_rank=self.ps.dp_rank, timestamp=time.time(), num_running_reqs=num_running_reqs, num_waiting_reqs=num_waiting_reqs, diff --git a/python/sglang/srt/ray/data_parallel_controller.py b/python/sglang/srt/ray/data_parallel_controller.py index bf118b734..a255ad049 100644 --- a/python/sglang/srt/ray/data_parallel_controller.py +++ b/python/sglang/srt/ray/data_parallel_controller.py @@ -147,7 +147,7 @@ class RayDataParallelController(DataParallelController): if server_args.enable_dp_attention: # DP attention: derive dp_rank from tp_rank - _, _, actual_dp_rank = compute_dp_attention_world_info( + _, _, actual_dp_rank, _ = compute_dp_attention_world_info( server_args.enable_dp_attention, tp_rank, server_args.tp_size, diff --git a/python/sglang/srt/ray/scheduler_actor.py b/python/sglang/srt/ray/scheduler_actor.py index 13f588ebb..e9090ec9a 100644 --- a/python/sglang/srt/ray/scheduler_actor.py +++ b/python/sglang/srt/ray/scheduler_actor.py @@ -126,7 +126,7 @@ class SchedulerActor: import torch # Need to set the GPU id for the event loop for nccl to work - torch.cuda.set_device(self.scheduler.gpu_id) + torch.cuda.set_device(self.scheduler.ps.gpu_id) self.scheduler.run_event_loop() except Exception as e: logger.error(f"Scheduler PP{self._pp_rank} TP{self._tp_rank} crashed: {e}") diff --git a/test/registered/unit/observability/test_forward_pass_metrics.py b/test/registered/unit/observability/test_forward_pass_metrics.py index 4eda3665c..d63c89993 100644 --- a/test/registered/unit/observability/test_forward_pass_metrics.py +++ b/test/registered/unit/observability/test_forward_pass_metrics.py @@ -7,12 +7,38 @@ import unittest from unittest.mock import patch from sglang.srt.disaggregation.utils import DisaggregationMode +from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.observability.scheduler_metrics_mixin import ( PrefillStats, SchedulerMetricsMixin, ) +def _make_ps(**overrides) -> ParallelState: + """Build a ParallelState with reasonable defaults for tests; override fields via kwargs.""" + defaults = dict( + tp_rank=0, + tp_size=1, + pp_rank=0, + pp_size=1, + dp_rank=None, + dp_size=1, + attn_tp_rank=0, + attn_tp_size=1, + attn_cp_rank=0, + attn_cp_size=1, + attn_dp_rank=0, + attn_dp_size=1, + moe_ep_rank=0, + moe_ep_size=1, + moe_dp_rank=None, + moe_dp_size=1, + gpu_id=0, + ) + defaults.update(overrides) + return ParallelState(**defaults) + + class _FakeReq: def __init__( self, @@ -221,10 +247,7 @@ class TestForwardPassMetrics(unittest.TestCase): forward_pass_metrics_ipc_name=None, kv_events_config=None, ) - scheduler.attn_tp_rank = 0 - scheduler.dp_rank = 2 - scheduler.pp_rank = 0 - scheduler.pp_size = 1 + scheduler.ps = _make_ps(attn_tp_rank=0, dp_rank=2, pp_rank=0, pp_size=1) scheduler.enable_kv_cache_events = False with patch( @@ -252,10 +275,7 @@ class TestForwardPassMetrics(unittest.TestCase): forward_pass_metrics_ipc_name=None, kv_events_config=None, ) - scheduler.attn_tp_rank = 0 - scheduler.dp_rank = 0 - scheduler.pp_rank = 0 - scheduler.pp_size = 2 + scheduler.ps = _make_ps(attn_tp_rank=0, dp_rank=0, pp_rank=0, pp_size=2) scheduler.enable_kv_cache_events = False with patch(