config: the DP/EP topology reads come from the parallel bag (#35025)

This commit is contained in:
Cheng Wan
2026-08-17 16:17:26 -07:00
committed by GitHub
parent d2bc697396
commit a97bc8db32
33 changed files with 332 additions and 304 deletions
@@ -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;
+7 -7
View File
@@ -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
+10 -6
View File
@@ -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:
+9 -6
View File
@@ -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),
)
+7 -3
View File
@@ -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
)
)
+26 -19
View File
@@ -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
+16 -19
View File
@@ -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(
+12 -9
View File
@@ -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:
+14 -10
View File
@@ -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()