diff --git a/python/sglang/benchmark/one_batch.py b/python/sglang/benchmark/one_batch.py index ea1863500..d6d12a499 100644 --- a/python/sglang/benchmark/one_batch.py +++ b/python/sglang/benchmark/one_batch.py @@ -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(), ) diff --git a/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py b/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py index 0eaad79a9..8567fe662 100644 --- a/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py +++ b/python/sglang/srt/distributed/device_communicators/mooncake_transfer_engine.py @@ -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) diff --git a/python/sglang/srt/elastic_ep/elastic_ep.py b/python/sglang/srt/elastic_ep/elastic_ep.py index 1dc43078a..73a3489a2 100644 --- a/python/sglang/srt/elastic_ep/elastic_ep.py +++ b/python/sglang/srt/elastic_ep/elastic_ep.py @@ -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 diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index 69575931f..13817d468 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -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, diff --git a/python/sglang/srt/eplb/eplb_manager.py b/python/sglang/srt/eplb/eplb_manager.py index 00738b40f..b591a9cd0 100644 --- a/python/sglang/srt/eplb/eplb_manager.py +++ b/python/sglang/srt/eplb/eplb_manager.py @@ -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 diff --git a/python/sglang/srt/eplb/expert_distribution.py b/python/sglang/srt/eplb/expert_distribution.py index 693d69264..ec6133956 100644 --- a/python/sglang/srt/eplb/expert_distribution.py +++ b/python/sglang/srt/eplb/expert_distribution.py @@ -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, diff --git a/python/sglang/srt/eplb/expert_location.py b/python/sglang/srt/eplb/expert_location.py index 19624be81..9a10e2693 100644 --- a/python/sglang/srt/eplb/expert_location.py +++ b/python/sglang/srt/eplb/expert_location.py @@ -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( diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index 0de53ce64..6d37cbb7c 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -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( ( diff --git a/python/sglang/srt/layers/attention/linear/gdn_backend.py b/python/sglang/srt/layers/attention/linear/gdn_backend.py index e8033a62d..40669014f 100644 --- a/python/sglang/srt/layers/attention/linear/gdn_backend.py +++ b/python/sglang/srt/layers/attention/linear/gdn_backend.py @@ -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, ) ) diff --git a/python/sglang/srt/layers/attention/linear/kda_backend.py b/python/sglang/srt/layers/attention/linear/kda_backend.py index a0dbba998..406d1c591 100644 --- a/python/sglang/srt/layers/attention/linear/kda_backend.py +++ b/python/sglang/srt/layers/attention/linear/kda_backend.py @@ -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, ) ) diff --git a/python/sglang/srt/layers/attention/linear/utils.py b/python/sglang/srt/layers/attention/linear/utils.py index 4e63068ec..c8e469e54 100644 --- a/python/sglang/srt/layers/attention/linear/utils.py +++ b/python/sglang/srt/layers/attention/linear/utils.py @@ -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( diff --git a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py index 1e2db45ff..fe2b659f4 100644 --- a/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py +++ b/python/sglang/srt/layers/deep_gemm_wrapper/compile_utils.py @@ -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 diff --git a/python/sglang/srt/lora/lora_manager.py b/python/sglang/srt/lora/lora_manager.py index ba07718a0..f83cededf 100644 --- a/python/sglang/srt/lora/lora_manager.py +++ b/python/sglang/srt/lora/lora_manager.py @@ -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, diff --git a/python/sglang/srt/managers/disagg_service.py b/python/sglang/srt/managers/disagg_service.py index be82b892b..d710c601a 100644 --- a/python/sglang/srt/managers/disagg_service.py +++ b/python/sglang/srt/managers/disagg_service.py @@ -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 diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py index 3ce1494d7..03772d725 100644 --- a/python/sglang/srt/managers/multi_tokenizer_mixin.py +++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py @@ -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() diff --git a/python/sglang/srt/managers/overlap_utils.py b/python/sglang/srt/managers/overlap_utils.py index ce78fa8da..21db5f396 100644 --- a/python/sglang/srt/managers/overlap_utils.py +++ b/python/sglang/srt/managers/overlap_utils.py @@ -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, diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 56b2200ca..24b2689c9 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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, diff --git a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py index 42245d088..2deb7bb92 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -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, diff --git a/python/sglang/srt/managers/scheduler_components/dp_attn.py b/python/sglang/srt/managers/scheduler_components/dp_attn.py index b443da627..595c55107 100644 --- a/python/sglang/srt/managers/scheduler_components/dp_attn.py +++ b/python/sglang/srt/managers/scheduler_components/dp_attn.py @@ -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, diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index c7713a0e4..0860a5e9f 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -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] diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 221877577..ee97103c8 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -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 diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index b1900ea65..e29de3a01 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -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, diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index 068fda733..ed0215bff 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -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 diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index be0a80da3..7a1b716bf 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -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 ) diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index b92d5ce07..85e8157e0 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -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] diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index cd3eb88ed..f7947fb61 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -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(), diff --git a/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py b/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py index 4051d5a25..921817903 100644 --- a/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py +++ b/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py @@ -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() diff --git a/python/sglang/srt/model_executor/model_runner_components/ngram_embedding_manager.py b/python/sglang/srt/model_executor/model_runner_components/ngram_embedding_manager.py index 4d034f781..d4054cce1 100644 --- a/python/sglang/srt/model_executor/model_runner_components/ngram_embedding_manager.py +++ b/python/sglang/srt/model_executor/model_runner_components/ngram_embedding_manager.py @@ -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, ): diff --git a/python/sglang/srt/model_executor/model_runner_components/spec_aux_hidden_state.py b/python/sglang/srt/model_executor/model_runner_components/spec_aux_hidden_state.py index 3e26fd8ad..61a2a573b 100644 --- a/python/sglang/srt/model_executor/model_runner_components/spec_aux_hidden_state.py +++ b/python/sglang/srt/model_executor/model_runner_components/spec_aux_hidden_state.py @@ -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: diff --git a/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py b/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py index 202e602c5..3633856a5 100644 --- a/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py +++ b/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py @@ -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, ), ) diff --git a/python/sglang/srt/model_executor/pool_configurator.py b/python/sglang/srt/model_executor/pool_configurator.py index e38553218..3f40b902b 100644 --- a/python/sglang/srt/model_executor/pool_configurator.py +++ b/python/sglang/srt/model_executor/pool_configurator.py @@ -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 ) diff --git a/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py index f58717e19..880ea3136 100644 --- a/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/base_cuda_graph_runner.py @@ -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. diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index 20a024aad..1ebc26310 100644 --- a/python/sglang/srt/model_executor/runner/base_runner.py +++ b/python/sglang/srt/model_executor/runner/base_runner.py @@ -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_: diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index 65323e63b..a954cd181 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -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(): diff --git a/python/sglang/srt/model_executor/runner/eager_runner.py b/python/sglang/srt/model_executor/runner/eager_runner.py index 6476c3a87..83caa44d6 100644 --- a/python/sglang/srt/model_executor/runner/eager_runner.py +++ b/python/sglang/srt/model_executor/runner/eager_runner.py @@ -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) diff --git a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py index 96d57bfca..b10ca8f09 100644 --- a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py @@ -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, diff --git a/python/sglang/srt/models/kimi_k3.py b/python/sglang/srt/models/kimi_k3.py index a7a59d87a..7c15a3ff1 100644 --- a/python/sglang/srt/models/kimi_k3.py +++ b/python/sglang/srt/models/kimi_k3.py @@ -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 = ( diff --git a/python/sglang/srt/multimodal/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index 2ba4e508c..ff3f60a8c 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -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 ) diff --git a/python/sglang/srt/speculative/dspark_components/dspark_config.py b/python/sglang/srt/speculative/dspark_components/dspark_config.py index 068e3ff7c..88bc14405 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_config.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_config.py @@ -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 diff --git a/python/sglang/srt/speculative/dspark_components/dspark_planner.py b/python/sglang/srt/speculative/dspark_components/dspark_planner.py index 1bdfa7c4a..edde1bbf1 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_planner.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_planner.py @@ -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" diff --git a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py index 53fe887c2..08c8d7896 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py @@ -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 ) diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py index dc21efc40..ab2f5ce27 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -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 ) diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index 3c4c63d7f..2ee2f7c28 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -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 ) diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py index 3c7352035..43914af2a 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py @@ -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 diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index aa15c97da..3dc1c6739 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -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 diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index ea75a3fd2..7807e8960 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -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 diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index d1fd11a02..c7972f925 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -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 diff --git a/python/sglang/srt/utils/cuda_vmm_transport_utils.py b/python/sglang/srt/utils/cuda_vmm_transport_utils.py index 8525e5812..d115701bf 100644 --- a/python/sglang/srt/utils/cuda_vmm_transport_utils.py +++ b/python/sglang/srt/utils/cuda_vmm_transport_utils.py @@ -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, ) diff --git a/python/sglang/srt/utils/offloader.py b/python/sglang/srt/utils/offloader.py index 63c4a8954..e38b0f2a1 100644 --- a/python/sglang/srt/utils/offloader.py +++ b/python/sglang/srt/utils/offloader.py @@ -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) diff --git a/test/registered/unit/entrypoints/test_http_server_warmup.py b/test/registered/unit/entrypoints/test_http_server_warmup.py index e4391e456..a74d02b85 100644 --- a/test/registered/unit/entrypoints/test_http_server_warmup.py +++ b/test/registered/unit/entrypoints/test_http_server_warmup.py @@ -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, diff --git a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py index 14dc72d4e..662da5b09 100644 --- a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py +++ b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py @@ -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, diff --git a/test/registered/unit/managers/test_batch_result_processor_hidden_states.py b/test/registered/unit/managers/test_batch_result_processor_hidden_states.py index 82d1beee5..f606842e0 100644 --- a/test/registered/unit/managers/test_batch_result_processor_hidden_states.py +++ b/test/registered/unit/managers/test_batch_result_processor_hidden_states.py @@ -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, diff --git a/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py b/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py index 2cc5f4d46..0c618ef04 100644 --- a/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py +++ b/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py @@ -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), diff --git a/test/registered/unit/managers/test_batch_result_processor_spec_grammar.py b/test/registered/unit/managers/test_batch_result_processor_spec_grammar.py index 144d2b13a..baee3db67 100644 --- a/test/registered/unit/managers/test_batch_result_processor_spec_grammar.py +++ b/test/registered/unit/managers/test_batch_result_processor_spec_grammar.py @@ -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, diff --git a/test/registered/unit/managers/test_mm_process_config.py b/test/registered/unit/managers/test_mm_process_config.py index 04a9e82f1..167837e5e 100644 --- a/test/registered/unit/managers/test_mm_process_config.py +++ b/test/registered/unit/managers/test_mm_process_config.py @@ -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() diff --git a/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py b/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py index ba5a924dc..dd183c379 100644 --- a/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py +++ b/test/registered/unit/model_executor/model_runner_components/test_startup_weight_load.py @@ -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, ) diff --git a/test/registered/unit/models/test_kimi_k25.py b/test/registered/unit/models/test_kimi_k25.py index dbb2046b2..3e4a11f56 100644 --- a/test/registered/unit/models/test_kimi_k25.py +++ b/test/registered/unit/models/test_kimi_k25.py @@ -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), diff --git a/test/registered/unit/multimodal/rust/qwen/_fixtures.py b/test/registered/unit/multimodal/rust/qwen/_fixtures.py index 28534a362..4873ea3d7 100644 --- a/test/registered/unit/multimodal/rust/qwen/_fixtures.py +++ b/test/registered/unit/multimodal/rust/qwen/_fixtures.py @@ -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:`, 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", ) diff --git a/test/registered/unit/spec/test_draft_per_runner_config.py b/test/registered/unit/spec/test_draft_per_runner_config.py index 274a71a51..29c7ce928 100644 --- a/test/registered/unit/spec/test_draft_per_runner_config.py +++ b/test/registered/unit/spec/test_draft_per_runner_config.py @@ -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, ) diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index aa9af73ce..3c366806f 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -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"),