config: read parallel config leaves via get_parallel() (#31816)

This commit is contained in:
Cheng Wan
2026-07-22 01:18:50 -07:00
committed by GitHub
parent 602b546a4e
commit 745b2ca45c
61 changed files with 151 additions and 192 deletions
+2 -2
View File
@@ -1198,7 +1198,7 @@ class Scheduler(
gloo_group=self.attn_tp_cpu_group,
tp_rank=self.ps.tp_rank,
tp_size=self.ps.tp_size,
dp_size=self.server_args.dp_size,
dp_size=get_parallel().dp_size,
gpu_id=self.ps.gpu_id,
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
max_total_num_tokens=self.max_total_num_tokens,
@@ -4313,7 +4313,7 @@ class Scheduler(
old_ep_size = ElasticEPStateManager.get_effective_ep_size()
new_ep_size = recv_req.new_ep_size
max_ep_size = self.server_args.max_ep_size or old_ep_size
max_ep_size = get_parallel().max_ep_size or old_ep_size
logger.debug(
"[Elastic EP][scale] request received: new_ep_size=%d "
@@ -24,7 +24,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
)
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.observability.metrics_collector import DPCooperationInfo
from sglang.srt.runtime_context import get_schedule
from sglang.srt.runtime_context import get_parallel, get_schedule
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils.common import require_mlp_tp_gather
@@ -377,7 +377,7 @@ class SchedulerDPAttnAdapter:
def prepare_mlp_sync_batch(self, local_batch: ScheduleBatch):
return prepare_mlp_sync_batch_raw(
local_batch,
dp_size=self.server_args.dp_size,
dp_size=get_parallel().dp_size,
attn_tp_size=self.ps.attn_tp_size,
attn_cp_size=self.ps.attn_cp_size,
tp_group=self.tp_group,
@@ -16,7 +16,7 @@ from sglang.srt.managers.io_struct import (
sock_recv,
)
from sglang.srt.managers.mm_utils import has_shm_features, unwrap_shm_features
from sglang.srt.runtime_context import get_disagg
from sglang.srt.runtime_context import get_disagg, get_parallel
from sglang.srt.utils import broadcast_pyobj, point_to_point_pyobj
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
@@ -127,7 +127,7 @@ class SchedulerRequestReceiver:
return recv_reqs
def _broadcast_reqs_across_ranks(self, recv_reqs: Optional[List]) -> List:
if self.server_args.enable_dp_attention:
if get_parallel().enable_dp_attention:
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:
@@ -156,7 +156,7 @@ class SchedulerRequestReceiver:
# instead of the full tp_group. This avoids an expensive
# all-ranks gloo sync.
_local_ctrl = (
self.server_args.enable_dp_attention_local_control_broadcast
get_parallel().enable_dp_attention_local_control_broadcast
or self.server_args.is_ep_scale_joiner
)
if _local_ctrl:
@@ -233,7 +233,7 @@ class SchedulerRequestReceiver:
# peer ranks may still be unpickling ShmPointerMMData
# (-> shm_open). Synchronize the same CPU groups that carried
# SHM-backed work requests before materialize() unlinks them.
if self.server_args.enable_dp_attention:
if get_parallel().enable_dp_attention:
if self.ps.attn_tp_size > 1:
barrier(group=self.attn_tp_cpu_group)
if self.ps.attn_cp_size > 1:
@@ -36,7 +36,7 @@ from sglang.srt.model_executor.forward_batch_info import (
PPProxyTensors,
)
from sglang.srt.observability.req_time_stats import set_time_batch
from sglang.srt.runtime_context import get_disagg
from sglang.srt.runtime_context import get_disagg, get_parallel
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.utils import DynamicGradMode, broadcast_pyobj, point_to_point_pyobj
from sglang.srt.utils.common import get_device_module, is_xpu
@@ -123,7 +123,7 @@ class SchedulerPPMixin:
next_pp_outputs = None
next_batch_result = None
d2h_event = None
if self.server_args.pp_async_batch_depth > 0:
if get_parallel().pp_async_batch_depth > 0:
next_pp_outputs, next_batch_result, d2h_event = (
self._pp_commit_send_output_work_and_preprocess_output_tensors(
next_first_rank_mb_id,
@@ -139,7 +139,7 @@ class SchedulerPPMixin:
self.mb_metadata,
self.last_rank_comm_queue,
)
if self.server_args.pp_async_batch_depth == 0:
if get_parallel().pp_async_batch_depth == 0:
next_pp_outputs, next_batch_result, d2h_event = (
self._pp_commit_send_output_work_and_preprocess_output_tensors(
next_first_rank_mb_id,
@@ -269,7 +269,7 @@ class SchedulerPPMixin:
server_is_idle = False
pp_proxy_tensors = self._pp_recv_proxy_tensors()
if self.server_args.pp_async_batch_depth > 0:
if get_parallel().pp_async_batch_depth > 0:
next_pp_outputs, next_batch_result, d2h_event = (
self._pp_commit_send_output_work_and_preprocess_output_tensors(
next_first_rank_mb_id,
@@ -285,7 +285,7 @@ class SchedulerPPMixin:
self.mb_metadata,
self.last_rank_comm_queue,
)
if self.server_args.pp_async_batch_depth == 0:
if get_parallel().pp_async_batch_depth == 0:
next_pp_outputs, next_batch_result, d2h_event = (
self._pp_commit_send_output_work_and_preprocess_output_tensors(
next_first_rank_mb_id,
@@ -428,7 +428,7 @@ class SchedulerPPMixin:
pp_proxy_tensors = self._pp_recv_proxy_tensors()
# early send output if possible
if self.server_args.pp_async_batch_depth > 0:
if get_parallel().pp_async_batch_depth > 0:
next_pp_outputs, next_batch_result, d2h_event = (
self._pp_commit_send_output_work_and_preprocess_output_tensors(
next_first_rank_mb_id,
@@ -446,7 +446,7 @@ class SchedulerPPMixin:
self.last_rank_comm_queue,
)
if self.server_args.pp_async_batch_depth == 0:
if get_parallel().pp_async_batch_depth == 0:
next_pp_outputs, next_batch_result, d2h_event = (
self._pp_commit_send_output_work_and_preprocess_output_tensors(
next_first_rank_mb_id,
@@ -557,10 +557,10 @@ class SchedulerPPMixin:
self.on_idle()
def init_pp_loop_state(self: Scheduler):
self.pp_loop_size: int = self.ps.pp_size + self.server_args.pp_async_batch_depth
self.pp_loop_size: int = self.ps.pp_size + get_parallel().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_dsa_prefill_context_parallel
not get_parallel().enable_dsa_prefill_context_parallel
)
self.mbs = [None] * self.pp_loop_size
self.last_mbs = [None] * self.pp_loop_size
@@ -74,7 +74,7 @@ from sglang.srt.managers.io_struct import (
UpdateWeightsFromTensorReqOutput,
)
from sglang.srt.managers.load_snapshot import LoadSnapshot
from sglang.srt.runtime_context import get_lora
from sglang.srt.runtime_context import get_lora, get_parallel
from sglang.srt.server_args import LoRARef, ServerArgs
from sglang.srt.utils import (
get_bool_env_var,
@@ -146,8 +146,8 @@ class TokenizerControlMixin:
def update_control_communicator_fan_out(self: TokenizerManager, worker_count: int):
primary_group_control = (
self.server_args.enable_dp_attention
and not self.server_args.enable_dp_attention_local_control_broadcast
get_parallel().enable_dp_attention
and not get_parallel().enable_dp_attention_local_control_broadcast
)
if primary_group_control:
control_fan_out = (
@@ -397,7 +397,7 @@ class TokenizerControlMixin:
) -> Tuple[bool, str]:
self.auto_create_handle_loop()
assert (
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention
), "dp_size must be 1 or dp attention must be enabled for update weights from distributed"
results = await self.init_weights_update_group_communicator(obj)
@@ -410,7 +410,7 @@ class TokenizerControlMixin:
) -> Tuple[bool, str]:
self.auto_create_handle_loop()
assert (
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention
), "dp_size must be 1 or dp attention must be enabled for destroy parameter update group"
results = await self.destroy_weights_update_group_communicator(obj)
@@ -423,7 +423,7 @@ class TokenizerControlMixin:
) -> Tuple[bool, str]:
self.auto_create_handle_loop()
assert (
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention
), "dp_size must be 1 or dp attention must be enabled for update weights from distributed"
if obj.abort_all_requests:
@@ -454,7 +454,7 @@ class TokenizerControlMixin:
self.auto_create_handle_loop()
# TODO: support DP
assert (
self.server_args.dp_size == 1
get_parallel().dp_size == 1
), "dp_size must be 1 for init_weights_send_group_for_remote_instance"
result = (
await self.init_weights_send_group_for_remote_instance_communicator(obj)
@@ -469,7 +469,7 @@ class TokenizerControlMixin:
self.auto_create_handle_loop()
# TODO: support DP
assert (
self.server_args.dp_size == 1
get_parallel().dp_size == 1
), "dp_size must be 1 for send_weights_to_remote_instance"
result = (await self.send_weights_to_remote_instance_communicator(obj))[0]
return result.success, result.message
@@ -481,7 +481,7 @@ class TokenizerControlMixin:
) -> Tuple[bool, str]:
self.auto_create_handle_loop()
assert (
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention
), "dp_size must be 1 or dp attention must be enabled for update weights from tensor"
if obj.abort_all_requests:
@@ -517,7 +517,7 @@ class TokenizerControlMixin:
try:
# For now, we only support single data parallel instance
assert (
self.server_args.dp_size == 1 or self.server_args.enable_dp_attention
get_parallel().dp_size == 1 or get_parallel().enable_dp_attention
), "dp_size must be 1 or dp attention must be enabled for update weights from IPC"
logger.info("Starting IPC weight update")
@@ -578,7 +578,7 @@ class TokenizerControlMixin:
# TODO (lifuhuang): Remove this after we verify that dynamic lora loading works
# with dp_size > 1.
assert (
self.server_args.dp_size == 1
get_parallel().dp_size == 1
), "dp_size must be 1 for dynamic lora loading"
logger.info(
"Start load Lora adapter. Lora name=%s, path=%s",
@@ -654,7 +654,7 @@ class TokenizerControlMixin:
)
assert (
self.server_args.dp_size == 1
get_parallel().dp_size == 1
), "dp_size must be 1 for dynamic lora loading"
logger.info(
"Start load Lora adapter from tensors. Lora name=%s",
@@ -730,7 +730,7 @@ class TokenizerControlMixin:
# TODO (lifuhuang): Remove this after we verify that dynamic lora loading works
# with dp_size > 1.
assert (
self.server_args.dp_size == 1
get_parallel().dp_size == 1
), "dp_size must be 1 for dynamic lora loading"
logger.info(
"Start unload Lora adapter. Lora name=%s",
@@ -750,7 +750,7 @@ class TokenizerControlMixin:
self.auto_create_handle_loop()
results = await self.get_weights_by_name_communicator(obj)
all_parameters = [r.parameter for r in results]
if self.server_args.dp_size == 1:
if get_parallel().dp_size == 1:
return all_parameters[0]
else:
return all_parameters
@@ -116,6 +116,7 @@ from sglang.srt.runtime_context import (
get_lora,
get_model,
get_observability,
get_parallel,
get_serving,
)
from sglang.srt.sampling.sampling_params import SamplingParams
@@ -1365,7 +1366,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
return batch_size > 0 and (
self.server_args.enable_tokenizer_batch_encode
or (
(not self.server_args.enable_dp_attention)
(not get_parallel().enable_dp_attention)
and (not self._batch_has_text(batch_size, requests))
)
)