config: stop handing the record to code that does not read it (#36252)

This commit is contained in:
Cheng Wan
2026-08-26 05:02:17 -07:00
committed by GitHub
parent d7b144f64e
commit ae5feb4b9c
60 changed files with 168 additions and 280 deletions
+2 -2
View File
@@ -538,7 +538,7 @@ def decode(input_token_ids, batch, model_runner):
def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
if require_mlp_sync(model_runner.server_args):
if require_mlp_sync():
prepare_mlp_sync_batch_raw(
batch,
model_runner=model_runner,
@@ -548,7 +548,7 @@ def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
tp_group=model_runner.tp_group,
get_idle_batch=None,
disable_cuda_graph=cuda_graph_fully_disabled(),
require_mlp_tp_gather=require_mlp_tp_gather(model_runner.server_args),
require_mlp_tp_gather=require_mlp_tp_gather(),
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
offload_tags=set(),
)
@@ -14,7 +14,7 @@ from sglang.srt.runtime_context import (
from sglang.srt.utils.network import NetworkAddress, get_free_port, get_local_ip_auto
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
pass
logger = logging.getLogger(__name__)
@@ -307,9 +307,7 @@ def get_mooncake_transfer_engine() -> Optional[MooncakeTransferEngine]:
return _mooncake_transfer_engine
def maybe_init_shared_mooncake_transfer_engine(
*, server_args: ServerArgs, gpu_id: int
) -> None:
def maybe_init_shared_mooncake_transfer_engine(*, gpu_id: int) -> None:
"""
Need MooncakeTransferEngine when:
1) PD disaggregation uses mooncake for KV transfer (prefill/decode)
+2 -2
View File
@@ -111,14 +111,14 @@ class ElasticEPStateManager:
inst.ep_join_rank_offset = get_parallel().config.ep_join_rank_offset
if server_args.is_ep_joiner:
cls._init_joiner_state(inst, server_args)
cls._init_joiner_state(inst)
cls._instance = inst
return cls._instance
@classmethod
def _init_joiner_state(cls, inst: ElasticEPState, server_args: ServerArgs) -> None:
def _init_joiner_state(cls, inst: ElasticEPState) -> None:
global_rank = torch.distributed.get_rank()
inst.active_ranks.zero_()
inst.active_ranks[global_rank] = 1
@@ -2142,7 +2142,6 @@ def _get_vlm_warmup_image_base64(model_info: dict) -> str:
async def _send_disaggregation_warmup_requests(
server_args: ServerArgs,
url: str,
headers: Dict[str, str],
ssl_verify: Union[bool, str],
@@ -2321,7 +2320,6 @@ def _execute_server_warmup(server_args: ServerArgs):
logger.info(f"Start of pd disaggregation warmup ...")
status_codes = asyncio.run(
_send_disaggregation_warmup_requests(
server_args=server_args,
url=url,
headers=headers,
ssl_verify=ssl_verify,
+10 -16
View File
@@ -18,11 +18,10 @@ from sglang.srt.eplb.expert_location import (
get_global_expert_location_metadata,
)
from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater
from sglang.srt.runtime_context import get_model
from sglang.srt.runtime_context import get_exec, get_model, get_parallel
if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__)
@@ -31,7 +30,6 @@ class EPLBManager:
def __init__(
self,
*,
server_args: ServerArgs,
model_config: ModelConfig,
ps: Any,
get_model: Callable[[], nn.Module],
@@ -43,7 +41,6 @@ class EPLBManager:
# These collaborators are set on ModelRunner AFTER EPLBManager is
# constructed (model load, expert_backup_client, weight_updater), so
# they are read through getters at rebalance time, not captured here.
self._server_args = server_args
self._model_config = model_config
self._ps = ps
self._get_model = get_model
@@ -51,16 +48,16 @@ class EPLBManager:
self._get_expert_backup_client = get_expert_backup_client
self._get_weight_updater = get_weight_updater
self._rebalance_layers_per_chunk = (
self._server_args.eplb_rebalance_layers_per_chunk
get_exec().moe.eplb_rebalance_layers_per_chunk
)
self._rebalance_num_iterations = self._server_args.eplb_rebalance_num_iterations
self._rebalance_num_iterations = get_exec().moe.eplb_rebalance_num_iterations
self._rebalance_disabled_reason = None
self._rebalance_disabled_logged = False
# Otherwise, the circular buffer will contain stale data. If the case is needed, it can be implemented.
assert (
self._server_args.eplb_rebalance_num_iterations
>= self._server_args.expert_distribution_recorder_buffer_size
get_exec().moe.eplb_rebalance_num_iterations
>= get_exec().moe.expert_distribution_recorder_buffer_size
), "eplb_rebalance_num_iterations must be greater than expert_distribution_recorder_buffer_size"
if not get_global_expert_distribution_recorder().recording:
@@ -160,7 +157,7 @@ class EPLBManager:
model=self._get_model(),
new_expert_location_metadata=expert_location_metadata,
update_layer_ids=chunk_layer_ids,
nnodes=self._server_args.nnodes,
nnodes=get_parallel().config.nnodes,
tp_rank=(
self._elastic_global_rank()
if is_post_scale_rebalance
@@ -169,7 +166,7 @@ class EPLBManager:
use_flat_topology=is_post_scale_rebalance,
expert_backup_client=self._get_expert_backup_client(),
update_weights_from_disk_callable=self._get_weight_updater().update_weights_from_disk,
ep_dispatch_algorithm=self._server_args.ep_dispatch_algorithm,
ep_dispatch_algorithm=get_exec().moe.ep_dispatch_algorithm,
init_lplb_solvers_callable=lambda: init_lplb_solvers(
model_config=self._model_config
),
@@ -193,7 +190,6 @@ class EPLBManager:
) -> ExpertLocationMetadata:
if not broadcast_over_world:
return ExpertLocationMetadata.init_by_eplb(
self._server_args,
self._model_config,
logical_count,
)
@@ -204,7 +200,6 @@ class EPLBManager:
# the mapping chosen for the expanded world.
if dist.get_rank() == 0:
computed_metadata = ExpertLocationMetadata.init_by_eplb(
self._server_args,
self._model_config,
logical_count,
# Arbitrary append topologies may not preserve node divisibility.
@@ -220,14 +215,13 @@ class EPLBManager:
dist.broadcast(physical_to_logical_map, src=0)
return ExpertLocationMetadata.init_by_mapping(
self._server_args,
self._model_config,
physical_to_logical_map,
moe_ep_rank=self._elastic_global_rank(),
)
def _elastic_global_rank(self) -> int:
return self._ps.tp_rank + self._server_args.ep_join_rank_offset
return self._ps.tp_rank + get_parallel().config.ep_join_rank_offset
def _check_rebalance_needed(self, average_utilization_rate_over_window):
if average_utilization_rate_over_window is None:
@@ -235,10 +229,10 @@ class EPLBManager:
if (
average_utilization_rate_over_window
> self._server_args.eplb_min_rebalancing_utilization_threshold
> get_exec().moe.eplb_min_rebalancing_utilization_threshold
):
logger.info(
f"[EPLBManager] Skipped ep rebalancing: current GPU utilization {average_utilization_rate_over_window:.2f} > minimum rebalance threshold {self._server_args.eplb_min_rebalancing_utilization_threshold:.2f}"
f"[EPLBManager] Skipped ep rebalancing: current GPU utilization {average_utilization_rate_over_window:.2f} > minimum rebalance threshold {get_exec().moe.eplb_min_rebalancing_utilization_threshold:.2f}"
)
return False
@@ -671,12 +671,10 @@ class _Accumulator(ABC):
expert_location_metadata: ExpertLocationMetadata,
rank: int,
) -> _Accumulator:
return _Accumulator.get_class(server_args)(
server_args, expert_location_metadata, rank
)
return _Accumulator.get_class()(server_args, expert_location_metadata, rank)
@staticmethod
def get_class(server_args: ServerArgs) -> Type[_Accumulator]:
def get_class() -> Type[_Accumulator]:
return {
"stat": _StatAccumulator,
"stat_approx": _StatAccumulator,
+7 -21
View File
@@ -32,7 +32,6 @@ from sglang.srt.runtime_context import (
if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__)
@@ -105,11 +104,9 @@ class ExpertLocationMetadata:
# -------------------------------- construction ------------------------------------
@staticmethod
def init_trivial(
server_args: ServerArgs, model_config: ModelConfig, moe_ep_rank: int
):
def init_trivial(model_config: ModelConfig, moe_ep_rank: int):
"""Trivial location - logical expert i corresponds to physical expert i"""
common = ExpertLocationMetadata._init_common(server_args, model_config)
common = ExpertLocationMetadata._init_common(model_config)
if common is None:
return None
@@ -131,7 +128,6 @@ class ExpertLocationMetadata:
)
return ExpertLocationMetadata.init_by_mapping(
server_args,
model_config,
physical_to_logical_map=physical_to_logical_map,
moe_ep_rank=moe_ep_rank,
@@ -139,7 +135,6 @@ class ExpertLocationMetadata:
@staticmethod
def init_by_mapping(
server_args: ServerArgs,
model_config: ModelConfig,
physical_to_logical_map,
moe_ep_rank: int = None,
@@ -148,7 +143,7 @@ class ExpertLocationMetadata:
physical_to_logical_map = torch.tensor(physical_to_logical_map)
physical_to_logical_map = physical_to_logical_map.to(get_device().device)
common = ExpertLocationMetadata._init_common(server_args, model_config)
common = ExpertLocationMetadata._init_common(model_config)
if common is None:
return None
@@ -179,7 +174,6 @@ class ExpertLocationMetadata:
@staticmethod
def init_by_eplb(
server_args: ServerArgs,
model_config: ModelConfig,
logical_count: torch.Tensor,
*,
@@ -193,7 +187,7 @@ class ExpertLocationMetadata:
from sglang.srt.runtime_context import get_parallel
common = ExpertLocationMetadata._init_common(server_args, model_config)
common = ExpertLocationMetadata._init_common(model_config)
if common is None:
return None
@@ -229,7 +223,7 @@ class ExpertLocationMetadata:
)
@staticmethod
def _init_common(server_args: ServerArgs, model_config: ModelConfig):
def _init_common(model_config: ModelConfig):
from sglang.srt.runtime_context import get_exec, get_parallel
model_config_for_expert_location = (
@@ -526,9 +520,6 @@ def broadcast_global_expert_location_metadata(
src_rank: int = 0,
group: Optional[torch.distributed.ProcessGroup] = None,
) -> ExpertLocationMetadata:
from sglang.srt.runtime_context import get_server_args
server_args = get_server_args()
metadata = get_global_expert_location_metadata()
assert metadata is not None
@@ -537,7 +528,6 @@ def broadcast_global_expert_location_metadata(
metadata.physical_to_logical_map, src=src_rank, group=group
)
metadata = ExpertLocationMetadata.init_by_mapping(
server_args,
model_config,
metadata.physical_to_logical_map,
moe_ep_rank=moe_ep_rank,
@@ -785,15 +775,12 @@ class ModelConfigForExpertLocation:
def compute_initial_expert_location_metadata(
server_args: ServerArgs,
model_config: ModelConfig,
moe_ep_rank: int,
) -> Optional[ExpertLocationMetadata]:
data = get_exec().moe.init_expert_location
if data == "trivial":
return ExpertLocationMetadata.init_trivial(
server_args, model_config, moe_ep_rank
)
return ExpertLocationMetadata.init_trivial(model_config, moe_ep_rank)
# TODO unify with the utils function
if data.endswith(".pt"):
@@ -808,7 +795,6 @@ def compute_initial_expert_location_metadata(
"init_expert_location from init_by_mapping using ServerArgs.init_expert_location"
)
return ExpertLocationMetadata.init_by_mapping(
server_args,
model_config,
**data_dict,
moe_ep_rank=moe_ep_rank,
@@ -818,7 +804,7 @@ def compute_initial_expert_location_metadata(
"init_expert_location from init_by_eplb using ServerArgs.init_expert_location"
)
return ExpertLocationMetadata.init_by_eplb(
server_args, model_config, logical_count=data_dict["logical_count"]
model_config, logical_count=data_dict["logical_count"]
)
else:
raise NotImplementedError(
@@ -442,9 +442,7 @@ class FlashInferAttnBackend(AttentionBackend):
)
else:
self.workspace_buffer = global_workspace_buffer
max_bs = get_cuda_graph_max_batch_size(
model_runner.server_args, model_runner.req_to_token_pool.size
)
max_bs = get_cuda_graph_max_batch_size(model_runner.req_to_token_pool.size)
if kv_indptr_buf is None:
self.kv_indptr = [
torch.zeros(
@@ -2254,7 +2252,7 @@ class FlashInferMultiStepDraftBackend:
self.page_size = model_runner.page_size
max_bs = get_cuda_graph_max_batch_size(
model_runner.server_args, model_runner.req_to_token_pool.size * self.topk
model_runner.req_to_token_pool.size * self.topk
)
self.kv_indptr = torch.zeros(
(
@@ -400,7 +400,6 @@ class GDNAttnBackend(MambaAttnBackendBase):
self.verify_intermediate_state_indices = (
build_verify_intermediate_state_indices(
self.req_to_token_pool.size,
model_runner.server_args,
model_runner.device,
)
)
@@ -408,7 +408,6 @@ class KDAAttnBackend(MambaAttnBackendBase):
self.verify_intermediate_state_indices = (
build_verify_intermediate_state_indices(
self.req_to_token_pool.size,
model_runner.server_args,
model_runner.device,
)
)
@@ -9,7 +9,7 @@ from sglang.srt.runtime_context import get_exec
from sglang.srt.utils.common import rank0_log
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
pass
class LinearAttnKernelBackend(Enum):
@@ -103,9 +103,7 @@ def resolve_linear_attn_backends(
return backends
def build_verify_intermediate_state_indices(
pool_size: int, server_args: ServerArgs, device
):
def build_verify_intermediate_state_indices(pool_size: int, device):
"""Per-request row index into the speculative intermediate scratch
(`intermediate_ssm` / `intermediate_conv_window`) for the MTP /
target_verify path: request slot i owns scratch row i.
@@ -123,7 +121,7 @@ def build_verify_intermediate_state_indices(
from sglang.srt.utils.common import get_eager_max_batch_size
padded_bs = max(get_eager_max_batch_size(server_args, pool_size), pool_size)
padded_bs = max(get_eager_max_batch_size(pool_size), pool_size)
indices = torch.arange(pool_size, dtype=torch.int32, device=device)
if padded_bs > pool_size:
indices = torch.cat(
@@ -477,7 +477,7 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
cp = max(get_cp_padding_align_size(), 1)
attn_tp_size = get_parallel().attn_tp_size
mlp_sync = require_mlp_sync(model_runner.server_args)
mlp_sync = require_mlp_sync()
def _align(bs: int) -> int:
# Align to lcm(cp, attn_tp_size) so the CP multiple isn't undone by a
+5 -8
View File
@@ -94,19 +94,17 @@ class LoRAManager:
self.attn_tp_size: int = get_parallel().attn_tp_size
self.lora_added_tokens_size: Optional[int] = None
self.enable_lora_overlap_loading: Optional[bool] = (
server_args.enable_lora_overlap_loading
get_lora().enable_lora_overlap_loading
)
self.pending_lora_load_events = {}
self.eviction_policy = server_args.lora_eviction_policy
self.eviction_policy = get_lora().lora_eviction_policy
self.enable_dp_attention: bool = get_parallel().config.enable_dp_attention
self._experts_shared_outer_override: Optional[bool] = (
server_args.experts_shared_outer_loras
)
self.lora_use_virtual_experts: bool = server_args.lora_use_virtual_experts
self.lora_strict_loading: bool = getattr(
server_args, "lora_strict_loading", False
get_lora().experts_shared_outer_loras
)
self.lora_use_virtual_experts: bool = get_lora().lora_use_virtual_experts
self.lora_strict_loading: bool = get_lora().lora_strict_loading
self.speculative_algorithm: Optional[str] = get_spec().speculative_algorithm
# LoRA backend for running sgemm kernels
@@ -1030,7 +1028,6 @@ class LoRAManager:
def init_lora_cuda_graph_moe_buffers(
*,
server_args: ServerArgs,
model: torch.nn.Module,
lora_manager: LoRAManager,
dtype: torch.dtype,
+3 -10
View File
@@ -13,12 +13,9 @@ from sglang.srt.runtime_context import (
get_parallel,
get_serving,
)
from sglang.srt.server_args import ServerArgs
def start_disagg_service(
server_args: ServerArgs,
):
def start_disagg_service():
# Start kv bootstrap server on prefill
disagg_mode = DisaggregationMode(get_disagg().disaggregation_mode)
transfer_backend = TransferBackend(get_disagg().disaggregation_transfer_backend)
@@ -32,16 +29,12 @@ def start_disagg_service(
host=get_serving().host,
port=get_disagg().disaggregation_bootstrap_port,
)
maybe_create_ascend_config_store(
server_args=server_args, transfer_backend=transfer_backend
)
maybe_create_ascend_config_store(transfer_backend=transfer_backend)
return bootstrap_server
def maybe_create_ascend_config_store(
server_args: ServerArgs, transfer_backend: TransferBackend
) -> None:
def maybe_create_ascend_config_store(transfer_backend: TransferBackend) -> None:
"""Also called directly by the rust-server scheduler: there the KV
bootstrap registry is served by the embedded rust server's api listener
(one rust implementation covers every transfer backend their
@@ -478,7 +478,7 @@ class MultiTokenizerRouter:
)
self._loop.call_soon_threadsafe(self._register_load_snapshot_reader)
self.disaggregation_bootstrap_server = start_disagg_service(self.server_args)
self.disaggregation_bootstrap_server = start_disagg_service()
# Worker IPC names for pause/continue broadcasting
self.all_worker_ipcs: set[str] = set()
+1 -3
View File
@@ -18,14 +18,12 @@ if TYPE_CHECKING:
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.eagle_info import EagleDraftInput
from sglang.srt.speculative.ngram_info import NgramVerifyInput
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
def decide_needs_cpu_seq_lens(
server_args: ServerArgs,
attn_backends: Sequence[AttentionBackend],
) -> bool:
"""Whether FutureMap must publish seq_lens_cpu / sum.
@@ -53,7 +51,7 @@ def decide_needs_cpu_seq_lens(
)
def decide_needs_confidence_relay(server_args: ServerArgs) -> bool:
def decide_needs_confidence_relay() -> bool:
from sglang.srt.speculative.ragged_verify import (
RaggedVerifyMode,
read_ragged_verify_mode,
+5 -8
View File
@@ -416,7 +416,7 @@ class Scheduler(
# init_soft_watchdog starts a daemon thread that reads these on its first tick.
self.forward_ct: int = 0
self.cur_batch_for_debug: Optional[ScheduleBatch] = None
self.init_soft_watchdog(server_args)
self.init_soft_watchdog()
# Parse args
self.server_args = server_args
@@ -911,7 +911,7 @@ class Scheduler(
initialize_bf16_gemm_config(self.server_args)
# This must be called after initialize_moe_config
self.require_mlp_sync = require_mlp_sync(self.server_args)
self.require_mlp_sync = require_mlp_sync()
def init_tp_model_worker(self):
worker_kwargs = dict(
@@ -1289,7 +1289,7 @@ class Scheduler(
self.new_token_ratio_tracker = NewTokenRatioTracker.from_config()
def init_soft_watchdog(self, server_args: ServerArgs):
def init_soft_watchdog(self):
if (x := get_device().soft_watchdog_timeout) is not None:
self.soft_watchdog = create_scheduler_watchdog(
self, watchdog_timeout=x, soft=True
@@ -1341,7 +1341,6 @@ class Scheduler(
and self._hosts_rust_server()
):
maybe_create_ascend_config_store(
server_args=self.server_args,
transfer_backend=self.transfer_backend,
)
@@ -1493,8 +1492,8 @@ class Scheduler(
)
else:
attn_backends = (self.tp_worker.model_runner.attn_backend,)
needs_cpu_seq_lens = decide_needs_cpu_seq_lens(self.server_args, attn_backends)
needs_confidence_relay = decide_needs_confidence_relay(self.server_args)
needs_cpu_seq_lens = decide_needs_cpu_seq_lens(attn_backends)
needs_confidence_relay = decide_needs_confidence_relay()
self.future_map = self.spec_algorithm.create_future_map(
self.device,
self.req_to_token_pool,
@@ -2089,7 +2088,6 @@ class Scheduler(
tree_cache=self.tree_cache,
offload_tags=self.weight_updater.offload_tags,
ps=self.ps,
server_args=self.server_args,
model_config=self.model_config,
enable_overlap=self.enable_overlap,
spec_algorithm=self.spec_algorithm,
@@ -2213,7 +2211,6 @@ class Scheduler(
disaggregation_mode=self.disaggregation_mode,
enable_overlap=self.enable_overlap,
enable_overlap_mlx=self.enable_overlap_mlx,
server_args=self.server_args,
model_config=self.model_config,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tree_cache=self.tree_cache,
@@ -70,7 +70,6 @@ if TYPE_CHECKING:
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.observability.metrics_collector import SchedulerMetricsCollector
from sglang.srt.sampling.sampling_observer import HostAuxiliaryOutput
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__)
@@ -81,7 +80,6 @@ class SchedulerBatchResultProcessor:
disaggregation_mode: DisaggregationMode
enable_overlap: bool
enable_overlap_mlx: bool
server_args: ServerArgs
model_config: ModelConfig
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator
tree_cache: BasePrefixCache
@@ -280,7 +278,6 @@ class SchedulerBatchResultProcessor:
hidden_state_offset = 0
prefill_hidden_capture_mode = self._get_prefill_hidden_capture_mode(
batch,
self.server_args,
)
# Check finish conditions
@@ -617,10 +614,7 @@ class SchedulerBatchResultProcessor:
)
@staticmethod
def _get_prefill_hidden_capture_mode(
batch: ScheduleBatch,
server_args: ServerArgs,
) -> CaptureHiddenMode:
def _get_prefill_hidden_capture_mode(batch: ScheduleBatch) -> CaptureHiddenMode:
return get_required_capture_hidden_mode(
max(
batch.return_hidden_states_mode,
@@ -27,7 +27,6 @@ from sglang.srt.model_executor.cuda_graph_config import (
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.observability.metrics_collector import DPCooperationInfo
from sglang.srt.runtime_context import get_parallel, get_schedule
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.utils.common import require_mlp_tp_gather
@@ -401,7 +400,6 @@ class SchedulerDPAttnAdapter:
tree_cache: BasePrefixCache
offload_tags: set[str]
ps: ParallelState
server_args: ServerArgs
model_config: ModelConfig
enable_overlap: bool
spec_algorithm: SpeculativeAlgorithm
@@ -417,7 +415,7 @@ class SchedulerDPAttnAdapter:
tp_group=self.tp_group,
get_idle_batch=self.get_idle_batch,
disable_cuda_graph=cuda_graph_fully_disabled(),
require_mlp_tp_gather=require_mlp_tp_gather(self.server_args),
require_mlp_tp_gather=require_mlp_tp_gather(),
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
offload_tags=self.offload_tags,
dwdp=get_parallel().config.dwdp_size > 1,
@@ -81,7 +81,7 @@ from sglang.srt.runtime_context import (
get_parallel,
get_spec,
)
from sglang.srt.server_args import LoRARef, ServerArgs
from sglang.srt.server_args import LoRARef
from sglang.srt.utils import (
get_bool_env_var,
normalize_serialized_named_tensor_payloads,
@@ -158,7 +158,7 @@ class TokenizerControlMixin:
FanOutCommunicator, as opposed to data-plane inference requests multiplexed by rid.
"""
def init_communicators(self: TokenizerManager, server_args: ServerArgs):
def init_communicators(self: TokenizerManager):
dispatch_pairs = []
for spec in _COMMUNICATOR_SPECS:
name, resp_type = spec[0], spec[1]
@@ -663,9 +663,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
# Keep a reference so the bootstrap server is not garbage-collected.
self.bootstrap_server = (
start_disagg_service(self.server_args)
if start_pd_bootstrap_service
else None
start_disagg_service() if start_pd_bootstrap_service else None
)
# Single-source counter for auto-assigning fake bootstrap_room.
self.fake_bootstrap_room_counter = 0
@@ -761,7 +759,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
(ElasticScaleUpdateReq, self.forward_elastic_scale_update),
]
)
self.init_communicators(self.server_args)
self.init_communicators()
self.sampling_params_class = SamplingParams
self.signal_handler_class = SignalHandler
@@ -410,7 +410,6 @@ def build_hybrid_swa_stack(
def _deepseek_v4_num_host_pages(
*,
params: CacheInitParams,
server_args: ServerArgs,
kvcache: Any,
page_size: int,
swa_page_size: int,
@@ -511,7 +510,6 @@ def build_deepseek_v4_hicache_stack(
}
num_host_pages, swa_num_host_pages = _deepseek_v4_num_host_pages(
params=params,
server_args=server_args,
kvcache=kvcache,
page_size=page_size,
swa_page_size=kvcache.swa_page_size,
@@ -1152,7 +1152,6 @@ class KVCacheConfigurator:
kv_cache_dim=calculate_mla_kv_cache_dim(
model_config=self.model_config,
kv_cache_dtype=self.kv_cache_dtype,
server_args=self.server_args,
),
enable_memory_saver=get_exec().features.enable_memory_saver,
start_layer=self.layer_info.start_layer,
@@ -1356,7 +1355,6 @@ class KVCacheConfigurator:
kv_cache_dim=calculate_mla_kv_cache_dim(
model_config=self.model_config,
kv_cache_dtype=self.kv_cache_dtype,
server_args=self.server_args,
),
enable_memory_saver=get_exec().features.enable_memory_saver,
start_layer=self.layer_info.start_layer,
@@ -1396,7 +1394,6 @@ class KVCacheConfigurator:
kv_cache_dim=calculate_mla_kv_cache_dim(
model_config=self.model_config,
kv_cache_dtype=self.kv_cache_dtype,
server_args=self.server_args,
),
)
@@ -2199,10 +2196,7 @@ class KVCacheConfigurator:
def calculate_mla_kv_cache_dim(
*,
model_config: ModelConfig,
kv_cache_dtype: torch.dtype,
server_args: ServerArgs,
*, model_config: ModelConfig, kv_cache_dtype: torch.dtype
) -> int:
is_dsa_model = is_deepseek_dsa(model_config.hf_config)
kv_cache_dtype = kv_cache_dtype
@@ -598,10 +598,10 @@ class CPUGraphRunner:
self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather()
self.enable_two_batch_overlap = (
model_runner.server_args.enable_two_batch_overlap
)
@@ -986,14 +986,14 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
pin_memory=is_pin_memory_available(batch.device),
).to(batch.device, non_blocking=True)
def adjust_num_token_non_padded_for_attn_tp(self, server_args) -> None:
def adjust_num_token_non_padded_for_attn_tp(self) -> None:
"""Make num_token_non_padded local to this attention-TP rank."""
from sglang.srt.utils.common import require_mlp_tp_gather
dp_rank = get_parallel().attn_dp_rank
assert self.global_num_tokens_cpu is not None
if require_mlp_tp_gather(server_args):
if require_mlp_tp_gather():
num_tokens_per_dp = self.global_num_tokens_cpu[dp_rank]
else:
num_tokens_per_dp = self.global_num_tokens_cpu[0]
@@ -223,7 +223,7 @@ from sglang.srt.utils.device_timer import device_timer_ctx
from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks
from sglang.srt.utils.nvtx_utils import profile_range
from sglang.srt.utils.offloader import (
create_offloader_from_server_args,
create_offloader,
get_offloader,
set_offloader,
)
@@ -270,7 +270,6 @@ class ModelRunnerOutput:
def resolve_draft_attention_backend(
*,
draft_attention_backend: Optional[str],
server_args: ServerArgs,
is_draft_worker: bool,
) -> Optional[str]:
"""The attention backend a runner uses because it is a draft runner.
@@ -342,7 +341,6 @@ class ModelRunner:
self.device = get_device().device
self.draft_attention_backend = resolve_draft_attention_backend(
draft_attention_backend=draft_attention_backend,
server_args=server_args,
is_draft_worker=is_draft_worker,
)
# This runner's own load format, resolved before anything keys off it:
@@ -430,9 +428,7 @@ class ModelRunner:
self.shared_read_done_event: Optional[torch.cuda.Event] = None
# CPU offload
set_offloader(
create_offloader_from_server_args(server_args, dp_rank=self.ps.dp_rank)
)
set_offloader(create_offloader(dp_rank=self.ps.dp_rank))
self._weight_checker = WeightChecker(get_model=lambda: self.model, ps=self.ps)
@@ -591,7 +587,6 @@ class ModelRunner:
model=self.model,
model_config=self.model_config,
req_to_token_pool=self.req_to_token_pool,
server_args=self.server_args,
max_running_requests=self.max_running_requests,
device=self.device,
)
@@ -702,7 +697,6 @@ class ModelRunner:
)
set_global_expert_location_metadata(
compute_initial_expert_location_metadata(
server_args=self.server_args,
model_config=self.model_config,
moe_ep_rank=expert_rank,
)
@@ -727,7 +721,6 @@ class ModelRunner:
def maybe_init_eplb_manager(self):
self.eplb_manager = (
EPLBManager(
server_args=self.server_args,
model_config=self.model_config,
ps=self.ps,
get_model=lambda: self.model,
@@ -1076,9 +1069,7 @@ class ModelRunner:
self.pre_model_load_memory = result.pre_model_load_memory
def init_shared_mooncake_transfer_engine(self):
maybe_init_shared_mooncake_transfer_engine(
server_args=self.server_args, gpu_id=self.gpu_id
)
maybe_init_shared_mooncake_transfer_engine(gpu_id=self.gpu_id)
def load_model(self):
tic_total = time.perf_counter()
@@ -1091,9 +1082,7 @@ class ModelRunner:
if self.device != "cpu":
torch.set_num_threads(1)
if self.device == "cuda":
maybe_downgrade_dtype_for_legacy_gpu(
server_args=self.server_args, model_config=self.model_config
)
maybe_downgrade_dtype_for_legacy_gpu(model_config=self.model_config)
set_cuda_arch()
@@ -1113,7 +1102,6 @@ class ModelRunner:
# and derive the per-rank daemon socket. Idempotent across reloads.
maybe_enable_ipc_weight_cache(
load_config=self.load_config,
server_args=self.server_args,
tp_size=self.ps.tp_size,
pp_rank=self.ps.pp_rank,
tp_rank=self.ps.tp_rank,
@@ -1124,7 +1112,6 @@ class ModelRunner:
)
maybe_trigger_remote_instance_nccl_send_group(
server_args=self.server_args,
tp_rank=self.ps.tp_rank,
load_format=draft_load_format,
)
@@ -1195,11 +1182,12 @@ class ModelRunner:
f"mem usage={self.weight_load_mem_usage:.2f} GB."
)
report_online_quantization(model=self.model, server_args=self.server_args)
report_online_quantization(
model=self.model,
)
maybe_register_debug_tensor_dump_hook(
model=self.model,
server_args=self.server_args,
spec_algorithm=self.spec_algorithm,
is_draft_worker=self.is_draft_worker,
tp_size=self.ps.tp_size,
@@ -1281,7 +1269,6 @@ class ModelRunner:
)
if not cuda_graph_fully_disabled():
init_lora_cuda_graph_moe_buffers(
server_args=self.server_args,
model=self.model,
lora_manager=self.lora_manager,
dtype=self.dtype,
@@ -1471,13 +1458,11 @@ class ModelRunner:
if (
forward_batch.num_token_non_padded is not None
and forward_batch.global_num_tokens_gpu is not None
and require_gathered_buffer(self.server_args)
and require_gathered_buffer()
and not is_dsa_enable_prefill_cp()
and not is_mla_prefill_cp_enabled()
):
forward_batch.adjust_num_token_non_padded_for_attn_tp(
server_args=self.server_args,
)
forward_batch.adjust_num_token_non_padded_for_attn_tp()
# Hisparse coordinator — backends now read it from self.model_runner.
if self.hisparse_coordinator is not None:
@@ -1926,7 +1911,6 @@ class ModelRunner:
start=old_num_physical - num_local * initial_ep_size,
)
new_metadata = ExpertLocationMetadata.init_by_mapping(
self.server_args,
self.model_config,
physical_to_logical_map=expanded_p2l,
moe_ep_rank=self._elastic_global_rank(),
@@ -67,9 +67,7 @@ class LoadedModel(msgspec.Struct, frozen=True, kw_only=True):
startup_weight_load: Optional[Any] = None
def maybe_downgrade_dtype_for_legacy_gpu(
*, server_args: ServerArgs, model_config: ModelConfig
) -> None:
def maybe_downgrade_dtype_for_legacy_gpu(*, model_config: ModelConfig) -> None:
if torch.cuda.get_device_capability()[0] < 8:
logger.info(
"Compute capability below sm80. Use float16 due to lack of bfloat16 support."
@@ -85,7 +83,7 @@ def maybe_downgrade_dtype_for_legacy_gpu(
def maybe_trigger_remote_instance_nccl_send_group(
*, server_args: ServerArgs, tp_rank: int, load_format: Optional[str] = None
*, tp_rank: int, load_format: Optional[str] = None
) -> None:
"""``load_format`` is this runner's effective format: a draft loading under
``--speculative-draft-draft-load-format`` needs its own send group, and the
@@ -151,7 +149,7 @@ def resolve_sliding_window_size(model, model_config: ModelConfig) -> Optional[in
return sliding_window_size
def report_online_quantization(*, model, server_args: ServerArgs) -> None:
def report_online_quantization(*, model) -> None:
# TODO: Make sure all models have `quant_config` attribute, and all online quantization methods register which layers they actually quantize.
quantized_layers = getattr(
getattr(model, "quant_config", None), "quantized_layers", None
@@ -170,7 +168,6 @@ def report_online_quantization(*, model, server_args: ServerArgs) -> None:
def maybe_register_debug_tensor_dump_hook(
*,
model,
server_args: ServerArgs,
spec_algorithm: SpeculativeAlgorithm,
is_draft_worker: bool,
tp_size: int,
@@ -237,7 +234,6 @@ def build_load_config(
def maybe_enable_ipc_weight_cache(
*,
load_config: LoadConfig,
server_args: ServerArgs,
tp_size: int,
pp_rank: int,
tp_rank: int,
@@ -311,12 +307,11 @@ def load_model_with_memory_saver(
StartupWeightLoadManager,
)
startup_weight_load = StartupWeightLoadManager.create_from_server_args(
startup_weight_load = StartupWeightLoadManager.create_from_published_config(
loader=loader,
model_config=model_config,
load_config=load_config,
device_config=device_config,
server_args=server_args,
is_draft_worker=is_draft_worker,
)
model = startup_weight_load.prepare()
@@ -12,7 +12,6 @@ from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.managers.schedule_batch import ForwardMode
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.runtime_context import get_schedule
from sglang.srt.server_args import ServerArgs
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
@@ -33,7 +32,6 @@ class NgramEmbeddingManager:
model: torch.nn.Module,
model_config: ModelConfig,
req_to_token_pool: ReqToTokenPool,
server_args: ServerArgs,
max_running_requests: int,
device: str,
):
@@ -203,7 +203,6 @@ def _resolve_dflash_aux_hidden_state(
config.dflash_draft_num_layers = int(draft_num_layers)
config.dflash_target_layer_ids = target_layer_ids
config.dflash_draft_cell_size_per_token = _resolve_dflash_draft_cell_size(
server_args=server_args,
draft_model_config=draft_model_config,
draft_num_layers=int(draft_num_layers),
)
@@ -211,7 +210,6 @@ def _resolve_dflash_aux_hidden_state(
def _resolve_dflash_draft_cell_size(
*,
server_args: ServerArgs,
draft_model_config: ModelConfig,
draft_num_layers: int,
) -> int | None:
@@ -31,7 +31,6 @@ from sglang.srt.runtime_context import (
if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__)
@@ -98,12 +97,16 @@ class StartupWeightLoadOptions:
prefetch_num_threads: int
@classmethod
def from_server_args(
def from_published_config(
cls,
*,
server_args: ServerArgs,
is_draft_worker: bool,
) -> StartupWeightLoadOptions:
"""Everything this needs is a published leaf; nothing comes off a record.
`is_draft_worker` is the exception and travels as an argument: it is
this runner's role, not the process's configuration.
"""
cuda_graph_config = get_exec().graph.cuda_graph_config
cuda_graph_enabled = any(
getattr(cuda_graph_config, phase).backend != Backend.DISABLED
@@ -267,29 +270,27 @@ class StartupWeightLoadManager:
self._prefetch_failure_reported = False
@classmethod
def create_from_server_args(
def create_from_published_config(
cls,
*,
loader,
model_config: ModelConfig,
load_config: LoadConfig,
device_config: DeviceConfig,
server_args: ServerArgs,
is_draft_worker: bool,
) -> StartupWeightLoadManager:
"""Build a manager straight from ``ServerArgs``.
"""Build a manager from the published configuration.
Callers on the model-loading path only decide *whether* to overlap; the
knowledge of which server arguments matter, and every support rule,
stays in this module.
knowledge of which config leaves matter, and every support rule, stays
in this module.
"""
return cls.create(
loader=loader,
model_config=model_config,
load_config=load_config,
device_config=device_config,
options=StartupWeightLoadOptions.from_server_args(
server_args=server_args,
options=StartupWeightLoadOptions.from_published_config(
is_draft_worker=is_draft_worker,
),
)
@@ -259,7 +259,6 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
calculate_mla_kv_cache_dim(
model_config=model_config,
kv_cache_dtype=kv_cache_dtype,
server_args=kvc.server_args,
)
* effective_num_layers
* kv_size
@@ -465,7 +464,6 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
calculate_mla_kv_cache_dim(
model_config=model_config,
kv_cache_dtype=kv_cache_dtype,
server_args=kvc.server_args,
)
* kv_size
)
@@ -70,11 +70,10 @@ def get_batch_sizes_to_capture(
constraints and clamps to req_to_token_pool.size.
"""
server_args = model_runner.server_args
capture_bs = list(get_exec().graph.cuda_graph_config.decode.bs)
num_max_requests = model_runner.req_to_token_pool.size
mul_base = get_cuda_graph_batch_size_alignment(server_args)
mul_base = get_cuda_graph_batch_size_alignment()
# TBO splits each request's rows across two micro-batches, so the
# alignment constraint applies per request rather than per token row.
alignment_width = captured_req_width
@@ -82,7 +81,7 @@ def get_batch_sizes_to_capture(
alignment_width = 1
# pad `num_max_requests` to avoid being filtered out
num_max_requests = get_cuda_graph_max_batch_size(server_args, num_max_requests)
num_max_requests = get_cuda_graph_max_batch_size(num_max_requests)
if max(capture_bs) > num_max_requests:
# In some cases (e.g., with a small GPU or --max-running-requests), the #max-running-requests
# is very small. We add more values here to make sure we capture the maximum bs.
@@ -351,7 +351,7 @@ class BaseRunner(ABC):
dp_size=get_parallel().config.dp_size,
pp_size=get_parallel().config.pp_size,
is_encoder_decoder=mr.model_config.is_encoder_decoder,
require_mlp_tp_gather=require_mlp_tp_gather(mr.server_args),
require_mlp_tp_gather=require_mlp_tp_gather(),
seq_len_fill_value=mr.attn_backend.get_cuda_graph_seq_len_fill_value(),
encoder_len_fill_value=(
getattr(mr.model_config.hf_config, "max_source_positions", 0)
@@ -535,9 +535,9 @@ class BaseRunner(ABC):
)
# TP-gather requirements for global token metadata.
require_mlp_tp_gather_ = require_mlp_tp_gather(mr.server_args)
require_attn_tp_gather_ = require_attn_tp_gather(mr.server_args)
if require_gathered_buffer(mr.server_args):
require_mlp_tp_gather_ = require_mlp_tp_gather()
require_attn_tp_gather_ = require_attn_tp_gather()
if require_gathered_buffer():
assert require_mlp_tp_gather_ or require_attn_tp_gather_
if require_mlp_tp_gather_:
@@ -227,10 +227,10 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder
self.require_mlp_tp_gather = require_mlp_tp_gather(
model_runner.server_args
) and not self._forward_is_dp_local(model_runner)
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.require_mlp_tp_gather = (
require_mlp_tp_gather() and not self._forward_is_dp_local(model_runner)
)
self.require_attn_tp_gather = require_attn_tp_gather()
# Composite predicates derive from the instance values so the dp-local
# draft exemption above stays consistent (require_gathered_buffer ==
# mlp_tp_gather or attn_tp_gather; require_mlp_sync adds dp attention).
@@ -597,7 +597,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
draft_is_deepseek_v4,
)
return not draft_is_deepseek_v4(server_args=model_runner.server_args)
return not draft_is_deepseek_v4()
def _ragged_capture_slots(self, num_tokens: int) -> int:
if envs.SGLANG_TEST_RAGGED_VERIFY_FORCE_UNIFORM_CAPTURE.get():
@@ -113,10 +113,10 @@ class EagerRunner(BaseRunner):
# (expand_for_topk_draft) before the eager fallback.
max_bs *= get_spec().speculative_eagle_topk
# Mirror prepare_mlp_sync_batch padding so the registry holds what load_batch copies.
max_bs = get_eager_max_batch_size(sa, max_bs)
max_bs = get_eager_max_batch_size(max_bs)
prefill_ceiling = max(mr.max_total_num_tokens, max_prefill_buffer_tokens())
max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_req)
if require_mlp_sync(sa):
if require_mlp_sync():
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
max_num_token = ceil_align(max_num_token, self.attn_tp_size)
@@ -332,7 +332,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
embed_dtype=self.model_runner.dtype,
enable_mamba_track=self.mamba_track_enabled,
enable_num_token_non_padded=enable_num_token_non_padded(),
require_gathered_buffer=require_gathered_buffer(model_runner.server_args),
require_gathered_buffer=require_gathered_buffer(),
enable_prefill_cp=(
is_dsa_enable_prefill_cp() or is_mla_prefill_cp_enabled()
),
@@ -349,8 +349,8 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
self.dsa_indexers = getattr(self.model_runner, "dsa_indexers", None)
self.dp_size = get_parallel().config.dp_size
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_attn_tp_gather = require_attn_tp_gather()
# --- backend ---------------------------------------------------
# TcPiecewise resolves by running a compile pass that calls back into
@@ -602,7 +602,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
buf = self.buffer_registry.get_slot("num_token_non_padded").buffer
buf.fill_(num_tokens)
if require_gathered_buffer(self.model_runner.server_args):
if require_gathered_buffer():
local = compute_local_num_token_non_padded(
global_num_token_non_padded=buf,
num_tokens_per_dp=num_tokens,
+2 -3
View File
@@ -121,7 +121,6 @@ from sglang.srt.multimodal.mm_utils import materialize_multimodal_features
from sglang.srt.runtime_context import (
get_exec,
get_parallel,
get_server_args,
)
from sglang.srt.utils import is_blackwell_supported, is_hip, is_npu, make_layers
from sglang.srt.utils.common import (
@@ -2146,7 +2145,7 @@ class KimiK3DecoderLayer(nn.Module):
self._dp_attention = is_dp_attention_enabled()
# mlp-sync (DP attention OR MoE a2a/EP) pads extend batches to
# attn_tp multiples; attention must then run on the real rows only.
self._trim_padded_attn = require_mlp_sync(get_server_args())
self._trim_padded_attn = require_mlp_sync()
# A layer runs MoE (vs a plain dense MLP) iff it is past the dense
# prefix and on the MoE cadence — same predicate the mlp construction
# below uses.
@@ -2614,7 +2613,7 @@ class KimiK3LinearModel(nn.Module):
self.pp_group = get_pp_group()
self.dspark_layers_to_capture: Optional[list[int]] = None
self._dp_attention = is_dp_attention_enabled()
self._trim_padded_attn = require_mlp_sync(get_server_args())
self._trim_padded_attn = require_mlp_sync()
if self.pp_group.is_first_rank:
embedding_quant_config = (
@@ -40,7 +40,7 @@ from sglang.srt.multimodal.transport.cuda_ipc import (
MmItemMemoryPool,
get_mm_feature_pool_size_per_worker,
)
from sglang.srt.runtime_context import get_mm
from sglang.srt.runtime_context import get_mm, get_serving
from sglang.srt.utils import (
CLIENT_MEDIA_EXCEPTIONS,
configure_media_url_security,
@@ -215,7 +215,7 @@ class BaseMultimodalProcessor(ABC):
self.transport_mode = transport_mode
configure_media_url_security(
get_mm().allowed_media_domains,
server_args.media_url_max_file_size_mb,
get_mm().media_url_max_file_size_mb,
)
configured_mm_feature_transport = get_mm().mm_feature_transport
self.mm_feature_transport = (
@@ -227,11 +227,11 @@ class BaseMultimodalProcessor(ABC):
self.use_ipc_pool_handle_cache = (
self.use_cuda_ipc and envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.get()
)
self.image_processor_backend = server_args.image_processor_backend
if server_args.disable_fast_image_processor:
self.image_processor_backend = get_mm().image_processor_backend
if get_mm().disable_fast_image_processor:
self.image_processor_backend = "pil"
self.disable_fast_image_processor = self.image_processor_backend == "pil"
self.skip_tokenizer_init = server_args.skip_tokenizer_init
self.skip_tokenizer_init = get_serving().skip_tokenizer_init
mm_process_config = get_mm().mm_process_config
self.image_config = mm_process_config.get("image", {})
@@ -241,19 +241,19 @@ class BaseMultimodalProcessor(ABC):
# Each tokenizer worker is a separate process with its own CPU cache.
# Split the requested service-wide budget so increasing worker count
# does not silently multiply host-memory usage.
requested_cache_mb = self.server_args.mm_preprocess_cache_size_mb
requested_cache_mb = get_mm().mm_preprocess_cache_size_mb
total_cache_mb = (
self.auto_mm_preprocess_cache_size_mb
if requested_cache_mb is None
else requested_cache_mb
)
tokenizer_worker_num = max(int(self.server_args.tokenizer_worker_num), 1)
tokenizer_worker_num = max(int(get_serving().tokenizer_worker_num), 1)
worker_cache_bytes = total_cache_mb * 1024 * 1024 // tokenizer_worker_num
self.mm_preprocess_cache = MultimodalPreprocessCache(
max_size_bytes=worker_cache_bytes,
max_entries=8192,
)
self.trust_mm_content_hashes = bool(self.server_args.trust_mm_content_hashes)
self.trust_mm_content_hashes = bool(get_mm().trust_mm_content_hashes)
# The fingerprint is needed only to build artifact keys. Avoid inspecting
# processor state when this processor will never retain artifacts.
self.processor_fingerprint = (
@@ -288,7 +288,7 @@ class BaseMultimodalProcessor(ABC):
# FIXME: not accurate, model and image specific
self.NUM_TOKEN_PER_FRAME = 330
requested_mm_io_worker_num = self.server_args.mm_io_worker_num
requested_mm_io_worker_num = get_mm().mm_io_worker_num
env_mm_io_worker_num = os.environ.get("SGLANG_IO_WORKERS")
if requested_mm_io_worker_num:
self.mm_io_worker_num = requested_mm_io_worker_num
@@ -310,7 +310,7 @@ class BaseMultimodalProcessor(ABC):
io_worker_source,
)
skip_mm_pool = kwargs.get("skip_mm_pool", False)
requested_mm_processor_worker_num = self.server_args.mm_processor_worker_num
requested_mm_processor_worker_num = get_mm().mm_processor_worker_num
self.mm_processor_worker_num = (
1
if skip_mm_pool
@@ -402,7 +402,7 @@ class BaseMultimodalProcessor(ABC):
# SGLANG_MM_FEATURE_CACHE_MB is the total pool budget across all
# tokenizer workers. Each worker gets an equal share so that adding
# workers doesn't multiply the GPU-side footprint.
worker_num = self.server_args.tokenizer_worker_num
worker_num = get_serving().tokenizer_worker_num
per_worker_pool_size = get_mm_feature_pool_size_per_worker(
MM_FEATURE_CACHE_SIZE, worker_num
)
@@ -30,7 +30,7 @@ def get_dspark_sample_from_anchor(draft_hf_config: Any) -> bool:
return bool(_cfg_get(draft_hf_config, "sample_from_anchor", True))
def draft_is_deepseek_v4(*, server_args: ServerArgs) -> bool:
def draft_is_deepseek_v4() -> bool:
from sglang.srt.configs.model_config import is_deepseek_v4
from sglang.srt.utils.hf_transformers_utils import get_config
@@ -172,7 +172,7 @@ class DSparkVerifyPlanner:
and is_dp_attention_enabled()
and get_parallel().attn_tp_size == 1
and get_parallel().attn_cp_size == 1
and require_mlp_tp_gather(self.server_args)
and require_mlp_tp_gather()
and not get_schedule().disable_overlap_schedule
and not get_spec().speculative_skip_dp_mlp_sync
and get_disagg().disaggregation_mode == "null"
@@ -103,7 +103,7 @@ class DSparkWorkerV2(BaseSpecWorker):
self.page_size = get_schedule().page_size
self.device = target_worker.device
self._draft_is_moe = draft_is_deepseek_v4(server_args=server_args)
self._draft_is_moe = draft_is_deepseek_v4()
self._draft_dp_context_enabled = (
get_parallel().config.enable_dp_attention and not self._draft_is_moe
)
@@ -116,10 +116,10 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
self.pp_size = get_parallel().config.pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather()
self.enable_profile_cuda_graph = (
model_runner.server_args.enable_profile_cuda_graph
)
@@ -102,10 +102,10 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
self.pp_size = get_parallel().config.pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather()
self.enable_profile_cuda_graph = (
model_runner.server_args.enable_profile_cuda_graph
)
@@ -93,10 +93,10 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
self.device_module = torch.get_device_module(self.device)
self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather()
self.tp_size = self.model_runner.ps.tp_size
self.attn_dp_size = self.model_runner.ps.attn_dp_size
self.pp_size = get_parallel().config.pp_size
@@ -158,10 +158,10 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
self.pp_size = get_parallel().config.pp_size
self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather()
self.enable_pdmux = model_runner.server_args.enable_pdmux
self.speculative_num_steps = get_spec().speculative_num_steps
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
@@ -846,7 +846,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
self.cuda_graph_runner_for_draft_extend.prune_draft_extend_logits
)
else:
prune_logits = not require_gathered_buffer(self.server_args)
prune_logits = not require_gathered_buffer()
if prune_logits:
forward_batch.spec_info.select_index = select_index
# Left unmarked on every platform: each de-tied runner has its own
+13 -15
View File
@@ -105,7 +105,7 @@ from sglang.srt.runtime_context import (
from sglang.srt.utils.video_decoder import _BACKEND, VideoDecoderWrapper
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
pass
logger = logging.getLogger(__name__)
torch_release = pkg_version.parse(torch.__version__).release
@@ -3719,7 +3719,7 @@ class Withable(Generic[T]):
self._value = None
def require_mlp_tp_gather(server_args: ServerArgs):
def require_mlp_tp_gather():
"""
Check if the input of MLP is obtained by all-gather rather than all-reduce. This only happens when each MLP TP group contains multiple attention DP groups.
"""
@@ -3763,7 +3763,7 @@ def require_mlp_tp_gather(server_args: ServerArgs):
return False
def require_attn_tp_gather(server_args: ServerArgs):
def require_attn_tp_gather():
"""
Check if the input of attention is scattered.
"""
@@ -3790,35 +3790,33 @@ def require_attn_tp_gather(server_args: ServerArgs):
return False
def require_gathered_buffer(server_args: ServerArgs):
return require_mlp_tp_gather(server_args) or require_attn_tp_gather(server_args)
def require_gathered_buffer():
return require_mlp_tp_gather() or require_attn_tp_gather()
def require_mlp_sync(server_args: ServerArgs):
def require_mlp_sync():
from sglang.srt.runtime_context import get_parallel
return get_parallel().config.enable_dp_attention or require_gathered_buffer(
server_args
)
return get_parallel().config.enable_dp_attention or require_gathered_buffer()
def get_cuda_graph_batch_size_alignment(server_args: ServerArgs) -> int:
def get_cuda_graph_batch_size_alignment() -> int:
alignment = 1
if get_exec().overlap.enable_two_batch_overlap:
alignment *= 2
if require_gathered_buffer(server_args):
if require_gathered_buffer():
alignment *= get_parallel().attn_tp_size
if alignment % get_parallel().attn_cp_size != 0:
alignment *= get_parallel().attn_cp_size
return alignment
def get_cuda_graph_max_batch_size(server_args: ServerArgs, max_batch_size: int) -> int:
return ceil_align(max_batch_size, get_cuda_graph_batch_size_alignment(server_args))
def get_cuda_graph_max_batch_size(max_batch_size: int) -> int:
return ceil_align(max_batch_size, get_cuda_graph_batch_size_alignment())
def get_eager_max_batch_size(server_args: ServerArgs, max_batch_size: int) -> int:
if not require_mlp_sync(server_args):
def get_eager_max_batch_size(max_batch_size: int) -> int:
if not require_mlp_sync():
return max_batch_size
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
@@ -161,7 +161,7 @@ def _contains_tensor_container(value) -> bool:
)
def get_vmm_feature_consumer_count(server_args) -> int:
def get_vmm_feature_consumer_count() -> int:
if get_parallel().config.enable_dp_attention:
return get_parallel().config.tp_size // get_parallel().config.dp_size
return get_parallel().config.tp_size
@@ -947,7 +947,7 @@ class CudaVmmFeatureTransport:
memory_size=per_worker_pool_size,
recycle_interval=MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL,
base_gpu_id=server_args.base_gpu_id,
consumer_count=get_vmm_feature_consumer_count(server_args),
consumer_count=get_vmm_feature_consumer_count(),
allow_posix_fallback=server_args.nnodes == 1,
)
+1 -2
View File
@@ -17,7 +17,6 @@ from sglang.srt.runtime_context import (
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 (
HostSharedMemoryManager,
@@ -66,7 +65,7 @@ def set_offloader(instance: BaseOffloader):
_instance = instance
def create_offloader_from_server_args(server_args: ServerArgs, dp_rank: int):
def create_offloader(dp_rank: int):
if get_exec().offload.cpu_offload_gb > 0:
return OffloaderV1(
cpu_offload_max_bytes=int(get_exec().offload.cpu_offload_gb * 1024**3)
@@ -57,7 +57,6 @@ class TestDisaggregationServerWarmup(unittest.IsolatedAsyncioTestCase):
with patch("sglang.srt.entrypoints.http_server.aiohttp.ClientSession", Session):
status_codes = await _send_disaggregation_warmup_requests(
server_args=server_args,
url="http://localhost:30000",
headers={"Authorization": "Bearer token"},
ssl_verify=False,
@@ -1200,10 +1200,6 @@ class TestMlxOverlapScheduler(unittest.TestCase):
disaggregation_mode=None,
enable_overlap=False,
enable_overlap_mlx=False,
server_args=SimpleNamespace(
disaggregation_decode_enable_offload_kvcache=False,
enable_hisparse=False,
),
model_config=None,
token_to_kv_pool_allocator=None,
tree_cache=tree_cache,
@@ -31,10 +31,6 @@ def _make_processor(case, server_mode: str = "full") -> SchedulerBatchResultProc
disaggregation_mode=None,
enable_overlap=False,
enable_overlap_mlx=False,
server_args=SimpleNamespace(
enable_metrics=False,
enable_hisparse=False,
),
model_config=SimpleNamespace(think_end_ids=None),
token_to_kv_pool_allocator=Mock(),
tree_cache=None,
@@ -60,7 +60,6 @@ def _make_processor() -> SchedulerBatchResultProcessor:
disaggregation_mode=None,
enable_overlap=True,
enable_overlap_mlx=False,
server_args=SimpleNamespace(),
model_config=SimpleNamespace(think_end_ids=None),
token_to_kv_pool_allocator=MagicMock(),
tree_cache=SimpleNamespace(page_size=TRACK_INTERVAL),
@@ -65,7 +65,6 @@ def _make_processor() -> SchedulerBatchResultProcessor:
disaggregation_mode=None,
enable_overlap=False,
enable_overlap_mlx=False,
server_args=SimpleNamespace(enable_metrics=False),
model_config=SimpleNamespace(think_end_ids=None),
token_to_kv_pool_allocator=None,
tree_cache=None,
@@ -81,21 +81,20 @@ class TestBaseProcessorConfigExtraction(CustomTestCase):
BaseMultimodalProcessor,
)
# The multimodal config comes from the bags.
override = get_context().override_server_args(
mm_process_config=mm_process_config,
allowed_media_domains=[],
mm_processor_worker_num=mm_processor_worker_num,
mm_io_worker_num=mm_io_worker_num,
mm_preprocess_cache_size_mb=None,
tokenizer_worker_num=1,
trust_mm_content_hashes=False,
media_url_max_file_size_mb=64,
)
override.install()
self.addCleanup(override.restore)
server_args = MagicMock()
server_args.mm_processor_worker_num = mm_processor_worker_num
server_args.mm_io_worker_num = mm_io_worker_num
server_args.mm_preprocess_cache_size_mb = None
server_args.tokenizer_worker_num = 1
server_args.trust_mm_content_hashes = False
server_args.media_url_max_file_size_mb = 64
hf_config = MagicMock()
mock_hf_processor = MagicMock()
@@ -209,8 +209,7 @@ class TestStartupWeightLoadSelector(CustomTestCase):
# The parallel sizes come from the bags, so the config has to be published.
publish(server_args, role="test")
self.addCleanup(reset_context)
options = StartupWeightLoadOptions.from_server_args(
server_args=server_args,
options = StartupWeightLoadOptions.from_published_config(
is_draft_worker=False,
)
+8 -6
View File
@@ -648,6 +648,14 @@ def _k3_preprocess_config(
)
def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls):
server_args = SimpleNamespace(
base_gpu_id=0,
rl_on_policy_target=None,
tp_size=1,
)
with get_context().override_server_args(
mm_feature_transport="cpu",
mm_process_config={},
allowed_media_domains=[],
image_processor_backend="auto",
disable_fast_image_processor=False,
skip_tokenizer_init=False,
@@ -656,13 +664,7 @@ def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls
tokenizer_worker_num=1,
mm_preprocess_cache_size_mb=0,
trust_mm_content_hashes=False,
base_gpu_id=0,
rl_on_policy_target=None,
media_url_max_file_size_mb=64,
)
# The multimodal config comes from the bags.
with get_context().override_server_args(
mm_feature_transport="cpu", mm_process_config={}, allowed_media_domains=[]
):
processor = processor_cls(
hf_config=SimpleNamespace(media_placeholder_token_id=42),
@@ -89,14 +89,15 @@ def make_processor(case, config, image_processor_cls=None):
allowed_media_domains=[],
media_url_max_file_size_mb=64,
)
# The processor reads its media policy, transport and per-modality limits
# from the mm bag, so the fixture publishes before building it.
# Left at the default backend, the fast image processor sends the tensor to
# `cuda:<base_gpu_id>`, which a CPU-only host cannot do.
publish(
ServerArgs(
model_path="dummy",
mm_feature_transport=server_args.mm_feature_transport,
mm_process_config=server_args.mm_process_config,
allowed_media_domains=server_args.allowed_media_domains,
disable_fast_image_processor=server_args.disable_fast_image_processor,
),
role="tokenizer",
)
@@ -134,14 +134,13 @@ class TestDraftPerRunnerConfig(CustomTestCase):
def test_an_unresolved_draft_falls_back_to_the_config_field(self):
"""The v2 workers pass no backend: --speculative-draft-attention-backend."""
server_args = self._seed(
self._seed(
attention_backend="fa3", speculative_draft_attention_backend="triton"
)
def effective(*, is_draft_worker, passed=None):
return resolve_draft_attention_backend(
draft_attention_backend=passed,
server_args=server_args,
is_draft_worker=is_draft_worker,
)
@@ -137,7 +137,6 @@ _PASSED = frozenset({"model_path", "device", "random_seed"})
_EXPOSED = {
("dllm/config.py", "max_running_requests"),
("dllm/config.py", "model_path"),
("multimodal/processors/base_processor.py", "image_processor_backend"),
("speculative/spec_registry.py", "disable_overlap_schedule"),
("disaggregation/encoder/server.py", "model_loader_extra_config"),
("layers/moe/utils.py", "deepep_mode"),
@@ -165,8 +164,6 @@ _EXPOSED = {
("entrypoints/engine.py", "enable_symm_mem"),
("entrypoints/engine.py", "reasoning_parser"),
("entrypoints/engine.py", "tool_call_parser"),
("eplb/eplb_manager.py", "ep_dispatch_algorithm"),
("eplb/eplb_manager.py", "expert_distribution_recorder_buffer_size"),
("layers/cp/base.py", "attn_cp_size"),
("layers/cp/base.py", "cp_strategy"),
("layers/cp/base.py", "enable_prefill_cp"),
@@ -178,11 +175,9 @@ _EXPOSED = {
("layers/moe/utils.py", "moe_runner_backend"),
("layers/moe/utils.py", "quantization"),
("layers/moe/utils.py", "speculative_moe_runner_backend"),
("lora/lora_manager.py", "enable_lora_overlap_loading"),
("lora/marlin_lora_temp/policy.py", "lora_paths"),
("model_loader/expert_pack_runtime.py", "model_path"),
("model_loader/expert_pack_runtime.py", "tokenizer_path"),
("multimodal/processors/base_processor.py", "image_processor_backend"),
("parser/template_detection.py", "model_path"),
("speculative/adaptive_spec_params.py", "speculative_algorithm"),
("speculative/adaptive_spec_params.py", "speculative_eagle_topk"),