config: the DP/EP topology reads come from the parallel bag (#35025)
This commit is contained in:
@@ -3957,7 +3957,7 @@ def launch_server(server_args: ServerArgs):
|
||||
# Publish before the launch path reads configuration; the encoder built
|
||||
# below re-projects the same object.
|
||||
publish(server_args, role="encoder")
|
||||
if server_args.dp_size > 1:
|
||||
if get_parallel().dp_size > 1:
|
||||
_launch_server_dp(server_args)
|
||||
return
|
||||
|
||||
@@ -4018,12 +4018,12 @@ def launch_server(server_args: ServerArgs):
|
||||
def _launch_server_dp(server_args: ServerArgs):
|
||||
global dp_dispatcher
|
||||
|
||||
if server_args.dp_size <= 1 or server_args.tp_size != 1:
|
||||
if get_parallel().dp_size <= 1 or server_args.tp_size != 1:
|
||||
raise ValueError(
|
||||
"Encoder DP mode requires --dp-size > 1 and --tp-size 1; got "
|
||||
f"dp_size={server_args.dp_size}, tp_size={server_args.tp_size}."
|
||||
f"dp_size={get_parallel().dp_size}, tp_size={server_args.tp_size}."
|
||||
)
|
||||
dp_size = server_args.dp_size
|
||||
dp_size = get_parallel().dp_size
|
||||
logger.info(f"Launching encoder in DP mode: dp_size={dp_size}")
|
||||
|
||||
# DP mode: workers (subprocesses) write metrics to the shared multiproc dir;
|
||||
|
||||
@@ -99,7 +99,7 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa
|
||||
from sglang.srt.parser.template_detection import resolve_auto_parsers
|
||||
from sglang.srt.parser.template_manager import TemplateManager
|
||||
from sglang.srt.plugins import load_plugins
|
||||
from sglang.srt.runtime_context import publish
|
||||
from sglang.srt.runtime_context import get_parallel, publish
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
from sglang.srt.utils import (
|
||||
MultiprocessingSerializer,
|
||||
@@ -336,7 +336,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
routed_dp_rank = data_parallel_rank
|
||||
|
||||
if routed_dp_rank is not None:
|
||||
dp_size = self.server_args.dp_size
|
||||
dp_size = get_parallel().dp_size
|
||||
if dp_size <= 1 and routed_dp_rank == 0:
|
||||
logger.debug(
|
||||
f"routed_dp_rank={routed_dp_rank} is ignored because dp_size={dp_size}"
|
||||
@@ -661,7 +661,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
(``python -m sglang.srt.weight_cache.daemon``) plus
|
||||
``--weight-cache-mode client``, where the daemon outlives the engine.
|
||||
"""
|
||||
if server_args.dp_size > 1:
|
||||
if get_parallel().dp_size > 1:
|
||||
raise ValueError(
|
||||
"Weight cache daemon mode does not support dp_size > 1. "
|
||||
"Please set --dp-size 1 when using --weight-cache-mode daemon."
|
||||
@@ -751,7 +751,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
"--dp-size",
|
||||
"1",
|
||||
"--ep-size",
|
||||
str(server_args.ep_size),
|
||||
str(get_parallel().ep_size),
|
||||
"--load-format",
|
||||
server_args.load_format,
|
||||
"--dtype",
|
||||
@@ -863,7 +863,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
"""
|
||||
scheduler_procs = []
|
||||
use_dp_controller = (
|
||||
server_args.dp_size > 1 or server_args.ep_join_mode == "scale"
|
||||
get_parallel().dp_size > 1 or server_args.ep_join_mode == "scale"
|
||||
)
|
||||
|
||||
if not use_dp_controller:
|
||||
@@ -1835,7 +1835,7 @@ def _compute_parallelism_ranks(
|
||||
server_args: ServerArgs, tp_rank: int
|
||||
) -> Tuple[int, int, int]:
|
||||
"""Compute attention-CP, MoE-DP, and MoE-EP ranks for a TP rank."""
|
||||
attn_dp_size = server_args.dp_size if server_args.enable_dp_attention else 1
|
||||
attn_dp_size = get_parallel().dp_size if get_parallel().enable_dp_attention else 1
|
||||
|
||||
# Parallelism hierarchy (outermost to innermost):
|
||||
# - Attention: Global(TP) -> DP -> ATTN_CP -> ATTN_TP (innermost)
|
||||
@@ -1846,6 +1846,6 @@ def _compute_parallelism_ranks(
|
||||
moe_ep_rank = (
|
||||
tp_rank
|
||||
% (server_args.tp_size // server_args.moe_dp_size)
|
||||
// (server_args.tp_size // server_args.moe_dp_size // server_args.ep_size)
|
||||
// (server_args.tp_size // server_args.moe_dp_size // get_parallel().ep_size)
|
||||
)
|
||||
return attn_cp_rank, moe_dp_rank, moe_ep_rank
|
||||
|
||||
@@ -480,6 +480,7 @@ v1_loads_router.route_class = ORJSONRoute
|
||||
app.include_router(v1_loads_router)
|
||||
|
||||
from sglang.srt.entrypoints.elastic_ep import router as elastic_ep_router
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
elastic_ep_router.route_class = ORJSONRoute
|
||||
app.include_router(elastic_ep_router)
|
||||
@@ -2156,7 +2157,10 @@ async def _send_disaggregation_warmup_requests(
|
||||
headers=headers,
|
||||
) as session:
|
||||
return await asyncio.gather(
|
||||
*(send_request(session, dp_rank) for dp_rank in range(server_args.dp_size))
|
||||
*(
|
||||
send_request(session, dp_rank)
|
||||
for dp_rank in range(get_parallel().dp_size)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -2214,9 +2218,9 @@ def _execute_server_warmup(server_args: ServerArgs):
|
||||
},
|
||||
}
|
||||
if server_args.skip_tokenizer_init:
|
||||
json_data["input_ids"] = [[10, 11, 12] for _ in range(server_args.dp_size)]
|
||||
json_data["input_ids"] = [[10, 11, 12] for _ in range(get_parallel().dp_size)]
|
||||
# TODO Workaround the bug that embedding errors for list of size 1
|
||||
if server_args.dp_size == 1:
|
||||
if get_parallel().dp_size == 1:
|
||||
json_data["input_ids"] = json_data["input_ids"][0]
|
||||
elif (
|
||||
is_vlm
|
||||
@@ -2260,9 +2264,9 @@ def _execute_server_warmup(server_args: ServerArgs):
|
||||
"temperature": 0.0,
|
||||
}
|
||||
else:
|
||||
json_data["text"] = ["The capital city of France is"] * server_args.dp_size
|
||||
json_data["text"] = ["The capital city of France is"] * get_parallel().dp_size
|
||||
# TODO Workaround the bug that embedding errors for list of size 1
|
||||
if server_args.dp_size == 1:
|
||||
if get_parallel().dp_size == 1:
|
||||
json_data["text"] = json_data["text"][0]
|
||||
|
||||
# Config debug dumping
|
||||
@@ -2304,7 +2308,7 @@ def _execute_server_warmup(server_args: ServerArgs):
|
||||
if not failed_status_codes:
|
||||
logger.info(
|
||||
"Disaggregation warmup requests completed for all %s DP ranks",
|
||||
server_args.dp_size,
|
||||
get_parallel().dp_size,
|
||||
)
|
||||
logger.info("End of disaggregation warmup")
|
||||
else:
|
||||
|
||||
@@ -31,7 +31,10 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
from sglang.srt.runtime_context import (
|
||||
configured_attn_cp_size,
|
||||
configured_moe_dp_size,
|
||||
get_device,
|
||||
get_exec,
|
||||
get_flags,
|
||||
get_parallel,
|
||||
)
|
||||
from sglang.srt.utils import get_bool_env_var, is_hip
|
||||
|
||||
@@ -345,9 +348,9 @@ def initialize_dp_attention(
|
||||
dp.max_len_with_idle = (
|
||||
getattr(model_config.hf_config, "hybrid_override_pattern", None) is not None
|
||||
)
|
||||
enable_dp_attention = server_args.enable_dp_attention
|
||||
dp_size = server_args.dp_size
|
||||
attn_cp_size = server_args.attn_cp_size
|
||||
enable_dp_attention = get_parallel().enable_dp_attention
|
||||
dp_size = get_parallel().dp_size
|
||||
attn_cp_size = configured_attn_cp_size()
|
||||
|
||||
dp.enabled = enable_dp_attention
|
||||
|
||||
@@ -359,15 +362,15 @@ def initialize_dp_attention(
|
||||
)
|
||||
_ATTN_DP_SIZE = dp_size if enable_dp_attention else 1
|
||||
|
||||
if server_args.elastic_ep_backend is not None and server_args.max_ep_size:
|
||||
_ATTN_DP_RANK = tp_rank + server_args.ep_join_rank_offset
|
||||
if get_exec().moe.elastic_ep_backend is not None and get_parallel().max_ep_size:
|
||||
_ATTN_DP_RANK = tp_rank + get_parallel().ep_join_rank_offset
|
||||
if server_args.is_ep_scale_joiner:
|
||||
dp.joiner_skip_all_gather = True
|
||||
|
||||
_DpGatheredBufferWrapper.set_metadata(
|
||||
hidden_size=model_config.hidden_size,
|
||||
dtype=model_config.dtype,
|
||||
device=torch.device(server_args.device),
|
||||
device=torch.device(get_device().device),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -46,7 +46,11 @@ from sglang.srt.lora.utils import (
|
||||
)
|
||||
from sglang.srt.managers.io_struct import LoRAUpdateOutput
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_lora,
|
||||
get_parallel,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import get_available_gpu_memory, replace_submodule
|
||||
from sglang.srt.utils.hf_transformers_utils import AutoConfig
|
||||
@@ -93,7 +97,7 @@ class LoRAManager:
|
||||
self.pending_lora_load_events = {}
|
||||
|
||||
self.eviction_policy = server_args.lora_eviction_policy
|
||||
self.enable_dp_attention: bool = server_args.enable_dp_attention
|
||||
self.enable_dp_attention: bool = get_parallel().enable_dp_attention
|
||||
self._experts_shared_outer_override: Optional[bool] = (
|
||||
server_args.experts_shared_outer_loras
|
||||
)
|
||||
@@ -1031,7 +1035,7 @@ def init_lora_cuda_graph_moe_buffers(
|
||||
from sglang.srt.lora.layers import FusedMoEWithLoRA
|
||||
|
||||
max_bs = get_exec().graph.cuda_graph_config.decode.max_bs
|
||||
max_loras = server_args.max_loras_per_batch
|
||||
max_loras = get_lora().max_loras_per_batch
|
||||
for module in model.modules():
|
||||
if isinstance(module, FusedMoEWithLoRA):
|
||||
lora_manager.init_cuda_graph_moe_buffers(max_bs, max_loras, dtype, module)
|
||||
|
||||
@@ -49,7 +49,7 @@ from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
|
||||
from sglang.srt.observability.req_time_stats import DPControllerReqTimeStats
|
||||
from sglang.srt.observability.startup_time import aggregate_scheduler_startup_times
|
||||
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
|
||||
from sglang.srt.runtime_context import get_exec, publish
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, publish
|
||||
from sglang.srt.server_args import (
|
||||
DP_ATTENTION_HANDSHAKE_PORT_DELTA,
|
||||
PortArgs,
|
||||
@@ -147,7 +147,7 @@ class DataParallelController:
|
||||
self.run_scheduler_process_func = run_scheduler_process_func
|
||||
|
||||
# Init inter-process communication
|
||||
self.context = zmq.Context(1 + server_args.dp_size)
|
||||
self.context = zmq.Context(1 + get_parallel().dp_size)
|
||||
if server_args.node_rank == 0:
|
||||
self.recv_from_tokenizer = get_zmq_socket(
|
||||
self.context, zmq.PULL, port_args.scheduler_input_ipc_name, False
|
||||
@@ -167,8 +167,8 @@ class DataParallelController:
|
||||
LoadBalanceMethod.TOTAL_TOKENS,
|
||||
)
|
||||
|
||||
self.launch_dp_size: int = server_args.dp_size
|
||||
self.max_dp_size: int = server_args.max_ep_size or server_args.dp_size
|
||||
self.launch_dp_size: int = get_parallel().dp_size
|
||||
self.max_dp_size: int = server_args.max_ep_size or get_parallel().dp_size
|
||||
assert self.max_dp_size >= self.launch_dp_size, (
|
||||
f"--max-ep-size ({self.max_dp_size}) must be >= "
|
||||
f"--dp ({self.launch_dp_size})."
|
||||
@@ -178,9 +178,8 @@ class DataParallelController:
|
||||
self.max_dp_size - self.launch_dp_size
|
||||
)
|
||||
|
||||
self.dp_budget = DPBudget(server_args.dp_size)
|
||||
self.dp_budget = DPBudget(get_parallel().dp_size)
|
||||
self.load_snapshot_reader = create_load_snapshot_reader(
|
||||
server_args,
|
||||
port_args,
|
||||
caller="DataParallelController",
|
||||
)
|
||||
@@ -196,14 +195,14 @@ class DataParallelController:
|
||||
self._active_workers: List[int] = list(range(self.launch_dp_size))
|
||||
self._active_count_cache: int = self.launch_dp_size
|
||||
|
||||
if server_args.enable_dp_attention:
|
||||
if get_parallel().enable_dp_attention:
|
||||
self.launch_dp_attention_schedulers(server_args, port_args)
|
||||
# When local control broadcast is enabled, send control messages to
|
||||
# every DP group leader (attn_tp_rank=0) so each leader broadcasts
|
||||
# within its own attn_tp_group instead of the full tp_group.
|
||||
# Otherwise fall back to the original behaviour: send to only the
|
||||
# first leader, which then broadcasts over the full tp_group.
|
||||
local_ctrl = server_args.enable_dp_attention_local_control_broadcast
|
||||
local_ctrl = get_parallel().enable_dp_attention_local_control_broadcast
|
||||
self.control_message_step = 1 if local_ctrl else server_args.tp_size
|
||||
else:
|
||||
self.launch_dp_schedulers(server_args, port_args)
|
||||
@@ -367,7 +366,7 @@ class DataParallelController:
|
||||
threads = []
|
||||
sockets = []
|
||||
ready_events = []
|
||||
for dp_rank in range(server_args.dp_size):
|
||||
for dp_rank in range(get_parallel().dp_size):
|
||||
tmp_port_args = PortArgs.init_new(server_args)
|
||||
tmp_port_args.tokenizer_ipc_name = port_args.tokenizer_ipc_name
|
||||
tmp_port_args.detokenizer_ipc_name = port_args.detokenizer_ipc_name
|
||||
@@ -569,7 +568,7 @@ class DataParallelController:
|
||||
bind_count = (
|
||||
self.max_dp_size
|
||||
if server_args.elastic_ep_backend is not None
|
||||
else server_args.dp_size
|
||||
else get_parallel().dp_size
|
||||
)
|
||||
for slot in range(bind_count):
|
||||
worker_port, worker_socket = get_zmq_socket_on_host(
|
||||
@@ -599,7 +598,7 @@ class DataParallelController:
|
||||
dp_rank: Optional[int],
|
||||
worker_ports: Optional[List[int]] = None,
|
||||
):
|
||||
if not server_args.enable_dp_attention:
|
||||
if not get_parallel().enable_dp_attention:
|
||||
logger.info(f"Launch DP{dp_rank} starting at GPU #{base_gpu_id}.")
|
||||
|
||||
memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||||
@@ -633,13 +632,13 @@ class DataParallelController:
|
||||
for tp_rank in tp_rank_range:
|
||||
rank_port_args = port_args
|
||||
|
||||
if server_args.enable_dp_attention:
|
||||
if get_parallel().enable_dp_attention:
|
||||
# dp attention has different sharding logic
|
||||
_, _, dp_rank, _ = compute_dp_attention_world_info(
|
||||
server_args.enable_dp_attention,
|
||||
get_parallel().enable_dp_attention,
|
||||
tp_rank,
|
||||
server_args.tp_size,
|
||||
server_args.dp_size,
|
||||
get_parallel().dp_size,
|
||||
server_args.attn_cp_size,
|
||||
)
|
||||
# compute zmq ports for this dp rank
|
||||
@@ -669,7 +668,7 @@ class DataParallelController:
|
||||
+ (tp_rank % tp_size_per_node) * server_args.gpu_id_step
|
||||
)
|
||||
attn_dp_size = (
|
||||
server_args.dp_size if server_args.enable_dp_attention else 1
|
||||
get_parallel().dp_size if get_parallel().enable_dp_attention else 1
|
||||
)
|
||||
|
||||
# Parallelism hierarchy (outermost to innermost):
|
||||
@@ -688,7 +687,7 @@ class DataParallelController:
|
||||
// (
|
||||
server_args.tp_size
|
||||
// server_args.moe_dp_size
|
||||
// server_args.ep_size
|
||||
// get_parallel().ep_size
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -52,6 +52,7 @@ import msgspec.msgpack
|
||||
import msgspec.structs
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.runtime_context import get_parallel, get_serving
|
||||
from sglang.srt.utils.network import is_zmq_endpoint_ipv6
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -61,7 +62,7 @@ logger = logging.getLogger(__name__)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def should_use_zmq(server_args) -> bool:
|
||||
def should_use_zmq() -> bool:
|
||||
"""Whether to use zmq PUSH/PULL instead of shared memory for load snapshots.
|
||||
|
||||
Shared memory (mmap) only works within a single node. When schedulers
|
||||
@@ -70,14 +71,14 @@ def should_use_zmq(server_args) -> bool:
|
||||
``SGLANG_LOAD_SNAPSHOT_USE_ZMQ`` forces zmq mode for testing.
|
||||
"""
|
||||
return (
|
||||
server_args.enable_dp_attention and server_args.nnodes > 1
|
||||
get_parallel().enable_dp_attention and get_parallel().nnodes > 1
|
||||
) or envs.SGLANG_LOAD_SNAPSHOT_USE_ZMQ.get()
|
||||
|
||||
|
||||
_LOAD_AWARE_METHODS = frozenset({"total_requests", "total_tokens"})
|
||||
|
||||
|
||||
def _tokenizer_load_snapshot_owner_caller(server_args) -> str:
|
||||
def _tokenizer_load_snapshot_owner_caller() -> str:
|
||||
"""The caller that plays the tokenizer-side zmq owner role.
|
||||
|
||||
In multi-tokenizer mode (``tokenizer_worker_num > 1``) there are N
|
||||
@@ -85,12 +86,12 @@ def _tokenizer_load_snapshot_owner_caller(server_args) -> str:
|
||||
same zmq PULL endpoint. Instead, the single ``MultiTokenizerRouter``
|
||||
process owns the socket (polls zmq -> SHM) and every worker reads SHM.
|
||||
"""
|
||||
if server_args.tokenizer_worker_num > 1:
|
||||
if get_serving().tokenizer_worker_num > 1:
|
||||
return "MultiTokenizerRouter"
|
||||
return "TokenizerManager"
|
||||
|
||||
|
||||
def zmq_reader_owner(server_args, caller: str) -> bool:
|
||||
def zmq_reader_owner(caller: str) -> bool:
|
||||
"""Decide which process owns the zmq PULL socket.
|
||||
|
||||
Exactly one of ``"DataParallelController"``, ``"TokenizerManager"``, or
|
||||
@@ -108,18 +109,25 @@ def zmq_reader_owner(server_args, caller: str) -> bool:
|
||||
load data -> tokenizer-side owner owns it (polls on /v1/loads calls).
|
||||
|
||||
The tokenizer-side owner is the ``"MultiTokenizerRouter"`` caller in
|
||||
multi-tokenizer mode, otherwise the ``"TokenizerManager"`` caller.
|
||||
multi-tokenizer mode, otherwise the ``"TokenizerManager"`` caller. Which of
|
||||
the two it is only matters to a tokenizer-side caller, and asking costs the
|
||||
DP controller a `serving` read its role is not audited for -- so the
|
||||
controller answers from the parallel leaves alone.
|
||||
"""
|
||||
if not should_use_zmq(server_args):
|
||||
if not should_use_zmq():
|
||||
return False
|
||||
if server_args.node_rank != 0:
|
||||
if get_parallel().node_rank != 0:
|
||||
return False
|
||||
tokenizer_owner = _tokenizer_load_snapshot_owner_caller(server_args)
|
||||
if server_args.dp_size == 1:
|
||||
return caller == tokenizer_owner
|
||||
if server_args.load_balance_method.lower() in _LOAD_AWARE_METHODS:
|
||||
return caller == "DataParallelController"
|
||||
return caller == tokenizer_owner
|
||||
if caller == "DataParallelController":
|
||||
return (
|
||||
get_parallel().dp_size > 1
|
||||
and get_parallel().load_balance_method.lower() in _LOAD_AWARE_METHODS
|
||||
)
|
||||
if get_parallel().dp_size > 1 and (
|
||||
get_parallel().load_balance_method.lower() in _LOAD_AWARE_METHODS
|
||||
):
|
||||
return False
|
||||
return caller == _tokenizer_load_snapshot_owner_caller()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -627,14 +635,13 @@ def _zmq_addr_for(port_args) -> str:
|
||||
|
||||
|
||||
def create_load_snapshot_writer(
|
||||
server_args,
|
||||
port_args,
|
||||
dp_size: int,
|
||||
dp_rank: int,
|
||||
publish_interval: int = 1,
|
||||
):
|
||||
"""Return a SHM or ZMQ writer based on server configuration."""
|
||||
if should_use_zmq(server_args):
|
||||
if should_use_zmq():
|
||||
return ZmqLoadSnapshotWriter(
|
||||
_zmq_addr_for(port_args), dp_size, dp_rank, publish_interval
|
||||
)
|
||||
@@ -643,7 +650,7 @@ def create_load_snapshot_writer(
|
||||
)
|
||||
|
||||
|
||||
def create_load_snapshot_reader(server_args, port_args, caller: str):
|
||||
def create_load_snapshot_reader(port_args, caller: str):
|
||||
"""Create a load snapshot reader.
|
||||
|
||||
Args:
|
||||
@@ -651,8 +658,8 @@ def create_load_snapshot_reader(server_args, port_args, caller: str):
|
||||
``"MultiTokenizerRouter"`` -- determines who binds the zmq PULL
|
||||
socket when zmq mode is active.
|
||||
"""
|
||||
dp_size = server_args.dp_size
|
||||
if zmq_reader_owner(server_args, caller):
|
||||
dp_size = get_parallel().dp_size
|
||||
if zmq_reader_owner(caller):
|
||||
return ZmqShmLoadSnapshotReader(
|
||||
_zmq_addr_for(port_args), shm_path_for(port_args.instance_id), dp_size
|
||||
)
|
||||
|
||||
@@ -467,9 +467,9 @@ class MultiTokenizerRouter:
|
||||
# read SHM only. Drain it event-driven via the socket's fd instead of
|
||||
# polling on a timer.
|
||||
self.load_snapshot_reader = None
|
||||
if zmq_reader_owner(server_args, "MultiTokenizerRouter"):
|
||||
if zmq_reader_owner("MultiTokenizerRouter"):
|
||||
self.load_snapshot_reader = create_load_snapshot_reader(
|
||||
server_args, port_args, caller="MultiTokenizerRouter"
|
||||
port_args, caller="MultiTokenizerRouter"
|
||||
)
|
||||
self._loop.call_soon_threadsafe(self._register_load_snapshot_reader)
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, NamedTuple, Optional
|
||||
import torch
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.utils import get_bool_env_var
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -103,7 +104,7 @@ class PrefillDelayer:
|
||||
f"queue_trigger_enabled={self._queue_trigger_enabled}"
|
||||
)
|
||||
self.dp_size = dp_size
|
||||
self.enable_dp_attention = server_args.enable_dp_attention
|
||||
self.enable_dp_attention = get_parallel().enable_dp_attention
|
||||
dp_size_dim = dp_size if self.enable_dp_attention else 1
|
||||
|
||||
# Mirror scheduler_dp_attn_mixin's NCCL all-gather path: when the
|
||||
|
||||
@@ -29,7 +29,10 @@ from typing import TYPE_CHECKING, Any, Deque, Dict, List, Optional, Set, Tuple,
|
||||
|
||||
from sglang.srt.runtime_context import (
|
||||
attention_backends,
|
||||
configured_attn_cp_size,
|
||||
configured_moe_dp_size,
|
||||
configured_pp_size,
|
||||
configured_tp_size,
|
||||
get_device,
|
||||
get_disagg,
|
||||
get_exec,
|
||||
@@ -285,12 +288,7 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa
|
||||
from sglang.srt.parser.reasoning_parser import ReasoningParser
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.plugins import load_plugins
|
||||
from sglang.srt.runtime_context import (
|
||||
get_context,
|
||||
get_device,
|
||||
get_parallel,
|
||||
publish,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_context, publish
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
from sglang.srt.sampling.sampling_params import TOP_K_ALL
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
@@ -451,16 +449,16 @@ class Scheduler(
|
||||
self.max_recv_per_poll = envs.SGLANG_SCHEDULER_MAX_RECV_PER_POLL.get()
|
||||
self.max_new_tokens_limit = envs.SGLANG_MAX_NEW_TOKENS_LIMIT.get()
|
||||
self.enable_hisparse = server_args.enable_hisparse
|
||||
self.enable_dp_attention = server_args.enable_dp_attention
|
||||
self.enable_dp_attention = get_parallel().enable_dp_attention
|
||||
self.enable_unified_memory = server_args.enable_unified_memory
|
||||
|
||||
# Distributed rank info
|
||||
attn_tp_rank, attn_tp_size, attn_dp_rank, attn_dp_size = (
|
||||
compute_dp_attention_world_info(
|
||||
server_args.enable_dp_attention,
|
||||
get_parallel().enable_dp_attention,
|
||||
tp_rank,
|
||||
server_args.tp_size,
|
||||
server_args.dp_size,
|
||||
get_parallel().dp_size,
|
||||
server_args.attn_cp_size,
|
||||
)
|
||||
)
|
||||
@@ -470,7 +468,7 @@ class Scheduler(
|
||||
pp_rank=pp_rank,
|
||||
pp_size=server_args.pp_size,
|
||||
dp_rank=dp_rank,
|
||||
dp_size=server_args.dp_size,
|
||||
dp_size=get_parallel().dp_size,
|
||||
attn_tp_rank=attn_tp_rank,
|
||||
attn_tp_size=attn_tp_size,
|
||||
attn_cp_rank=attn_cp_rank,
|
||||
@@ -480,7 +478,7 @@ class Scheduler(
|
||||
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_ep_size=get_parallel().ep_size,
|
||||
moe_dp_rank=moe_dp_rank,
|
||||
moe_dp_size=server_args.moe_dp_size,
|
||||
gpu_id=gpu_id,
|
||||
@@ -768,7 +766,6 @@ class Scheduler(
|
||||
dp_rank = self.ps.dp_rank if self.ps.dp_rank is not None else 0
|
||||
try:
|
||||
self.load_snapshot_writer = create_load_snapshot_writer(
|
||||
self.server_args,
|
||||
port_args,
|
||||
self.ps.dp_size,
|
||||
dp_rank,
|
||||
@@ -1277,7 +1274,7 @@ class Scheduler(
|
||||
)
|
||||
|
||||
# Init recv skipper and input blocker
|
||||
self.recv_skipper = SchedulerRecvSkipper.maybe_create(self.server_args)
|
||||
self.recv_skipper = SchedulerRecvSkipper.maybe_create()
|
||||
self.input_blocker = (
|
||||
SchedulerInputBlocker(noop=self.ps.attn_tp_rank != 0)
|
||||
if get_bool_env_var("SGLANG_ENABLE_COLOCATED_BATCH_GEN")
|
||||
@@ -4967,15 +4964,15 @@ def configure_scheduler_process(
|
||||
prefix = ""
|
||||
if shown_dp is not None:
|
||||
prefix += f" DP{shown_dp}"
|
||||
if server_args.pp_size > 1:
|
||||
if configured_pp_size() > 1:
|
||||
prefix += f" PP{pp_rank}"
|
||||
if server_args.attn_cp_size > 1:
|
||||
if configured_attn_cp_size() > 1:
|
||||
prefix += f" ATTN_CP{attn_cp_rank}"
|
||||
if server_args.moe_dp_size > 1:
|
||||
if configured_moe_dp_size() > 1:
|
||||
prefix += f" MOE_DP{moe_dp_rank}"
|
||||
if server_args.tp_size > 1:
|
||||
if configured_tp_size() > 1:
|
||||
prefix += f" TP{shown_tp}"
|
||||
if server_args.ep_size > 1:
|
||||
if get_parallel().ep_size > 1:
|
||||
prefix += f" EP{shown_moe_ep}"
|
||||
|
||||
# Config the process
|
||||
@@ -4989,7 +4986,7 @@ def configure_scheduler_process(
|
||||
# Set cpu affinity to this gpu process
|
||||
if envs.SGLANG_SET_CPU_AFFINITY.get():
|
||||
set_gpu_proc_affinity(
|
||||
server_args.pp_size, server_args.tp_size, server_args.nnodes, gpu_id
|
||||
configured_pp_size(), configured_tp_size(), get_parallel().nnodes, gpu_id
|
||||
)
|
||||
if not envs.SGLANG_NUMA_BIND_V2.get():
|
||||
numa_node = get_numa_node_if_available(server_args, gpu_id)
|
||||
|
||||
@@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, List, Optional
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.runtime_context import get_parallel, get_schedule
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||
@@ -12,10 +12,10 @@ if TYPE_CHECKING:
|
||||
|
||||
class SchedulerRecvSkipper:
|
||||
@staticmethod
|
||||
def maybe_create(server_args: ServerArgs):
|
||||
if server_args.scheduler_recv_interval <= 1:
|
||||
def maybe_create():
|
||||
if get_schedule().scheduler_recv_interval <= 1:
|
||||
return None
|
||||
return SchedulerRecvSkipper(server_args)
|
||||
return SchedulerRecvSkipper()
|
||||
|
||||
@staticmethod
|
||||
def derive_forward_mode(gathered_modes: List[int]) -> Optional[ForwardMode]:
|
||||
@@ -33,10 +33,10 @@ class SchedulerRecvSkipper:
|
||||
return ForwardMode.TARGET_VERIFY
|
||||
return ForwardMode.DECODE
|
||||
|
||||
def __init__(self, server_args: ServerArgs):
|
||||
self._use_synced_mode = server_args.enable_dp_attention
|
||||
def __init__(self):
|
||||
self._use_synced_mode = get_parallel().enable_dp_attention
|
||||
self._counter = 0
|
||||
self._threshold = server_args.scheduler_recv_interval
|
||||
self._threshold = get_schedule().scheduler_recv_interval
|
||||
# All can be tuned if needed
|
||||
self._default_weight = envs.SGLANG_SCHEDULER_RECV_SKIPPER_WEIGHT_DEFAULT.get()
|
||||
self._weight_of_forward_mode = {
|
||||
|
||||
@@ -74,6 +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_parallel
|
||||
from sglang.srt.server_args import LoRARef, ServerArgs
|
||||
from sglang.srt.utils import (
|
||||
get_bool_env_var,
|
||||
@@ -157,7 +158,7 @@ class TokenizerControlMixin:
|
||||
mode = spec[2] if len(spec) > 2 else "queueing"
|
||||
comm = FanOutCommunicator(
|
||||
self._dispatch_to_scheduler,
|
||||
server_args.dp_size,
|
||||
get_parallel().dp_size,
|
||||
mode,
|
||||
)
|
||||
setattr(self, f"{name}_communicator", comm)
|
||||
@@ -166,8 +167,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 = (
|
||||
@@ -420,7 +421,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)
|
||||
@@ -433,7 +434,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)
|
||||
@@ -446,7 +447,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:
|
||||
@@ -479,7 +480,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)
|
||||
@@ -494,7 +495,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
|
||||
@@ -506,7 +507,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:
|
||||
@@ -544,7 +545,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")
|
||||
|
||||
@@ -607,7 +608,7 @@ class TokenizerControlMixin:
|
||||
)
|
||||
|
||||
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 dynamic lora loading"
|
||||
logger.info(
|
||||
"Start load Lora adapter. Lora name=%s, path=%s",
|
||||
@@ -685,7 +686,7 @@ class TokenizerControlMixin:
|
||||
)
|
||||
|
||||
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 dynamic lora loading"
|
||||
logger.info(
|
||||
"Start load Lora adapter from tensors. Lora name=%s",
|
||||
@@ -765,7 +766,7 @@ class TokenizerControlMixin:
|
||||
), "lora_name must be provided to unload LoRA adapter"
|
||||
|
||||
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 dynamic lora loading"
|
||||
logger.info(
|
||||
"Start unload Lora adapter. Lora name=%s",
|
||||
@@ -785,7 +786,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
|
||||
|
||||
@@ -121,6 +121,7 @@ from sglang.srt.observability.request_metrics_exporter import (
|
||||
RequestMetricsExporterManager,
|
||||
)
|
||||
from sglang.srt.observability.trace import SpanAttributes, extract_trace_headers
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||
from sglang.srt.server_args import (
|
||||
PortArgs,
|
||||
@@ -398,7 +399,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
set_global_server_args_for_tokenizer(server_args)
|
||||
self.startup_time: Optional[Dict[str, Any]] = None
|
||||
self._config_updates: List[Tuple[str, Dict[str, Any]]] = []
|
||||
self.elastic_worker_count = server_args.dp_size
|
||||
self.elastic_worker_count = get_parallel().dp_size
|
||||
self.elastic_pending_ep_size = None
|
||||
self.elastic_scale_phase = "idle"
|
||||
self.elastic_last_error = None
|
||||
@@ -551,7 +552,6 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
self.tokenizer_ipc_name = port_args.tokenizer_ipc_name
|
||||
|
||||
self.load_snapshot_reader = create_load_snapshot_reader(
|
||||
self.server_args,
|
||||
port_args,
|
||||
caller="TokenizerManager",
|
||||
)
|
||||
@@ -1583,7 +1583,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))
|
||||
)
|
||||
)
|
||||
|
||||
@@ -240,10 +240,10 @@ def build_kv_cache(
|
||||
|
||||
retraction_backup = resolve_decode_retraction_backup(tp_worker=tp_worker)
|
||||
|
||||
disable_radix_cache = server_args.disable_radix_cache or (
|
||||
disable_radix_cache = get_memory().disable_radix_cache or (
|
||||
model_config.is_multimodal and uses_transformers_backend
|
||||
)
|
||||
if disable_radix_cache and not server_args.disable_radix_cache:
|
||||
if disable_radix_cache and not get_memory().disable_radix_cache:
|
||||
logger.warning(
|
||||
"Radix cache is disabled for multimodal models with the "
|
||||
"Transformers backend to avoid multimodal prefix-cache mismatches."
|
||||
@@ -253,8 +253,8 @@ def build_kv_cache(
|
||||
# these use specialized memory pools incompatible with the
|
||||
# prefix-match-and-lock allocation path.
|
||||
if (
|
||||
server_args.disaggregation_decode_enable_radix_cache
|
||||
and server_args.disaggregation_mode == "decode"
|
||||
get_disagg().disaggregation_decode_enable_radix_cache
|
||||
and get_disagg().disaggregation_mode == "decode"
|
||||
):
|
||||
if is_hybrid_swa:
|
||||
raise ValueError(
|
||||
@@ -285,15 +285,15 @@ def build_kv_cache(
|
||||
),
|
||||
is_eagle=spec_algorithm.is_eagle(),
|
||||
tp_cache_group=(
|
||||
attn_tp_cpu_group if server_args.enable_dp_attention else tp_cpu_group
|
||||
attn_tp_cpu_group if get_parallel().enable_dp_attention else tp_cpu_group
|
||||
),
|
||||
attn_cp_cache_group=attn_cp_cpu_group,
|
||||
attn_tp_cache_group=attn_tp_cpu_group,
|
||||
pp_cache_group=pp_group.cpu_group,
|
||||
eviction_policy=server_args.radix_eviction_policy,
|
||||
eviction_policy=get_memory().radix_eviction_policy,
|
||||
enable_metrics=enable_metrics,
|
||||
enable_kv_cache_events=enable_kv_cache_events,
|
||||
enable_session_radix_cache=server_args.enable_session_radix_cache,
|
||||
enable_session_radix_cache=get_memory().enable_session_radix_cache,
|
||||
enable_mamba_extra_buffer=server_args.enable_mamba_extra_buffer(),
|
||||
enable_mamba_extra_buffer_lazy=server_args.enable_mamba_extra_buffer_lazy(),
|
||||
pp_rank=ps.pp_rank,
|
||||
|
||||
@@ -20,6 +20,7 @@ from sglang.srt.model_loader.weight_utils import (
|
||||
CheckpointFilePrefetchHandle,
|
||||
)
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
@@ -112,8 +113,8 @@ class StartupWeightLoadOptions:
|
||||
attn_cp_size=server_args.attn_cp_size,
|
||||
dcp_size=server_args.dcp_size,
|
||||
pp_size=server_args.pp_size,
|
||||
dp_size=server_args.dp_size,
|
||||
ep_size=server_args.ep_size,
|
||||
dp_size=get_parallel().dp_size,
|
||||
ep_size=get_parallel().ep_size,
|
||||
cpu_offload_gb=server_args.cpu_offload_gb,
|
||||
offload_group_size=server_args.offload_group_size,
|
||||
enable_memory_saver=server_args.enable_memory_saver,
|
||||
|
||||
@@ -31,6 +31,7 @@ from sglang.srt.ray.engine import (
|
||||
_get_bundle_node_ip,
|
||||
_resolve_bundle_indices,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
from sglang.srt.utils.network import bind_port, get_zmq_socket, get_zmq_socket_on_host
|
||||
|
||||
@@ -74,7 +75,7 @@ class RayDataParallelController(DataParallelController):
|
||||
sockets = []
|
||||
dp_port_args_list = []
|
||||
|
||||
for dp_rank in range(server_args.dp_size):
|
||||
for dp_rank in range(get_parallel().dp_size):
|
||||
tmp_port_args = PortArgs.init_new(server_args)
|
||||
tmp_port_args.tokenizer_ipc_name = port_args.tokenizer_ipc_name
|
||||
tmp_port_args.detokenizer_ipc_name = port_args.detokenizer_ipc_name
|
||||
@@ -98,7 +99,7 @@ class RayDataParallelController(DataParallelController):
|
||||
sock.close()
|
||||
|
||||
# Create actors for each DP rank sequentially
|
||||
for dp_rank in range(server_args.dp_size):
|
||||
for dp_rank in range(get_parallel().dp_size):
|
||||
self._launch_ray_tp_group(server_args, dp_port_args_list[dp_rank], dp_rank)
|
||||
|
||||
def launch_dp_attention_schedulers(
|
||||
@@ -109,7 +110,7 @@ class RayDataParallelController(DataParallelController):
|
||||
# rank-0 node IP instead of tcp://* to avoid exposing unauthenticated
|
||||
# ZMQ sockets (CVE-2026-3060).
|
||||
worker_ports = []
|
||||
for dp_rank in range(server_args.dp_size):
|
||||
for dp_rank in range(get_parallel().dp_size):
|
||||
worker_port, worker_socket = get_zmq_socket_on_host(
|
||||
self.context, zmq.PUSH, host=self.rank0_node_ip
|
||||
)
|
||||
@@ -154,12 +155,12 @@ class RayDataParallelController(DataParallelController):
|
||||
tp_rank % tp_per_node
|
||||
)
|
||||
|
||||
if server_args.enable_dp_attention:
|
||||
if get_parallel().enable_dp_attention:
|
||||
_, _, actual_dp_rank, _ = compute_dp_attention_world_info(
|
||||
server_args.enable_dp_attention,
|
||||
get_parallel().enable_dp_attention,
|
||||
tp_rank,
|
||||
server_args.tp_size,
|
||||
server_args.dp_size,
|
||||
get_parallel().dp_size,
|
||||
server_args.attn_cp_size,
|
||||
)
|
||||
rank_port_args = PortArgs.init_new(
|
||||
@@ -225,12 +226,12 @@ class RayDataParallelController(DataParallelController):
|
||||
|
||||
bundle_idx = bundle_indices[global_rank]
|
||||
|
||||
if server_args.enable_dp_attention:
|
||||
if get_parallel().enable_dp_attention:
|
||||
_, _, actual_dp_rank, _ = compute_dp_attention_world_info(
|
||||
server_args.enable_dp_attention,
|
||||
get_parallel().enable_dp_attention,
|
||||
tp_rank,
|
||||
server_args.tp_size,
|
||||
server_args.dp_size,
|
||||
get_parallel().dp_size,
|
||||
server_args.attn_cp_size,
|
||||
)
|
||||
rank_port_args = PortArgs.init_new(
|
||||
|
||||
@@ -32,6 +32,7 @@ from sglang.srt.entrypoints.engine import (
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.ray.scheduler_actor import SchedulerActor
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -107,9 +108,9 @@ def _compute_world_size(server_args: ServerArgs) -> int:
|
||||
|
||||
Normal: dp_size * tp_size * pp_size; DP attention: tp_size * pp_size.
|
||||
"""
|
||||
if server_args.enable_dp_attention:
|
||||
if get_parallel().enable_dp_attention:
|
||||
return server_args.tp_size * server_args.pp_size
|
||||
return server_args.dp_size * server_args.tp_size * server_args.pp_size
|
||||
return get_parallel().dp_size * server_args.tp_size * server_args.pp_size
|
||||
|
||||
|
||||
def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]:
|
||||
@@ -267,11 +268,11 @@ class RayEngine(Engine):
|
||||
placement_group as create_placement_group,
|
||||
)
|
||||
|
||||
if server_args.enable_dp_attention:
|
||||
if get_parallel().enable_dp_attention:
|
||||
total_gpus = server_args.tp_size * server_args.pp_size
|
||||
else:
|
||||
total_gpus = (
|
||||
server_args.dp_size * server_args.tp_size * server_args.pp_size
|
||||
get_parallel().dp_size * server_args.tp_size * server_args.pp_size
|
||||
)
|
||||
|
||||
nnodes = server_args.nnodes
|
||||
@@ -313,7 +314,7 @@ class RayEngine(Engine):
|
||||
rank0_bundle_idx = int(indices_str.split(",")[0]) if indices_str else 0
|
||||
rank0_node_ip = _get_bundle_node_ip(pg, rank0_bundle_idx)
|
||||
|
||||
if server_args.dp_size == 1:
|
||||
if get_parallel().dp_size == 1:
|
||||
dist_init_addr = f"{rank0_node_ip}:{port_args.nccl_port}"
|
||||
logger.info(f"dist_init_addr: {dist_init_addr}")
|
||||
|
||||
@@ -446,17 +447,19 @@ class RayEngine(Engine):
|
||||
RayDataParallelController,
|
||||
)
|
||||
|
||||
if server_args.enable_dp_attention:
|
||||
if get_parallel().enable_dp_attention:
|
||||
# DP attention folds DP into TP — total GPUs = tp_size * pp_size
|
||||
total_gpus = server_args.tp_size * server_args.pp_size
|
||||
else:
|
||||
total_gpus = server_args.dp_size * server_args.tp_size * server_args.pp_size
|
||||
total_gpus = (
|
||||
get_parallel().dp_size * server_args.tp_size * server_args.pp_size
|
||||
)
|
||||
gpus_per_node = total_gpus // server_args.nnodes
|
||||
logger.info(
|
||||
f"Ray DP cluster: {server_args.nnodes} nodes, "
|
||||
f"{gpus_per_node} GPUs/node, dp_size={server_args.dp_size}, "
|
||||
f"{gpus_per_node} GPUs/node, dp_size={get_parallel().dp_size}, "
|
||||
f"tp_size={server_args.tp_size}, pp_size={server_args.pp_size}, "
|
||||
f"enable_dp_attention={server_args.enable_dp_attention}"
|
||||
f"enable_dp_attention={get_parallel().enable_dp_attention}"
|
||||
)
|
||||
|
||||
# Set dist_init_addr on server_args so PortArgs.init_new() can compute
|
||||
|
||||
@@ -96,14 +96,14 @@ class DSparkWorkerV2(BaseSpecWorker):
|
||||
|
||||
self._draft_is_moe = draft_is_deepseek_v4(server_args=server_args)
|
||||
self._draft_dp_context_enabled = (
|
||||
server_args.enable_dp_attention and not self._draft_is_moe
|
||||
get_parallel().enable_dp_attention and not self._draft_is_moe
|
||||
)
|
||||
self._is_pd_prefill = server_args.disaggregation_mode == "prefill"
|
||||
self._decode_graph_allowed = (
|
||||
not server_args.disable_cuda_graph and not self._is_pd_prefill
|
||||
not get_exec().graph.disable_cuda_graph and not self._is_pd_prefill
|
||||
)
|
||||
if (
|
||||
server_args.enable_dp_attention
|
||||
get_parallel().enable_dp_attention
|
||||
and self._draft_is_moe
|
||||
and ps.attn_tp_size > 1
|
||||
):
|
||||
@@ -202,7 +202,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
||||
verify_num_draft_tokens=self.verify_num_draft_tokens,
|
||||
)
|
||||
if (
|
||||
server_args.enable_dp_attention
|
||||
get_parallel().enable_dp_attention
|
||||
and not self._draft_is_moe
|
||||
and self._verify_planner.is_compact_mode
|
||||
and self._decode_graph_allowed
|
||||
@@ -228,7 +228,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
||||
gamma=self.gamma,
|
||||
mask_token_id=self._mask_token_id,
|
||||
draft_block_spec_info=self._draft_block_spec_info,
|
||||
dp_moe_sync=self._draft_is_moe and server_args.enable_dp_attention,
|
||||
dp_moe_sync=self._draft_is_moe and get_parallel().enable_dp_attention,
|
||||
)
|
||||
self._verify_epilogue = None
|
||||
if (
|
||||
|
||||
@@ -158,7 +158,10 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
||||
self._rebuild_topk1_chain_buffers()
|
||||
|
||||
# Load draft model weights only.
|
||||
if server_args.enable_dp_attention and self.speculative_algorithm.is_eagle3():
|
||||
if (
|
||||
get_parallel().enable_dp_attention
|
||||
and self.speculative_algorithm.is_eagle3()
|
||||
):
|
||||
ctx = draft_tp_context(get_parallel().attn_tp_group)
|
||||
else:
|
||||
ctx = empty_context()
|
||||
@@ -182,7 +185,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
||||
# Eager draft-extend seed buffer (graph paths use their own static ones).
|
||||
self.dsa_extend_topk_buf: Optional[torch.Tensor] = None
|
||||
self.draft_tp_context = (
|
||||
draft_tp_context if server_args.enable_dp_attention else empty_context
|
||||
draft_tp_context if get_parallel().enable_dp_attention else empty_context
|
||||
)
|
||||
self.tree_mask_mode = default_tree_mask_mode()
|
||||
|
||||
|
||||
@@ -44,7 +44,12 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
)
|
||||
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
|
||||
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
|
||||
from sglang.srt.runtime_context import attention_backends, get_schedule, get_spec
|
||||
from sglang.srt.runtime_context import (
|
||||
attention_backends,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase
|
||||
from sglang.srt.speculative.eagle_utils import (
|
||||
@@ -157,7 +162,7 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
|
||||
self.kv_context: Optional[FrozenKVMTPContext] = None
|
||||
|
||||
self.draft_tp_context = (
|
||||
draft_tp_context if server_args.enable_dp_attention else empty_context
|
||||
draft_tp_context if get_parallel().enable_dp_attention else empty_context
|
||||
)
|
||||
|
||||
self.draft_attn_backend = None
|
||||
|
||||
@@ -43,7 +43,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
CaptureHiddenMode,
|
||||
ForwardBatch,
|
||||
)
|
||||
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.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase
|
||||
from sglang.srt.speculative.draft_utils import DraftBackendFactory
|
||||
@@ -180,7 +180,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
||||
"InklingForConditionalGenerationMTP",
|
||||
]
|
||||
self.draft_tp_context = (
|
||||
draft_tp_context if server_args.enable_dp_attention else empty_context
|
||||
draft_tp_context if get_parallel().enable_dp_attention else empty_context
|
||||
)
|
||||
self.tree_mask_mode = default_tree_mask_mode()
|
||||
self.plan_stream, self.plan_stream_ctx = get_plan_stream(self.device)
|
||||
|
||||
@@ -10,7 +10,7 @@ from sglang.srt.layers.moe.utils import (
|
||||
speculative_moe_backend_context,
|
||||
)
|
||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||
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.adaptive_runtime_state import (
|
||||
AdaptiveController,
|
||||
@@ -89,7 +89,7 @@ class StandaloneDraftWorker(EagleDraftWorker):
|
||||
# Alias for better readability
|
||||
self.draft_runner = self.draft_worker.model_runner
|
||||
self.draft_tp_context = (
|
||||
draft_tp_context if server_args.enable_dp_attention else empty_context
|
||||
draft_tp_context if get_parallel().enable_dp_attention else empty_context
|
||||
)
|
||||
self.tree_mask_mode = default_tree_mask_mode()
|
||||
self.plan_stream, self.plan_stream_ctx = get_plan_stream(self.device)
|
||||
|
||||
@@ -31,7 +31,10 @@ from sglang.srt.managers.schedule_batch import (
|
||||
MultimodalDataItem,
|
||||
MultimodalProcessorOutput,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
configured_tp_size,
|
||||
get_parallel,
|
||||
)
|
||||
from sglang.srt.utils.cuda_ipc_transport_utils import (
|
||||
DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY,
|
||||
MM_FEATURE_CACHE_SIZE,
|
||||
@@ -159,9 +162,9 @@ def _contains_tensor_container(value) -> bool:
|
||||
|
||||
|
||||
def get_vmm_feature_consumer_count(server_args) -> int:
|
||||
if server_args.enable_dp_attention:
|
||||
return server_args.tp_size // server_args.dp_size
|
||||
return server_args.tp_size
|
||||
if get_parallel().enable_dp_attention:
|
||||
return configured_tp_size() // get_parallel().dp_size
|
||||
return configured_tp_size()
|
||||
|
||||
|
||||
class CudaVmmMemoryPool:
|
||||
|
||||
@@ -12,7 +12,11 @@ from sglang.srt.distributed.naive_distributed import (
|
||||
set_naive_distributed,
|
||||
)
|
||||
from sglang.srt.layers.parameter import ModelWeightParameter
|
||||
from sglang.srt.runtime_context import get_parallel, get_stream
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import MultiprocessingSerializer, is_pin_memory_available
|
||||
from sglang.srt.utils.host_shared_memory import (
|
||||
@@ -63,21 +67,21 @@ def set_offloader(instance: BaseOffloader):
|
||||
|
||||
|
||||
def create_offloader_from_server_args(server_args: ServerArgs, dp_rank: int):
|
||||
if server_args.cpu_offload_gb > 0:
|
||||
if get_exec().offload.cpu_offload_gb > 0:
|
||||
return OffloaderV1(
|
||||
cpu_offload_max_bytes=int(server_args.cpu_offload_gb * 1024**3)
|
||||
cpu_offload_max_bytes=int(get_exec().offload.cpu_offload_gb * 1024**3)
|
||||
)
|
||||
if server_args.offload_group_size > 0:
|
||||
if get_exec().offload.offload_group_size > 0:
|
||||
assert (
|
||||
server_args.cpu_offload_gb == 0
|
||||
get_exec().offload.cpu_offload_gb == 0
|
||||
), "V2 offload does not support cpu_offload_gb yet"
|
||||
return OffloaderV2(
|
||||
group_size=server_args.offload_group_size,
|
||||
num_in_group=server_args.offload_num_in_group,
|
||||
prefetch_step=server_args.offload_prefetch_step,
|
||||
mode=server_args.offload_mode,
|
||||
group_size=get_exec().offload.offload_group_size,
|
||||
num_in_group=get_exec().offload.offload_num_in_group,
|
||||
prefetch_step=get_exec().offload.offload_prefetch_step,
|
||||
mode=get_exec().offload.offload_mode,
|
||||
dp_rank=dp_rank,
|
||||
dp_size=server_args.dp_size,
|
||||
dp_size=get_parallel().dp_size,
|
||||
)
|
||||
return NoopOffloader()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user