config: read parallel config leaves via get_parallel() (#31816)
This commit is contained in:
@@ -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))
|
||||
)
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user