diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index 3d7e40662..776095e5c 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -261,19 +261,14 @@ def mamba_extra_buffer_of(cfg: Any) -> bool: def declare_load_time_override(source: str, declared: Dict[str, Any]) -> None: """Declare a load-time resolved field (model-file config overrides, - weight-resolved dtypes) on the published ``server_args``: resolution has - already materialized, so the declaration writes through, joining the - declaration stash for provenance and republish consistency.""" + weight-resolved dtypes): validated against the resolvable whitelist, then + written to the config bags via ``get_context().override``; ``server_args`` + stays the pristine startup record.""" from sglang.srt.runtime_context import get_context - server_args = get_context().server_args - validate_declarations(server_args, [(source, dict(declared))]) - override = getattr(server_args, "override", None) - if override is not None: - override(source, **declared) - else: - # Config-shaped fixtures without the mutation entry point. - _apply_fields(server_args, declared) + context = get_context() + validate_declarations(context.server_args, [(source, dict(declared))]) + context.override(source, **declared) def collect_model_override_declarations( diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index 2d447081c..772d9af0e 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -40,7 +40,7 @@ from sglang.srt.model_executor.forward_batch_info import ( compute_position, ) from sglang.srt.model_executor.forward_context import get_attn_backend -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_device, get_parallel, get_server_args from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip @@ -184,7 +184,7 @@ def _update_device_and_sum_field_from_cpu_field( cpu_value if isinstance(cpu_value, torch.Tensor) else torch.tensor(cpu_value, dtype=old_device_value.dtype) - ).to(device=get_server_args().device, non_blocking=True) + ).to(device=get_device().device, non_blocking=True) setattr(batch, device_field, new_device_value) if sum_field is not None: @@ -336,7 +336,7 @@ def compute_split_indices_for_cuda_graph_replay( class TboCudaGraphRunnerPlugin: def __init__(self): self._tbo_children_num_token_non_padded = torch.zeros( - (2,), dtype=torch.int32, device=get_server_args().device + (2,), dtype=torch.int32, device=get_device().device ) def capture_one_batch_size(self, batch: ForwardBatch, num_tokens: int): @@ -835,7 +835,7 @@ class TboForwardBatchPreparer: value_a = min(tbo_split_token_index, num_token_non_padded) value_b = max(0, num_token_non_padded - tbo_split_token_index) return torch.tensor([value_a, value_b], dtype=torch.int32).to( - device=get_server_args().device, non_blocking=True + device=get_device().device, non_blocking=True ) @classmethod diff --git a/python/sglang/srt/configs/inkling.py b/python/sglang/srt/configs/inkling.py index b4f24592e..c963ca9e9 100644 --- a/python/sglang/srt/configs/inkling.py +++ b/python/sglang/srt/configs/inkling.py @@ -8,6 +8,7 @@ from transformers import CONFIG_MAPPING from transformers.configuration_utils import PretrainedConfig from sglang.srt.configs.mamba_utils import BaseLinearStateParams +from sglang.srt.runtime_context import get_exec class InklingModelConfig(PretrainedConfig): @@ -224,9 +225,8 @@ class InklingModelConfig(PretrainedConfig): self.swa_num_key_value_heads, self.swa_head_dim ) stream_dim = self.hidden_size - from sglang.srt.runtime_context import get_server_args - if get_server_args().enable_scattered_sconv: + if get_exec().comm.enable_scattered_sconv: # Scattered sconv: the attn/mlp output sconvs run on the [T, H/P] # hidden shard, so their conv-state caches shard with them. assert ( diff --git a/python/sglang/srt/debug_utils/dumper.py b/python/sglang/srt/debug_utils/dumper.py index b73f75b34..e6edb3f1f 100644 --- a/python/sglang/srt/debug_utils/dumper.py +++ b/python/sglang/srt/debug_utils/dumper.py @@ -23,6 +23,7 @@ import torch.distributed as dist import zmq from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle +from sglang.srt.runtime_context import get_serving # -------------------------------------- config base ------------------------------------------ @@ -1798,7 +1799,7 @@ class _SGLangPlugin(_FrameworkPlugin): if args is None: return None - return args.tokenizer_path + return get_serving().tokenizer_path except Exception: return None diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index 9ce7b1c7e..f88eda509 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -41,6 +41,7 @@ class KVArgs: kv_data_lens: List[int] kv_item_lens: List[int] kv_layer_ids: List[int] + kv_cache_dtype_str: str aux_data_ptrs: List[int] aux_data_lens: List[int] aux_item_lens: List[int] diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index f5d6982a4..43b9bbbbb 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -36,7 +36,7 @@ from sglang.srt.layers.dp_attention import ( get_attention_dp_rank, get_attention_dp_size, ) -from sglang.srt.runtime_context import get_model, get_parallel +from sglang.srt.runtime_context import get_parallel, get_serving from sglang.srt.server_args import ServerArgs from sglang.srt.utils.network import ( NetworkAddress, @@ -148,6 +148,7 @@ class CommonKVManager(BaseKVManager): is_mla_backend: Optional[bool] = False, ): self.kv_args = args + self.kv_cache_dtype_str = args.kv_cache_dtype_str self.kv_item_lens_sum = sum(args.kv_item_lens) self.state_item_lens_sum = sum(x for comp in args.state_item_lens for x in comp) self.is_mla_backend = is_mla_backend @@ -533,11 +534,11 @@ class CommonKVManager(BaseKVManager): if ( info.kv_cache_dtype is not None - and info.kv_cache_dtype != get_model().kv_cache_dtype + and info.kv_cache_dtype != self.kv_cache_dtype_str ): raise RuntimeError( f"KV cache dtype mismatch: prefill server has kv_cache_dtype={info.kv_cache_dtype}, " - f"but decode server has kv_cache_dtype={get_model().kv_cache_dtype}. " + f"but decode server has kv_cache_dtype={self.kv_cache_dtype_str}. " f"Both servers must use the same --kv-cache-dtype value." ) @@ -701,7 +702,7 @@ class CommonKVManager(BaseKVManager): "rank_ip": self.local_ip, "rank_port": self.rank_port, "page_size": self.kv_args.page_size, - "kv_cache_dtype": get_model().kv_cache_dtype, + "kv_cache_dtype": self.kv_cache_dtype_str, "load_balance_method": self.server_args.load_balance_method, "enable_dsa_cache_layer_split": getattr( self.server_args, "enable_dsa_cache_layer_split", False @@ -709,7 +710,7 @@ class CommonKVManager(BaseKVManager): # Self-register the HTTP API port so the decode can derive the PD # retract rebootstrap /generate URL from bootstrap info instead of a # router-injected pd_rebootstrap_prefill_url. - "prefill_http_port": self.server_args.port, + "prefill_http_port": get_serving().port, } max_retries, initial_delay, max_delay = 5, 1.0, 30.0 diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 8bd0bb3b3..69d07fb03 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -87,7 +87,7 @@ from sglang.srt.observability.req_time_stats import ( set_schedule_time_batch, set_time_batch, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_disagg, get_parallel from sglang.srt.utils import get_num_new_pages, is_npu from sglang.srt.utils.network import NetworkAddress from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method @@ -423,6 +423,9 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin): kv_args.pp_rank = self.pp_rank kv_args.system_dp_rank = self.scheduler.ps.dp_rank + kv_args.kv_cache_dtype_str = ( + self.scheduler.tp_worker.model_runner.kv_cache_dtype_str + ) transfer_kv_pool = ( self.scheduler.hisparse_coordinator.mem_pool_host if self.scheduler.enable_hisparse @@ -2244,7 +2247,7 @@ class SchedulerDisaggregationDecodeMixin: # Decode-radix path: new requests already matched in # `pop_preallocated`. Retracted requests reset `last_node`, # so re-match only when that state is missing. - if self.server_args.disaggregation_decode_enable_radix_cache: + if get_disagg().disaggregation_decode_enable_radix_cache: tree_cache = self.tree_cache if req.last_node is None else None else: tree_cache = self.tree_cache @@ -2284,7 +2287,7 @@ class SchedulerDisaggregationDecodeMixin: if self.enable_decode_hicache: self.tree_cache.check_hicache_events() - if self.server_args.disaggregation_decode_enable_offload_kvcache: + if get_disagg().disaggregation_decode_enable_offload_kvcache: self.decode_offload_manager.check_offload_progress() # try to resume retracted requests if there are enough space for another `num_reserved_decode_tokens` decode steps @@ -2296,9 +2299,7 @@ class SchedulerDisaggregationDecodeMixin: if not hasattr(self, "polling_count"): self.polling_count = 0 - self.polling_interval = ( - self.server_args.disaggregation_decode_polling_interval - ) + self.polling_interval = get_disagg().disaggregation_decode_polling_interval self.polling_count = (self.polling_count + 1) % self.polling_interval diff --git a/python/sglang/srt/disaggregation/encode_grpc_server.py b/python/sglang/srt/disaggregation/encode_grpc_server.py index 5abb27473..3516f56d4 100644 --- a/python/sglang/srt/disaggregation/encode_grpc_server.py +++ b/python/sglang/srt/disaggregation/encode_grpc_server.py @@ -28,6 +28,7 @@ from sglang.srt.disaggregation.encode_server import ( ) from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle from sglang.srt.managers.schedule_batch import Modality +from sglang.srt.runtime_context import get_disagg from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.utils import random_uuid from sglang.srt.utils.network import NetworkAddress, get_zmq_socket @@ -117,13 +118,13 @@ class SGLangEncoderServer(SGLangEncoderServicer): context.set_details(error_msg) return sglang_encoder_pb2.EncodeResponse() - if self.server_args.encoder_transfer_backend == "mooncake": + if get_disagg().encoder_transfer_backend == "mooncake": return sglang_encoder_pb2.EncodeResponse( embedding_size=nbytes, embedding_len=embedding_len, embedding_dim=embedding_dim, ) - elif self.server_args.encoder_transfer_backend == "zmq_to_scheduler": + elif get_disagg().encoder_transfer_backend == "zmq_to_scheduler": embedding_ports = list(request.embedding_port) logger.info(f"embedding_port = {embedding_ports}") if not embedding_ports: @@ -141,7 +142,7 @@ class SGLangEncoderServer(SGLangEncoderServicer): await asyncio.gather(*tasks) self.encoder.embedding_to_send.pop(request.req_id, None) return sglang_encoder_pb2.EncodeResponse() - elif self.server_args.encoder_transfer_backend == "zmq_to_tokenizer": + elif get_disagg().encoder_transfer_backend == "zmq_to_tokenizer": embedding_port = ( request.embedding_port[0] if request.embedding_port else 0 ) diff --git a/python/sglang/srt/disaggregation/encode_server.py b/python/sglang/srt/disaggregation/encode_server.py index f73b37d06..76c434db5 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -63,7 +63,7 @@ from sglang.srt.observability.trace import ( process_tracing_init, trace_set_thread_info, ) -from sglang.srt.runtime_context import publish +from sglang.srt.runtime_context import get_disagg, get_exec, get_mm, publish from sglang.srt.server_args import ( PortArgs, ServerArgs, @@ -352,7 +352,7 @@ class MMEncoder: [], dtype=self._embedding_dtype ).element_size() - if self.server_args.enable_mm_global_cache: + if get_mm().enable_mm_global_cache: from sglang.srt.mem_cache.storage.mooncake_store.embedding_cache_controller import ( EmbeddingCacheController, ) @@ -370,15 +370,15 @@ class MMEncoder: self.mm_global_cache = None # Pre-compute embedding metadata (needed by all ranks for mooncake) - if self.server_args.encoder_transfer_backend == "mooncake": + if get_disagg().encoder_transfer_backend == "mooncake": self._embedding_dims = self._infer_embedding_dims() if self.rank == 0: logger.info( - f"Using transfer backend: {self.server_args.encoder_transfer_backend}" + f"Using transfer backend: {get_disagg().encoder_transfer_backend}" ) - if self.server_args.encoder_transfer_backend == "mooncake": + if get_disagg().encoder_transfer_backend == "mooncake": self.local_ip = get_local_ip_auto() self.engine = get_mooncake_transfer_engine() @@ -391,8 +391,8 @@ class MMEncoder: hostname=self.local_ip, gpu_id=self.gpu_id, ib_device=( - self.server_args.disaggregation_ib_device - or self.server_args.mooncake_ib_device + get_disagg().disaggregation_ib_device + or get_exec().moe.mooncake_ib_device ), ) @@ -401,7 +401,7 @@ class MMEncoder: self.encode_dispatch_lock = asyncio.Lock() # Async mooncake state: track background VIT forward completion - if self.server_args.encoder_transfer_backend == "mooncake": + if get_disagg().encoder_transfer_backend == "mooncake": self._forward_ready_events: Dict[str, asyncio.Event] = {} self._forward_results: Dict[str, dict] = {} # when multiple decoder TP ranks call @@ -415,12 +415,12 @@ class MMEncoder: # Bind unified encode entry point based on backend and cache config if self.mm_global_cache is not None: - if self.server_args.encoder_transfer_backend == "mooncake": + if get_disagg().encoder_transfer_backend == "mooncake": self._encode_fn = self.encode_with_global_cache_mooncake else: self._encode_fn = self.encode_with_global_cache else: - if self.server_args.encoder_transfer_backend == "mooncake": + if get_disagg().encoder_transfer_backend == "mooncake": self._encode_fn = self.encode_with_mooncake else: self._encode_fn = self.encode @@ -1710,7 +1710,7 @@ class MMEncoder: mm_item.set(k, _convert(v)) cache_hit = False - use_mm_cache = self.server_args.enable_prefix_mm_cache and log_metrics + use_mm_cache = get_mm().enable_prefix_mm_cache and log_metrics if use_mm_cache: mm_item.set_pad_value() mm_hash = MultiModalStaticCache.combine_hashes([mm_item.hash]) @@ -1806,7 +1806,7 @@ class MMEncoder: embedding_port=None, url=None, ): - if self.server_args.encoder_transfer_backend == "mooncake": + if get_disagg().encoder_transfer_backend == "mooncake": # Wait for async VIT forward completion if needed req_id = mm_data.req_id if req_id in self._forward_ready_events: @@ -1878,7 +1878,7 @@ class MMEncoder: logger.info(f"{endpoint = }") # Serialize data - if self.server_args.encoder_transfer_backend == "mooncake": + if get_disagg().encoder_transfer_backend == "mooncake": # Mooncake already pushed the embedding via RDMA; new_mm_data = mm_data.copy_without_embedding() serialized_data = pickle.dumps(new_mm_data) @@ -1910,11 +1910,11 @@ class MMEncoder: await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket) if ( encoder_metrics_collector is not None - and self.server_args.encoder_transfer_backend != "mooncake" + and get_disagg().encoder_transfer_backend != "mooncake" ): encoder_metrics_collector.observe_transfer( time.perf_counter() - _zmq_xfer_start, - backend=self.server_args.encoder_transfer_backend, + backend=get_disagg().encoder_transfer_backend, ) async def encode( diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index be0d3f121..1e23c54ef 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -60,6 +60,7 @@ from sglang.srt.observability.trace import ( TraceReqContext, trace_set_thread_info, ) +from sglang.srt.runtime_context import get_schedule from sglang.srt.server_args import ServerArgs from sglang.srt.utils.network import NetworkAddress @@ -327,7 +328,7 @@ class MooncakeKVManager(CommonKVManager): lambda ptr, size: self.engine.batch_register([ptr], [size]), self.kv_args, count, - self.server_args.chunked_prefill_size, + get_schedule().chunked_prefill_size, ) self.kv_buffer_tensors = None @@ -498,7 +499,7 @@ class MooncakeKVManager(CommonKVManager): room, self.transfer_infos, self.kv_buffer_tensors, - self.server_args.chunked_prefill_size, + get_schedule().chunked_prefill_size, self._staging_ctx.prefetch_requested, self._staging_ctx.prefetch_sockets, ) diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index ff22ecfde..86a5b2904 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -44,6 +44,7 @@ from sglang.srt.disaggregation.utils import ( resolve_dcp_dst_entry_indices, ) from sglang.srt.environ import envs +from sglang.srt.runtime_context import get_schedule from sglang.srt.server_args import ServerArgs try: @@ -538,7 +539,7 @@ class NixlKVManager(CommonKVManager): lambda ptr, size: self._register_staging_memory(ptr, size, gpu_id), self.kv_args, count, - self.server_args.chunked_prefill_size, + get_schedule().chunked_prefill_size, ) def _init_staging_allocator(self): @@ -670,7 +671,7 @@ class NixlKVManager(CommonKVManager): room, self.transfer_infos, self.kv_buffer_tensors, - self.server_args.chunked_prefill_size, + get_schedule().chunked_prefill_size, self._staging_ctx.prefetch_requested, self._staging_ctx.prefetch_sockets, ) diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 87acd602c..9a28d739d 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -65,6 +65,7 @@ from sglang.srt.mem_cache.common import ( ) from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.observability.req_time_stats import set_schedule_time_batch +from sglang.srt.runtime_context import get_disagg from sglang.srt.utils import is_npu from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method @@ -183,6 +184,9 @@ class PrefillBootstrapQueue: kv_args.engine_rank = self.tp_rank kv_args.pp_rank = self.pp_rank kv_args.system_dp_rank = self.scheduler.ps.dp_rank + kv_args.kv_cache_dtype_str = ( + self.scheduler.tp_worker.model_runner.kv_cache_dtype_str + ) layer_shard_enabled = getattr( self.token_to_kv_pool, "layer_shard_enabled", False ) @@ -1249,7 +1253,7 @@ class SchedulerDisaggregationPrefillMixin: def optimistic_release_and_requeue(self: Scheduler, req: Req) -> None: """Release KV cache and requeue an optimistic prefill request.""" - max_attempts = self.server_args.optimistic_prefill_attempts + max_attempts = get_disagg().optimistic_prefill_attempts maybe_cache_unfinished_req(req, self.tree_cache) release_kv_cache(req, self.tree_cache) req.reset_for_retract() diff --git a/python/sglang/srt/distributed/device_communicators/pymscclpp.py b/python/sglang/srt/distributed/device_communicators/pymscclpp.py index 261e8d6cd..95d9411e4 100644 --- a/python/sglang/srt/distributed/device_communicators/pymscclpp.py +++ b/python/sglang/srt/distributed/device_communicators/pymscclpp.py @@ -14,7 +14,7 @@ from sglang.srt.compilation.compile_phase import ( from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( is_in_tc_piecewise_cuda_graph, ) -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec logger = logging.getLogger(__name__) @@ -25,7 +25,7 @@ class PyMscclppCommunicator: def _is_symm_mem_enabled(self) -> bool: try: - return get_server_args().enable_symm_mem + return get_exec().comm.enable_symm_mem except ValueError: return False diff --git a/python/sglang/srt/distributed/device_communicators/pynccl_allocator.py b/python/sglang/srt/distributed/device_communicators/pynccl_allocator.py index 3e833824e..5ce034f5b 100644 --- a/python/sglang/srt/distributed/device_communicators/pynccl_allocator.py +++ b/python/sglang/srt/distributed/device_communicators/pynccl_allocator.py @@ -15,7 +15,7 @@ from torch.cuda.memory import ( from sglang.srt.distributed.parallel_state import GroupCoordinator from sglang.srt.environ import envs -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec from sglang.srt.utils.common import torch_release after_2_8_0 = torch_release >= (2, 8) @@ -159,7 +159,7 @@ _register_func = None def is_symmetric_memory_enabled(): try: - return get_server_args().enable_symm_mem + return get_exec().comm.enable_symm_mem except ValueError: return False diff --git a/python/sglang/srt/distributed/device_communicators/torch_symm_mem.py b/python/sglang/srt/distributed/device_communicators/torch_symm_mem.py index 3ba756c7d..9c81599f2 100644 --- a/python/sglang/srt/distributed/device_communicators/torch_symm_mem.py +++ b/python/sglang/srt/distributed/device_communicators/torch_symm_mem.py @@ -12,6 +12,7 @@ from sglang.srt.distributed.device_communicators.all_reduce_utils import ( TORCH_SYMM_MEM_ALL_REDUCE_MAX_SIZES, ) from sglang.srt.environ import envs +from sglang.srt.runtime_context import get_exec from sglang.srt.utils import is_cuda, is_hip try: @@ -98,10 +99,9 @@ class TorchSymmMemCommunicator: # ([16384, 6144] bf16 = 192 MiB), including room for tail regions. if envs.SGLANG_OPT_USE_INKLING_CUSTOM_AR.get(): self.max_size = max(self.max_size, 256 * 1024 * 1024) - from sglang.srt.runtime_context import get_server_args if ( - get_server_args().enable_scattered_sconv + get_exec().comm.enable_scattered_sconv or envs.SGLANG_OPT_USE_INKLING_FUSED_AR_SCONV.get() ): # Fused extend kernels are out-of-place, so OUT must hold the diff --git a/python/sglang/srt/dllm/mixin/scheduler.py b/python/sglang/srt/dllm/mixin/scheduler.py index 6d532531e..abc4b3ed5 100644 --- a/python/sglang/srt/dllm/mixin/scheduler.py +++ b/python/sglang/srt/dllm/mixin/scheduler.py @@ -11,6 +11,7 @@ from sglang.srt.managers.schedule_policy import AddReqResult, PrefillAdder from sglang.srt.mem_cache.common import release_kv_cache from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.observability.req_time_stats import set_time_batch +from sglang.srt.runtime_context import get_exec, get_schedule logger = logging.getLogger(__name__) @@ -22,7 +23,7 @@ class SchedulerDllmMixin: def init_diffusion_llm(self: Scheduler): self.dllm_config = ( DllmConfig.from_server_args(self.server_args) - if self.server_args.dllm_algorithm is not None + if get_exec().dllm.dllm_algorithm is not None else None ) self.dllm_manager = DllmManager(dllm_config=self.dllm_config) @@ -200,7 +201,7 @@ class SchedulerDllmMixin: self.chunked_prefill_size, running_bs if self.is_mixed_chunk else 0, self.priority_scheduling_preemption_threshold, - prefill_max_requests=self.server_args.prefill_max_requests, + prefill_max_requests=get_schedule().prefill_max_requests, dllm_config=self.dllm_config, ) diff --git a/python/sglang/srt/elastic_ep/expert_backup_client.py b/python/sglang/srt/elastic_ep/expert_backup_client.py index 8b77f7f07..1b82477ba 100644 --- a/python/sglang/srt/elastic_ep/expert_backup_client.py +++ b/python/sglang/srt/elastic_ep/expert_backup_client.py @@ -14,6 +14,7 @@ from sglang.srt.distributed.parallel_state import ( from sglang.srt.environ import envs from sglang.srt.eplb.expert_location import get_global_expert_location_metadata from sglang.srt.managers.io_struct import UpdateExpertBackupReq, sock_recv, sock_send +from sglang.srt.runtime_context import get_exec from sglang.srt.server_args import ServerArgs from sglang.srt.utils.network import get_local_ip_auto @@ -111,7 +112,7 @@ class ExpertBackupClient: global_expert_location_metadata = get_global_expert_location_metadata() num_experts = ( self.model_config.hf_config.n_routed_experts - + self.server_args.ep_num_redundant_experts + + get_exec().moe.ep_num_redundant_experts ) num_local_experts = num_experts // self.moe_ep_size for i in range(self.engine_num): diff --git a/python/sglang/srt/entrypoints/grpc_bridge.py b/python/sglang/srt/entrypoints/grpc_bridge.py index ac8c25872..f2adfb8f8 100644 --- a/python/sglang/srt/entrypoints/grpc_bridge.py +++ b/python/sglang/srt/entrypoints/grpc_bridge.py @@ -377,9 +377,9 @@ class RuntimeHandle: model_config = self.tokenizer_manager.model_config result = { "model_path": self.tokenizer_manager.model_path, - "tokenizer_path": self.server_args.tokenizer_path, + "tokenizer_path": self.tokenizer_manager.server_args.tokenizer_path, "is_generation": self.tokenizer_manager.is_generation, - "weight_version": self.server_args.weight_version, + "weight_version": self.tokenizer_manager.server_args.weight_version, "model_type": getattr(model_config.hf_config, "model_type", None), "architectures": getattr(model_config.hf_config, "architectures", None), } @@ -432,7 +432,7 @@ class RuntimeHandle: "max_model_len": self.tokenizer_manager.model_config.context_len, } ] - if self.server_args.enable_lora and hasattr( + if self.tokenizer_manager.server_args.enable_lora and hasattr( self.tokenizer_manager, "lora_registry" ): lora_registry = self.tokenizer_manager.lora_registry diff --git a/python/sglang/srt/eplb/eplb_manager.py b/python/sglang/srt/eplb/eplb_manager.py index 4dc97b23e..00738b40f 100644 --- a/python/sglang/srt/eplb/eplb_manager.py +++ b/python/sglang/srt/eplb/eplb_manager.py @@ -18,7 +18,7 @@ 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_server_args +from sglang.srt.runtime_context import get_model if TYPE_CHECKING: from sglang.srt.configs.model_config import ModelConfig @@ -343,8 +343,8 @@ def update_expert_location_with_recovery( else: # Load the missing weights from disk update_weights_from_disk_callable( - get_server_args().model_path, - get_server_args().load_format, + get_model().model_path, + get_model().load_format, weight_name_filter=weight_name_filter, ) diff --git a/python/sglang/srt/eplb/expert_location_dispatch.py b/python/sglang/srt/eplb/expert_location_dispatch.py index 60b01f29e..6325417b1 100644 --- a/python/sglang/srt/eplb/expert_location_dispatch.py +++ b/python/sglang/srt/eplb/expert_location_dispatch.py @@ -18,7 +18,7 @@ from typing import Literal, Optional import torch from sglang.srt.eplb.expert_location import get_global_expert_location_metadata -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec @dataclass @@ -40,8 +40,7 @@ class ExpertLocationDispatchInfo: @classmethod def init_new(cls, layer_id: int): - server_args = get_server_args() - ep_dispatch_algorithm = server_args.ep_dispatch_algorithm + ep_dispatch_algorithm = get_exec().moe.ep_dispatch_algorithm expert_location_metadata = get_global_expert_location_metadata() assert expert_location_metadata is not None @@ -50,7 +49,7 @@ class ExpertLocationDispatchInfo: return cls( ep_dispatch_algorithm=ep_dispatch_algorithm, - rank_invariant=server_args.moe_a2a_backend == "none", + rank_invariant=get_exec().moe.moe_a2a_backend == "none", partial_logical_to_rank_dispatch_physical_map=( expert_location_metadata.logical_to_rank_dispatch_physical_map[ layer_id, : diff --git a/python/sglang/srt/eplb/expert_location_updater.py b/python/sglang/srt/eplb/expert_location_updater.py index 46bd5cdf2..1fd0022ef 100644 --- a/python/sglang/srt/eplb/expert_location_updater.py +++ b/python/sglang/srt/eplb/expert_location_updater.py @@ -26,7 +26,7 @@ from sglang.srt.eplb.expert_location import ( ExpertLocationMetadata, get_global_expert_location_metadata, ) -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_device from sglang.srt.utils import get_bool_env_var logger = logging.getLogger(__name__) @@ -109,7 +109,7 @@ def _update_expert_weights_with_canary( canary_tensor = ( _get_canary_value(old_expert_location_metadata, layer_id) .clone() - .to(device=get_server_args().device, non_blocking=True) + .to(device=get_device().device, non_blocking=True) ) routed_experts_weights_of_layer[layer_id].append(canary_tensor) diff --git a/python/sglang/srt/hardware_backend/mlx/model_runner_stub.py b/python/sglang/srt/hardware_backend/mlx/model_runner_stub.py index 920efb0d0..114d711d9 100644 --- a/python/sglang/srt/hardware_backend/mlx/model_runner_stub.py +++ b/python/sglang/srt/hardware_backend/mlx/model_runner_stub.py @@ -19,6 +19,7 @@ from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner_components.layer_setup import ( ModelLayerInfo, ) +from sglang.srt.runtime_context import get_exec, get_memory, get_schedule logger = logging.getLogger(__name__) @@ -144,13 +145,13 @@ class MlxModelRunnerStub(ModelRunner): (``MlxAuxiliaryStateComponent``) raises ``NotImplementedError`` for the mode. """ - if self.server_args.disable_radix_cache: + if get_memory().disable_radix_cache: return 1 return MLX_AUX_STATE_SIZE_MAX_RUNNING_REQUESTS_RATIO def _explicit_aux_state_size_per_worker(self) -> int | None: """Return the explicit auxiliary-state cap for this attention-DP owner.""" - aux_state_size = self.server_args.max_mamba_cache_size + aux_state_size = get_schedule().max_mamba_cache_size if aux_state_size is None: return None return aux_state_size // self.ps.attn_dp_size @@ -173,7 +174,7 @@ class MlxModelRunnerStub(ModelRunner): Requires ``self.max_total_num_tokens`` to already be set. """ capacity_cap = self.max_total_num_tokens // 2 - requested = self.server_args.max_running_requests + requested = get_schedule().max_running_requests if requested is None: requested_per_worker = None resolved = min(capacity_cap, 4096) @@ -189,7 +190,7 @@ class MlxModelRunnerStub(ModelRunner): ratio = self._aux_state_slots_per_request() resolved = min(resolved, aux_state_size // ratio) if resolved <= 0: - global_aux_state_size = self.server_args.max_mamba_cache_size + global_aux_state_size = get_schedule().max_mamba_cache_size min_global_aux_state_size = ratio * self.ps.attn_dp_size raise RuntimeError( f"MLX auxiliary-state cache is too small to serve any " @@ -221,7 +222,7 @@ class MlxModelRunnerStub(ModelRunner): from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter self.memory_saver_adapter = TorchMemorySaverAdapter.create( - enable=self.server_args.enable_memory_saver + enable=get_exec().features.enable_memory_saver ) # Load model (sets metadata only) @@ -267,7 +268,7 @@ class MlxModelRunnerStub(ModelRunner): # With the radix cache disabled no tree component exists to # release auxiliary slots, so the pool owns their release # (see MlxAuxiliaryStateReqToTokenPool docstring). - owns_auxiliary_state_release=self.server_args.disable_radix_cache, + owns_auxiliary_state_release=get_memory().disable_radix_cache, ) else: self.req_to_token_pool = ReqToTokenPool( diff --git a/python/sglang/srt/hardware_backend/mlx/tp_worker.py b/python/sglang/srt/hardware_backend/mlx/tp_worker.py index 53f9b88c1..d79ac6784 100644 --- a/python/sglang/srt/hardware_backend/mlx/tp_worker.py +++ b/python/sglang/srt/hardware_backend/mlx/tp_worker.py @@ -31,6 +31,7 @@ from sglang.srt.model_executor.forward_batch_info import ( ForwardBatch, PPProxyTensors, ) +from sglang.srt.runtime_context import get_memory, get_model, get_schedule logger = logging.getLogger(__name__) @@ -53,19 +54,19 @@ class MlxTpModelWorker(TpModelWorker): logger.info("Initializing MlxModelRunner for end-to-end MLX inference") init_kwargs = dict( - model_path=self.server_args.model_path, - trust_remote_code=self.server_args.trust_remote_code, - disable_radix_cache=self.server_args.disable_radix_cache, - mem_fraction_static=self.server_args.mem_fraction_static, - quantization=self.server_args.quantization, + model_path=get_model().model_path, + trust_remote_code=get_model().trust_remote_code, + disable_radix_cache=get_memory().disable_radix_cache, + mem_fraction_static=get_schedule().mem_fraction_static, + quantization=get_model().quantization, ) - if self.server_args.max_total_tokens is not None: - init_kwargs["pool_size"] = self.server_args.max_total_tokens + if get_schedule().max_total_tokens is not None: + init_kwargs["pool_size"] = get_schedule().max_total_tokens self._mlx_runner = MlxModelRunner(**init_kwargs) self._model_runner = MlxModelRunnerStub( model_config=self.model_config, - mem_fraction_static=self.server_args.mem_fraction_static, + mem_fraction_static=get_schedule().mem_fraction_static, gpu_id=self.gpu_id, ps=self.ps, nccl_port=self.nccl_port, diff --git a/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py b/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py index 6b823eaec..403c7f1f9 100644 --- a/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py +++ b/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py @@ -23,7 +23,7 @@ from sglang.srt.layers.utils.cp_utils import ( cp_allgather_and_save_kv_cache, ) from sglang.srt.mem_cache.memory_pool import KVWriteLoc -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_schedule if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention @@ -515,7 +515,7 @@ class MusaFlashAttentionBackend(FlashAttentionBackend): and not forward_batch.forward_mode.is_draft_extend_v2() ): if forward_batch.attn_attend_prefix_cache: - assert not get_server_args().disable_chunked_prefix_cache + assert not get_schedule().disable_chunked_prefix_cache assert forward_batch.prefix_chunk_idx is not None assert forward_batch.prefix_chunk_cu_seq_lens is not None assert forward_batch.prefix_chunk_max_seq_lens is not None diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py index d70872e67..89e9de06d 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py @@ -26,7 +26,7 @@ from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_kv_cache from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode -from sglang.srt.runtime_context import get_flags +from sglang.srt.runtime_context import get_flags, get_spec from sglang.srt.speculative.spec_info import SpecInput, SpecInputType from sglang.srt.utils import get_bool_env_var, get_current_device_stream_fast @@ -336,9 +336,7 @@ class AscendAttnBackend(AttentionBackend): self.use_fa = get_bool_env_var("ASCEND_USE_FA", "False") self.use_fia = get_bool_env_var("ASCEND_USE_FIA", "False") self.enable_torch_compile = get_flags().capture.enable_torch_compile - self.speculative_num_draft_tokens = ( - model_runner.server_args.speculative_num_draft_tokens - ) + self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens self.ascend_attn_mask_builder = AscendAttnMaskBuilder( model_runner, self.device, self.use_fia, self.use_mla ) diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_dsv4_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_dsv4_backend.py index eb5ee46ce..5c147a2af 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_dsv4_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_dsv4_backend.py @@ -13,7 +13,7 @@ from sglang.srt.layers.attention.dsv4.compressor import CompressorBackendMixin from sglang.srt.layers.attention.dsv4.indexer import C4IndexerBackendMixin from sglang.srt.model_executor.forward_batch_info import DSV4OutCacheLoc, ForwardMode from sglang.srt.model_executor.forward_context import get_attn_backend -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_parallel, get_spec if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention @@ -1493,9 +1493,8 @@ class DeepseekV4AscendAttnBackend( or forward_batch.forward_mode.is_draft_extend_v2() ): B = forward_batch.batch_size - from sglang.srt.runtime_context import get_server_args - n_draft = get_server_args().speculative_num_draft_tokens or 1 + n_draft = get_spec().speculative_num_draft_tokens or 1 actual_q = torch.arange( n_draft, B * n_draft + 1, n_draft, dtype=torch.int32, device=device ) @@ -1540,9 +1539,8 @@ class DeepseekV4AscendAttnBackend( forward_batch.forward_mode.is_target_verify() or forward_batch.forward_mode.is_draft_extend_v2() ): - from sglang.srt.runtime_context import get_server_args - max_seqlen_q = get_server_args().speculative_num_draft_tokens or 1 + max_seqlen_q = get_spec().speculative_num_draft_tokens or 1 else: max_seqlen_q = 1 return self._kernel_metadata_from_parts( diff --git a/python/sglang/srt/hardware_backend/npu/graph_runner/vit_npu_graph_runner.py b/python/sglang/srt/hardware_backend/npu/graph_runner/vit_npu_graph_runner.py index 3738c0d36..46dd3a101 100644 --- a/python/sglang/srt/hardware_backend/npu/graph_runner/vit_npu_graph_runner.py +++ b/python/sglang/srt/hardware_backend/npu/graph_runner/vit_npu_graph_runner.py @@ -27,7 +27,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import ( ) from sglang.srt.layers.attention.vision import VisionAttention from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_mm class ViTNpuGraphRunner(ViTCudaGraphRunner): @@ -70,7 +70,7 @@ class ViTNpuGraphRunner(ViTCudaGraphRunner): graph = torch_npu.npu.NPUGraph() vit = self.vit - override_backend = get_server_args().mm_attention_backend + override_backend = get_mm().mm_attention_backend with torch_npu.npu.graph(graph, pool=ViTNpuGraphRunner._graph_memory_pool): y = None deepstack_outs: List[torch.Tensor] = [] diff --git a/python/sglang/srt/hardware_backend/npu/moe/fuseep.py b/python/sglang/srt/hardware_backend/npu/moe/fuseep.py index 9c5c16c9d..609f6b6ca 100644 --- a/python/sglang/srt/hardware_backend/npu/moe/fuseep.py +++ b/python/sglang/srt/hardware_backend/npu/moe/fuseep.py @@ -17,7 +17,7 @@ from sglang.srt.environ import envs from sglang.srt.hardware_backend.npu.utils import npu_format_cast from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPBuffer from sglang.srt.layers.moe.utils import DeepEPMode -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec if TYPE_CHECKING: from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE @@ -57,7 +57,7 @@ def forward_fuseep( envs.SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get() ), num_experts=layer.num_experts, - fuse_mode=get_server_args().fuseep_mode, + fuse_mode=get_exec().moe.fuseep_mode, ) return hidden_states @@ -126,7 +126,7 @@ def process_fuseep_weights(layer: torch.nn.Module, weight_prefix: str) -> None: Invoked by ``maybe_apply_fuseep_weights`` for both ``"w13"`` and ``"w2"``. """ - if get_server_args().fuseep_mode == 1: + if get_exec().moe.fuseep_mode == 1: # -- The fused MoE optimization mode "1": dispatch_gmm_combine_decode -- if weight_prefix == "w13": cpu_w13 = layer.w13_weight.data.transpose(1, 2).cpu() @@ -143,7 +143,7 @@ def process_fuseep_weights(layer: torch.nn.Module, weight_prefix: str) -> None: layer.w2_weight_scale = torch.nn.Parameter( w2_scale.to(torch.float32), requires_grad=False ) - elif get_server_args().fuseep_mode == 2: + elif get_exec().moe.fuseep_mode == 2: # -- The fused MoE optimization mode "2": dispatch_ffn_combine -- if weight_prefix == "w13": w13_weight = _release_weight_cache(layer.w13_weight) diff --git a/python/sglang/srt/layers/activation.py b/python/sglang/srt/layers/activation.py index 0e97fe5fb..8062eb465 100644 --- a/python/sglang/srt/layers/activation.py +++ b/python/sglang/srt/layers/activation.py @@ -33,7 +33,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( Phase, check_cuda_graph_backend, ) -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -130,7 +130,7 @@ logger = logging.getLogger(__name__) class SiluAndMul(MultiPlatformOp): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - if get_server_args().rl_on_policy_target is not None: + if get_exec().deterministic.rl_on_policy_target is not None: self._forward_method = self.forward_native elif _use_aiter and envs.SGLANG_OPT_USE_AITER_SILU_MUL.get(): self._forward_method = self.forward_aiter diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 513461823..a213cc5ab 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -1,6 +1,6 @@ from __future__ import annotations -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_parallel, get_spec """ end to end attention solution with aiter kernels @@ -148,8 +148,8 @@ class AiterAttnBackend(AttentionBackend): self.device = model_runner.device self.is_multimodal = model_runner.model_config.is_multimodal - self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens - self.speculative_num_steps = model_runner.server_args.speculative_num_steps + self.num_draft_tokens = get_spec().speculative_num_draft_tokens + self.speculative_num_steps = get_spec().speculative_num_steps self.topk = topk self.num_head = ( model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index bd4385dfa..f85b3bfc6 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -60,7 +60,7 @@ from sglang.srt.layers.attention.dsv4.sparse_prefill_utils import ( from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_parallel, get_spec from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc from sglang.srt.speculative.ragged_verify import ( RaggedVerifyMode, @@ -537,9 +537,7 @@ class DeepseekV4AttnBackend( assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4" self.mtp_enabled = self.topk > 0 self.speculative_num_steps = speculative_num_steps - self.speculative_num_draft_tokens: int = ( - model_runner.server_args.speculative_num_draft_tokens - ) + self.speculative_num_draft_tokens: int = get_spec().speculative_num_draft_tokens if self.speculative_num_draft_tokens is not None: # Persistent target-verify metadata buffers. Allocated here (not # lazily) so they are ordinary tensors: the first touch of a lazy diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py index 16b307a7e..a6adfc929 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend_hip_radix.py @@ -39,7 +39,7 @@ from sglang.srt.layers.attention.dsv4.metadata import ( ) from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_parallel, get_spec from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout from sglang.srt.utils import ceil_align @@ -455,9 +455,7 @@ class DeepseekV4HipRadixBackend( assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4" self.mtp_enabled = self.topk > 0 self.speculative_num_steps = speculative_num_steps - self.speculative_num_draft_tokens: int = ( - model_runner.server_args.speculative_num_draft_tokens - ) + self.speculative_num_draft_tokens: int = get_spec().speculative_num_draft_tokens self.speculative_step_id = speculative_step_id self.forward_metadata: Union[ DSV4Metadata, diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index cee6b5bd0..2ba465ba3 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -37,7 +37,14 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo get_tc_piecewise_forward_context, is_in_tc_piecewise_cuda_graph, ) -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_device, + get_exec, + get_parallel, + get_schedule, + get_server_args, + get_spec, +) from sglang.srt.state_capturer.indexer_topk import ( maybe_capture_indexer_topk, ) @@ -163,7 +170,7 @@ def _uses_dsa_attention_backend(forward_batch: ForwardBatch) -> bool: ): backend_name = ( decode_backend - if server_args.speculative_attention_mode == "decode" + if get_spec().speculative_attention_mode == "decode" else prefill_backend ) else: @@ -460,7 +467,7 @@ class Indexer(MultiPlatformOp): base=rope_theta, # type: ignore rope_scaling=rope_scaling, is_neox_style=is_neox_style, - device=get_server_args().device, + device=get_device().device, ) self.block_size = block_size self.scale_fmt = scale_fmt @@ -471,7 +478,7 @@ class Indexer(MultiPlatformOp): self.num_local_tokens = getattr(config, "index_local_tokens", 0) self.paged_mqa_logits_backend = DSAPagedMQALogitsBackend.resolve( - get_server_args().dsa_paged_mqa_logits_backend + get_exec().kernel.dsa_paged_mqa_logits_backend ) @contextlib.contextmanager @@ -1066,7 +1073,7 @@ class Indexer(MultiPlatformOp): total_mem = torch.cuda.get_device_properties(device_index).total_memory total_mem_budget = int(total_mem * self._MQA_LOGITS_TOTAL_MEM_FRACTION) - mem_fraction_static = get_server_args().mem_fraction_static + mem_fraction_static = get_schedule().mem_fraction_static if mem_fraction_static is None: static_budget = total_mem_budget else: diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index 3fe440859..b9094679b 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -15,7 +15,7 @@ from typing import ( import torch from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_parallel, get_spec logger = logging.getLogger(__name__) from sglang.kernels.ops.attention.dsa.dequant_k_cache import ( @@ -469,9 +469,7 @@ class DeepseekSparseAttnBackend( # Speculative decoding self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.speculative_num_steps = speculative_num_steps - self.speculative_num_draft_tokens = ( - model_runner.server_args.speculative_num_draft_tokens - ) + self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens self.speculative_step_id = speculative_step_id self.use_fused_topk = should_use_dsa_fused_topk( model_runner.server_args, seed_dsa_topk_from_draft_extend diff --git a/python/sglang/srt/layers/attention/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index bf46609e1..f596caade 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -40,7 +40,7 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( is_in_tc_piecewise_cuda_graph, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer from sglang.srt.utils import add_prefix, is_cuda, is_hip, is_xpu from sglang.srt.utils.common import is_sm120_supported @@ -922,9 +922,8 @@ class C4Indexer(nn.Module): self.rotary_emb = rotary_emb self.freqs_cis = freqs_cis self.weight_scale: float = self.softmax_scale * self.n_heads**-0.5 - from sglang.srt.runtime_context import get_server_args - self.use_fp4_indexer = get_server_args().enable_deepseek_v4_fp4_indexer + self.use_fp4_indexer = get_exec().kernel.enable_deepseek_v4_fp4_indexer self.alt_streams = alt_streams def compute_q( diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index 7535a8d8e..41ebbdea8 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -30,7 +30,7 @@ from sglang.srt.layers.utils.cp_utils import ( from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_schedule, get_spec from sglang.srt.speculative.ragged_verify import build_ragged_target_verify_geometry from sglang.srt.speculative.spec_info import SpecInput, SpeculativeAlgorithm from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req @@ -204,9 +204,7 @@ class FlashAttentionBackend(AttentionBackend): self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.speculative_num_steps = speculative_num_steps - self.speculative_num_draft_tokens = ( - model_runner.server_args.speculative_num_draft_tokens - ) + self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens if ( self.speculative_num_draft_tokens is not None and model_runner.is_draft_worker @@ -1513,7 +1511,7 @@ class FlashAttentionBackend(AttentionBackend): ): # Do multi-head attention with chunked prefix cache if forward_batch.attn_attend_prefix_cache: - assert not get_server_args().disable_chunked_prefix_cache + assert not get_schedule().disable_chunked_prefix_cache # MHA for chunked prefix kv cache when running model with MLA assert forward_batch.prefix_chunk_idx is not None assert forward_batch.prefix_chunk_cu_seq_lens is not None diff --git a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py index ebb5b40ef..62c43eb7c 100644 --- a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py @@ -1,6 +1,6 @@ from __future__ import annotations -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_disagg, get_exec, get_parallel, get_schedule """ Support attention backend for flashinfer MLA. @@ -33,7 +33,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( is_in_tc_piecewise_cuda_graph, ) -from sglang.srt.runtime_context import get_buffer, get_server_args +from sglang.srt.runtime_context import get_buffer from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_utils import ( draft_kv_indices_buffer_width, @@ -224,9 +224,9 @@ class FlashInferMLAAttnBackend(AttentionBackend): self.token_to_kv_pool = model_runner.token_to_kv_pool self.enable_chunk_kv = ( not skip_prefill - and get_server_args().disaggregation_mode != "decode" - and not get_server_args().disable_chunked_prefix_cache - and not get_server_args().flashinfer_mla_disable_ragged + and get_disagg().disaggregation_mode != "decode" + and not get_schedule().disable_chunked_prefix_cache + and not get_exec().kernel.flashinfer_mla_disable_ragged ) self.page_size = model_runner.page_size @@ -402,7 +402,7 @@ class FlashInferMLAAttnBackend(AttentionBackend): prefix_lens = forward_batch.extend_prefix_lens extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu) use_ragged = ( - not get_server_args().flashinfer_mla_disable_ragged + not get_exec().kernel.flashinfer_mla_disable_ragged and extend_no_prefix # Piecewise cuda graph should use paged prefill to be compatible with prefix cache and not is_in_tc_piecewise_cuda_graph() diff --git a/python/sglang/srt/layers/attention/flashmla_backend.py b/python/sglang/srt/layers/attention/flashmla_backend.py index ce628b818..65106cbf5 100644 --- a/python/sglang/srt/layers/attention/flashmla_backend.py +++ b/python/sglang/srt/layers/attention/flashmla_backend.py @@ -20,7 +20,7 @@ from sglang.kernels.ops.quantization.fp8_kernel import scaled_fp8_quant from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_parallel, get_spec logger = logging.getLogger(__name__) @@ -94,7 +94,7 @@ class FlashMLABackend(FlashInferMLAAttnBackend): torch.float8_e5m2, } - self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens + self.num_draft_tokens = get_spec().speculative_num_draft_tokens self.cuda_graph_kv_indices = None self.cuda_graph_mla_metadata = None diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index 0db8ff036..0c6d9adb9 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -24,7 +24,7 @@ from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.model_runner import ModelRunner -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec, get_memory, get_server_args from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput from sglang.srt.speculative.spec_info import SpecInput @@ -392,7 +392,7 @@ class MambaAttnBackendBase(AttentionBackend): """Per-row (length bs) bool flush mask = the radix track's seq_lens_cpu % mamba_track_interval == 0, so force-flush and snapshot fire on the same steps (no off-by-one).""" - interval = get_server_args().mamba_track_interval + interval = get_exec().mamba.mamba_track_interval if seq_lens_cpu is None: # Should not happen for the supported config; stay safe and never flush. return torch.zeros((bs,), dtype=torch.bool) @@ -823,7 +823,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase): # Page-major stores state strided; only the stride-aware Triton causal-conv # reads it (CUDA causal_conv1d garbles it). A model may also force Triton. use_triton_causal_conv = ( - use_triton_causal_conv or get_server_args().enable_page_major_kv_layout + use_triton_causal_conv or get_memory().enable_page_major_kv_layout ) layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id) mixer_out, intermediate_states = mixer.forward( diff --git a/python/sglang/srt/layers/attention/intel_amx_backend.py b/python/sglang/srt/layers/attention/intel_amx_backend.py index 20769b0cb..c54dd395a 100644 --- a/python/sglang/srt/layers/attention/intel_amx_backend.py +++ b/python/sglang/srt/layers/attention/intel_amx_backend.py @@ -8,6 +8,7 @@ from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch +from sglang.srt.runtime_context import get_spec if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention @@ -59,7 +60,7 @@ class IntelAMXAttnBackend(AttentionBackend): self.num_kv_splits = 8 # speculative decoding params - self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens + self.num_draft_tokens = get_spec().speculative_num_draft_tokens def _build_extend_metadata(self, forward_batch: ForwardBatch): """Resolve (seq_lens, extend_seq_lens, extend_start_loc, tree_mask) for diff --git a/python/sglang/srt/layers/attention/linear/inkling_sconv_backend.py b/python/sglang/srt/layers/attention/linear/inkling_sconv_backend.py index 80f5c2ff3..2130ff563 100644 --- a/python/sglang/srt/layers/attention/linear/inkling_sconv_backend.py +++ b/python/sglang/srt/layers/attention/linear/inkling_sconv_backend.py @@ -60,7 +60,7 @@ from sglang.srt.models.inkling_common.kernels.sconv import ( fused_extend_sconv_metadata, precompute_helion_extend_metadata, ) -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec, get_server_args, get_spec from sglang.srt.speculative.eagle_info import EagleDraftExtendInput if TYPE_CHECKING: @@ -117,7 +117,7 @@ class InklingShortConvAttnBackend(ShortConvAttnBackend): growing a buffer after a graph captured it moves the address that graph reads, and prefill captures before the decode runner reports its bounds.""" server_args = get_server_args() - cuda_graph_config = server_args.cuda_graph_config + cuda_graph_config = get_exec().graph.cuda_graph_config decode_bs: list[int] = [] prefill_tokens: list[int] = [] decode_max_bs = 0 @@ -125,7 +125,7 @@ class InklingShortConvAttnBackend(ShortConvAttnBackend): decode_bs = list(cuda_graph_config.decode.bs or []) prefill_tokens = list(cuda_graph_config.prefill.bs or []) decode_max_bs = cuda_graph_config.decode.max_bs or 0 - draft_token_num = server_args.speculative_num_draft_tokens or 1 + draft_token_num = get_spec().speculative_num_draft_tokens or 1 # req_to_token_pool.size is the runner's max_bs for both graph phases. max_bs = max([self.req_to_token_pool.size, decode_max_bs, *decode_bs]) max_tokens = max([max_bs, *prefill_tokens, max_bs * draft_token_num]) diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index 02fe634b0..57e90f2b4 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -37,7 +37,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( cuda_graph_fully_disabled, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_parallel, get_spec from sglang.srt.speculative.spec_utils import ( draft_kv_indices_buffer_width, draft_kv_indices_used_len, @@ -168,9 +168,9 @@ class TritonAttnBackend(AttentionBackend): self._translate_kv_loc = getattr( self.token_to_kv_pool_allocator, "translate_kv_loc_dense", None ) or getattr(self.token_to_kv_pool_allocator, "translate_kv_loc", None) - self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens - self.speculative_num_steps = model_runner.server_args.speculative_num_steps - self.topk = model_runner.server_args.speculative_eagle_topk or 0 + self.num_draft_tokens = get_spec().speculative_num_draft_tokens + self.speculative_num_steps = get_spec().speculative_num_steps + self.topk = get_spec().speculative_eagle_topk or 0 # Split-KV verify is bit-equivalent only for a pure-causal chain (topk==1) # and is gfx95-only; else fall back to extend_attention_fwd. self.use_verify_splitkv = ( diff --git a/python/sglang/srt/layers/attention/trtllm_mha_backend.py b/python/sglang/srt/layers/attention/trtllm_mha_backend.py index ea05b28d1..cf68608e3 100644 --- a/python/sglang/srt/layers/attention/trtllm_mha_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mha_backend.py @@ -34,7 +34,7 @@ from sglang.srt.layers.radix_attention import AttentionType from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode -from sglang.srt.runtime_context import get_buffer +from sglang.srt.runtime_context import get_buffer, get_spec from sglang.srt.speculative.ragged_verify import ( build_ragged_target_verify_geometry, resolve_ragged_verify_layout, @@ -162,9 +162,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): self.speculative_step_id = speculative_step_id self.target_verify_metadata = {} - self.speculative_num_draft_tokens = ( - model_runner.server_args.speculative_num_draft_tokens - ) + self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens # True iff the model declares ENCODER_ONLY (bidirectional) layers, which # need the expanded TARGET_VERIFY metadata (TRTLLMMHAMetadata.encoder_*). self.expand_encoder_only_verify = any( diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index bb66ddc58..5c3a76a2e 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -41,7 +41,12 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( is_in_tc_piecewise_cuda_graph, ) -from sglang.srt.runtime_context import get_buffer, get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_buffer, + get_parallel, + get_schedule, + get_spec, +) from sglang.srt.utils import is_flashinfer_available, is_float4_e2m1fn_x2 if is_flashinfer_available(): @@ -238,11 +243,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): self.forward_prefill_metadata: Optional[TRTLLMMLAPrefillMetadata] = None self.forward_decode_metadata: Union[TRTLLMMLADecodeMetadata, None] = None - self.disable_chunked_prefix_cache = ( - get_server_args().disable_chunked_prefix_cache - ) + self.disable_chunked_prefix_cache = get_schedule().disable_chunked_prefix_cache - self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens + self.num_draft_tokens = get_spec().speculative_num_draft_tokens self._verify_mask = None # Tree-mask scratch is fetched from the target backend only. self.is_draft_runner = model_runner.is_draft_worker diff --git a/python/sglang/srt/layers/attention/vision.py b/python/sglang/srt/layers/attention/vision.py index 8eeb0807c..25d065228 100644 --- a/python/sglang/srt/layers/attention/vision.py +++ b/python/sglang/srt/layers/attention/vision.py @@ -17,7 +17,7 @@ from sglang.kernels.ops.layernorm.norm import ( ) from sglang.srt.environ import envs from sglang.srt.models.utils import apply_qk_norm -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_exec, get_mm, get_parallel from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -88,7 +88,6 @@ from sglang.srt.layers.linear import ( from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.layers.rotary_embedding import apply_rotary_pos_emb from sglang.srt.layers.rotary_embedding.utils import apply_rotary_pos_emb_native_eager -from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import add_prefix _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip @@ -1047,7 +1046,7 @@ class VisionAttention(nn.Module): # Select attention backend via a unified method _passed_backend = qkv_backend qkv_backend = self._determine_attention_backend(_passed_backend) - if get_server_args().mm_attention_backend is None and _passed_backend is None: + if get_mm().mm_attention_backend is None and _passed_backend is None: print_info_once(f"Multimodal attention backend not set. Use {qkv_backend}.") print_info_once(f"Using {qkv_backend} as multimodal attention backend.") @@ -1126,7 +1125,7 @@ class VisionAttention(nn.Module): weight_dtype=torch.float32, cast_x_before_out_mul=True, ) - if get_server_args().rl_on_policy_target is not None + if get_exec().deterministic.rl_on_policy_target is not None else {} ) q_norm = RMSNorm( @@ -1154,7 +1153,7 @@ class VisionAttention(nn.Module): - CUDA (other): "triton_attn" - Non-CUDA: "sdpa" """ - override_backend = get_server_args().mm_attention_backend + override_backend = get_mm().mm_attention_backend if override_backend is not None: backend = override_backend elif passed_backend is not None: @@ -1259,7 +1258,7 @@ class VisionAttention(nn.Module): x = x.unsqueeze(0) assert x.dim() == 3, x.shape if ( - get_server_args().rl_on_policy_target is not None + get_exec().deterministic.rl_on_policy_target is not None and position_embeddings is not None ): assert isinstance(position_embeddings, tuple), ( diff --git a/python/sglang/srt/layers/attention/wave_backend.py b/python/sglang/srt/layers/attention/wave_backend.py index 8381e1330..9764b7e41 100644 --- a/python/sglang/srt/layers/attention/wave_backend.py +++ b/python/sglang/srt/layers/attention/wave_backend.py @@ -13,7 +13,7 @@ from sglang.kernels.ops.kvcache.kv_indices import ( ) from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_parallel, get_spec from sglang.srt.utils import get_bool_env_var, get_device_core_count if TYPE_CHECKING: @@ -92,7 +92,7 @@ class WaveAttnBackend(AttentionBackend): (max_bs + 1,), dtype=torch.int64, device=model_runner.device ) - self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens + self.num_draft_tokens = get_spec().speculative_num_draft_tokens self.num_head = ( model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size diff --git a/python/sglang/srt/layers/attention/xpu_backend.py b/python/sglang/srt/layers/attention/xpu_backend.py index 2b6b469ca..9a4561921 100644 --- a/python/sglang/srt/layers/attention/xpu_backend.py +++ b/python/sglang/srt/layers/attention/xpu_backend.py @@ -15,7 +15,7 @@ from sglang.srt.layers.attention.flashattention_backend import ( from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_schedule, get_spec if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention @@ -86,9 +86,7 @@ class XPUAttentionBackend(AttentionBackend): ) self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.speculative_num_steps = speculative_num_steps - self.speculative_num_draft_tokens = ( - model_runner.server_args.speculative_num_draft_tokens - ) + self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens self.speculative_step_id = speculative_step_id # Local attention settings @@ -638,7 +636,7 @@ class XPUAttentionBackend(AttentionBackend): ): # Do multi-head attention with chunked prefix cache if forward_batch.attn_attend_prefix_cache: - assert not get_server_args().disable_chunked_prefix_cache + assert not get_schedule().disable_chunked_prefix_cache # MHA for chunked prefix kv cache when running model with MLA assert forward_batch.prefix_chunk_idx is not None assert forward_batch.prefix_chunk_cu_seq_lens is not None diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index 92ab37d80..dd81f9489 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -73,7 +73,13 @@ from sglang.srt.model_executor.cuda_graph_config import ( check_cuda_graph_backend, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch -from sglang.srt.runtime_context import get_forward, get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_exec, + get_forward, + get_parallel, + get_server_args, + get_spec, +) from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils import ( get_bool_env_var, @@ -169,7 +175,7 @@ def apply_flashinfer_allreduce_fusion(batch_size: int): (_is_sm90_supported or _is_sm100_supported) and _is_flashinfer_available and not is_dp_attention_enabled() - and get_server_args().flashinfer_allreduce_fusion_backend is not None + and get_exec().comm.flashinfer_allreduce_fusion_backend is not None and not is_flashinfer_allreduce_unavailable() # Symbolic size checks stay last: under Dynamo tracing they guard on # the dynamic token dim, so statically-off configs must short-circuit @@ -190,7 +196,7 @@ def apply_aiter_all_reduce_fusion(input_tensor: torch.Tensor): and total_bytes <= 8 * 1024 * 8192 and get_parallel().tp_size != 6 and not is_dp_attention_enabled() - and get_server_args().enable_aiter_allreduce_fusion + and get_exec().comm.enable_aiter_allreduce_fusion ) @@ -278,7 +284,7 @@ class AttnTpContext: and get_moe_a2a_backend().is_none() and not enable_moe_dense_fully_dp() and not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE) - and get_server_args().speculative_algorithm != "EAGLE3" + and get_spec().speculative_algorithm != "EAGLE3" ) if get_server_args().enable_attn_tp_input_scattered: if not self.allow_input_scattered: @@ -411,7 +417,7 @@ class LayerScatterModes: not context.is_layer_sparse and context.is_next_layer_sparse and enable_moe_dense_fully_dp() - and get_server_args().enable_two_batch_overlap + and get_exec().overlap.enable_two_batch_overlap ) @classmethod @@ -475,7 +481,7 @@ class LayerCommunicator: ) self._post_init_communicate() self._speculative_algo = SpeculativeAlgorithm.from_string( - get_server_args().speculative_algorithm + get_spec().speculative_algorithm ) def _post_init_communicate(self): @@ -846,7 +852,7 @@ class LayerCommunicator: and get_parallel().tp_size != 6 and not is_dp_attention_enabled() and get_moe_a2a_backend().is_none() - and get_server_args().enable_aiter_allreduce_fusion + and get_exec().comm.enable_aiter_allreduce_fusion ) ) and (not self.is_last_layer) @@ -1151,7 +1157,7 @@ class CommunicateWithAllReduceAndLayerNormFn: if not handled: quantize_communications = ( not forward_batch.forward_mode.is_decode_or_idle() - and get_server_args().enable_quant_communications + and get_exec().comm.enable_quant_communications ) if quantize_communications: hidden_states = attention_tensor_model_parallel_quant_all_reduce( diff --git a/python/sglang/srt/layers/cp/zigzag.py b/python/sglang/srt/layers/cp/zigzag.py index 7b5154a96..6ad74e186 100644 --- a/python/sglang/srt/layers/cp/zigzag.py +++ b/python/sglang/srt/layers/cp/zigzag.py @@ -53,7 +53,7 @@ from sglang.srt.layers.dp_attention import ( ) from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.model_executor.forward_context import get_token_to_kv_pool -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_device, get_parallel @dataclass @@ -208,10 +208,8 @@ class ZigzagCPStrategy(ContextParallelStrategy): actual_seq_q_prev_list.append(block_sizes[cp_rank]) actual_seq_q_next_list.append(block_sizes[cp_segment_num - cp_rank - 1]) - from sglang.srt.runtime_context import get_server_args - try: - device = torch.device(get_server_args().device) + device = torch.device(get_device().device) except Exception: device = torch.device("cuda" if torch.cuda.is_available() else "cpu") cu_prev = [0] + list(accumulate(actual_seq_q_prev_list)) diff --git a/python/sglang/srt/layers/dcp/planner.py b/python/sglang/srt/layers/dcp/planner.py index 1a9caba76..d41d23630 100644 --- a/python/sglang/srt/layers/dcp/planner.py +++ b/python/sglang/srt/layers/dcp/planner.py @@ -26,7 +26,7 @@ from sglang.kernels.ops.attention.dcp_kernels import ( ) from sglang.srt.layers.dcp.layout import update_local_kv_lens_for_dcp from sglang.srt.layers.dcp.metadata import DecodeContextParallelMetadata -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_device, get_parallel def prepare_decode_context_parallel_metadata( @@ -53,12 +53,12 @@ def prepare_decode_context_parallel_metadata( extend_prefix_starts = torch.zeros( len(seq_lens), dtype=torch.int32, - device=get_server_args().device, + device=get_device().device, ) extend_cu_prefix_lens = torch.zeros( len(seq_lens) + 1, dtype=torch.int32, - device=get_server_args().device, + device=get_device().device, ) extend_cu_prefix_lens[1:] = torch.cumsum(extend_prefix_lens, dim=0) extend_cu_prefix_lens = extend_cu_prefix_lens[:-1] @@ -67,7 +67,7 @@ def prepare_decode_context_parallel_metadata( dcp_prefix_kv_indices = torch.empty( sum(extend_prefix_lens_cpu), dtype=torch.int32, - device=get_server_args().device, + device=get_device().device, ) create_chunked_prefix_cache_kv_indices_fn[(len(seq_lens),)]( req_to_token, @@ -81,20 +81,20 @@ def prepare_decode_context_parallel_metadata( dcp_kv_indptr = torch.zeros( len(seq_lens) + 1, dtype=torch.int32, - device=get_server_args().device, + device=get_device().device, ) dcp_kv_indptr[1:] = seq_lens.cumsum(dim=0) dcp_kv_indptr = dcp_kv_indptr[: (len(seq_lens) + 1)] dcp_kv_indices = torch.zeros( seq_lens_sum, dtype=torch.int32, - device=get_server_args().device, + device=get_device().device, ) extend_cu_lens = torch.zeros( len(seq_lens) + 1, dtype=torch.int32, - device=get_server_args().device, + device=get_device().device, ) extend_cu_lens[1:] = torch.cumsum(extend_seq_lens, dim=0) extend_cu_lens = extend_cu_lens[:-1] diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index f7643efe2..67661147f 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -32,7 +32,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( Phase, check_cuda_graph_backend, ) -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -223,7 +223,7 @@ def _forward_with_allreduce_fusion( return fused_result # For AITER route, preserve correctness when fused path is unavailable. - if _use_aiter and get_server_args().enable_aiter_allreduce_fusion: + if _use_aiter and get_exec().comm.enable_aiter_allreduce_fusion: x = tensor_model_parallel_all_reduce(x) return norm_module.forward(x, residual, None) @@ -425,7 +425,7 @@ class RMSNorm(MultiPlatformOp): if ( residual is not None or self.cast_x_before_out_mul - or get_server_args().rl_on_policy_target == "fsdp" + or get_exec().deterministic.rl_on_policy_target == "fsdp" ): return self.forward_native(x, residual, post_residual_addition) out = rms_norm_batch_invariant( @@ -532,7 +532,7 @@ class RMSNorm(MultiPlatformOp): if ( residual is not None or self.cast_x_before_out_mul - or get_server_args().rl_on_policy_target == "fsdp" + or get_exec().deterministic.rl_on_policy_target == "fsdp" or (self._fused_pad_kernel is not None and self.x_pad_to_multiple > 0) ): return self.forward_native(x, residual, post_residual_addition) @@ -593,7 +593,7 @@ class RMSNorm(MultiPlatformOp): if ( residual is not None or self.cast_x_before_out_mul - or get_server_args().rl_on_policy_target == "fsdp" + or get_exec().deterministic.rl_on_policy_target == "fsdp" ): return self.forward_native(x, residual, post_residual_addition) return rms_norm_batch_invariant( @@ -720,7 +720,10 @@ class RMSNorm(MultiPlatformOp): if self.variance_size_override is not None: return self.forward_native(x, residual, post_residual_addition) if is_batch_invariant_mode_enabled(): - if residual is not None or get_server_args().rl_on_policy_target == "fsdp": + if ( + residual is not None + or get_exec().deterministic.rl_on_policy_target == "fsdp" + ): return self.forward_native(x, residual, post_residual_addition) return rms_norm_batch_invariant( x, diff --git a/python/sglang/srt/layers/linear.py b/python/sglang/srt/layers/linear.py index 12acd2fbd..2886cb01e 100644 --- a/python/sglang/srt/layers/linear.py +++ b/python/sglang/srt/layers/linear.py @@ -39,7 +39,7 @@ from sglang.srt.layers.parameter import ( _ColumnvLLMParameter, ) from sglang.srt.layers.utils import pad_or_narrow_weight -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.utils import get_bool_env_var, is_cpu, is_hip, is_npu, set_weight_attrs if TYPE_CHECKING: @@ -1597,7 +1597,7 @@ class RowParallelLinear(LinearBase): quantize_communications = ( ( not forward_batch.forward_mode.is_decode_or_idle() - and get_server_args().enable_quant_communications + and get_exec().comm.enable_quant_communications ) if forward_batch is not None else False diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index 641644c7d..499840fe2 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -51,7 +51,7 @@ from sglang.srt.model_executor.forward_batch_info import ( ForwardBatch, ForwardMode, ) -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args from sglang.srt.utils.common import ( is_cpu, is_npu, @@ -350,7 +350,7 @@ class LogitsProcessor(nn.Module): self.vocab_size = config.vocab_size self.logit_scale = logit_scale self.use_attn_tp_group = get_server_args().enable_dp_lm_head - self.use_fp32_lm_head = get_server_args().enable_fp32_lm_head + self.use_fp32_lm_head = get_exec().features.enable_fp32_lm_head if self.use_attn_tp_group: self.attn_tp_size = get_parallel().attn_tp_size self.do_tensor_parallel_all_gather = ( @@ -374,8 +374,8 @@ class LogitsProcessor(nn.Module): self.final_logit_softcapping = None self.return_full_logits = return_full_logits - self.enable_mis = get_server_args().enable_mis - self.rl_on_policy_target = get_server_args().rl_on_policy_target + self.enable_mis = get_exec().features.enable_mis + self.rl_on_policy_target = get_exec().deterministic.rl_on_policy_target self._logits_gatherer = triton_symm_mem_ag.MultimemAllGatherer( max_tokens=triton_symm_mem_ag.recommended_max_tokens( diff --git a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py index 1160b8637..03d0ccc2d 100644 --- a/python/sglang/srt/layers/moe/fused_moe_triton/layer.py +++ b/python/sglang/srt/layers/moe/fused_moe_triton/layer.py @@ -71,6 +71,7 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo ) from sglang.srt.model_loader.weight_utils import narrow_padded_param_and_loaded_weight from sglang.srt.runtime_context import ( + get_exec, get_global_dwdp_manager, get_parallel, get_server_args, @@ -260,7 +261,7 @@ class FusedMoE(torch.nn.Module): self._num_global_routed = num_experts - num_shared_slots server_args = get_server_args() - if server_args.ep_join_mode == "scale": + if get_exec().moe.ep_join_mode == "scale": storage_ep_size = server_args.elastic_ep_initial_size assert storage_ep_size is not None self._expert_storage_rank = ( @@ -359,7 +360,7 @@ class FusedMoE(torch.nn.Module): print_info_once( "FlashInfer TRTLLM MoE deferred finalize is " f"{'enabled' if self.supports_deferred_finalize else 'disabled'} " - f"(moe_runner_backend={server_args.moe_runner_backend}, " + f"(moe_runner_backend={get_exec().moe.moe_runner_backend}, " f"quant_method={type(self.quant_method).__name__})." ) diff --git a/python/sglang/srt/layers/moe/hash_topk.py b/python/sglang/srt/layers/moe/hash_topk.py index ebcf45702..972510be7 100644 --- a/python/sglang/srt/layers/moe/hash_topk.py +++ b/python/sglang/srt/layers/moe/hash_topk.py @@ -22,6 +22,7 @@ from sglang.srt.layers.moe.topk import ( remap_topk_for_per_rank_shared_slots, ) from sglang.srt.layers.moe.utils import has_per_rank_fused_shared_slots +from sglang.srt.runtime_context import get_exec from sglang.srt.utils import is_hip, is_npu logger = logging.getLogger(__name__) @@ -44,10 +45,9 @@ class HashTopK(nn.Module): ): super().__init__() self.layer_id = layer_id - from sglang.srt.runtime_context import get_server_args self.enable_waterfill = ( - num_fused_shared_experts > 0 and get_server_args().enable_waterfill + num_fused_shared_experts > 0 and get_exec().moe.enable_waterfill ) self.waterfill_balancer = None diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py index a22880bf9..964682ab8 100644 --- a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe.py @@ -28,7 +28,7 @@ from sglang.srt.environ import envs from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig from sglang.srt.layers.moe.utils import get_moe_padding_size -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -531,7 +531,7 @@ def _fused_moe_kernel_sequence( out_hidden_states = torch.empty_like(hidden_states) use_fused_moe_sum_all_reduce = ( - get_server_args().enable_fused_moe_sum_all_reduce + get_exec().moe.enable_fused_moe_sum_all_reduce and (not no_combine) and (topk > 2) and (not use_int8_w8a16) diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe_triton_config.py b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe_triton_config.py index 210247f86..c4ad7d49d 100644 --- a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe_triton_config.py +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe_triton_config.py @@ -9,7 +9,7 @@ from typing import Any, Dict, List, Optional, Tuple import torch import triton -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec from sglang.srt.utils import get_device_name, is_hip logger = logging.getLogger(__name__) @@ -69,7 +69,7 @@ def get_moe_configs( kernel on a given batch size bs, the closest batch size in the grid should be picked and the associated configuration chosen to invoke the kernel. """ - if get_server_args().enable_deterministic_inference: + if get_exec().deterministic.enable_deterministic_inference: logger.warning( "Deterministic inference is enabled, using default MoE kernel config." ) @@ -187,7 +187,7 @@ def get_default_config( is_marlin: bool, block_shape: Optional[List[int]] = None, ) -> Dict[str, int]: - if get_server_args().enable_deterministic_inference: + if get_exec().deterministic.enable_deterministic_inference: config = { "BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 64, diff --git a/python/sglang/srt/layers/moe/token_dispatcher/flashinfer.py b/python/sglang/srt/layers/moe/token_dispatcher/flashinfer.py index 53e1b086a..670fb428c 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/flashinfer.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/flashinfer.py @@ -27,7 +27,7 @@ from sglang.srt.layers.moe.topk import ( TopKOutputChecker, ) from sglang.srt.layers.moe.utils import get_moe_runner_backend -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_schedule, get_spec from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils import get_int_env_var @@ -123,7 +123,7 @@ class FlashinferDispatcher(BaseDispatcher): # max_running_requests is not yet resolved at model-construction time, # so we use 4096 as a floor to cover decode batches and _dummy_run # (which warms up at batch_size = req_to_token_pool.size). - cps = get_server_args().chunked_prefill_size + cps = get_schedule().chunked_prefill_size default_max_tokens = max(cps if cps and cps > 0 else 4096, 4096) self.max_num_tokens = get_int_env_var( "SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK", @@ -132,7 +132,7 @@ class FlashinferDispatcher(BaseDispatcher): # Calculate workspace size. For eagle mode, use the larger workspace size since nextn layer will be unquantized. speculative_algo = SpeculativeAlgorithm.from_string( - get_server_args().speculative_algorithm + get_spec().speculative_algorithm ) if MOE_NVFP4_DISPATCH and not speculative_algo.is_eagle(): total_dispatch_payload_size_per_token = ( diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py index bd166b363..768e03137 100644 --- a/python/sglang/srt/layers/moe/topk.py +++ b/python/sglang/srt/layers/moe/topk.py @@ -32,7 +32,7 @@ from typing import ( import torch import torch.nn.functional as F -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_exec, get_lora, get_parallel try: from triton_kernels.matmul_ogs import GatherIndx, RoutingData, ScatterIndx @@ -430,10 +430,9 @@ class TopK(MultiPlatformOp): assert num_expert_group is not None and topk_group is not None self.layer_id = layer_id - from sglang.srt.runtime_context import get_server_args self.enable_waterfill = ( - num_fused_shared_experts > 0 and get_server_args().enable_waterfill + num_fused_shared_experts > 0 and get_exec().moe.enable_waterfill ) self.waterfill_balancer = None @@ -507,9 +506,8 @@ class TopK(MultiPlatformOp): # ===== TO BE REFACTORED ==== elif get_moe_runner_backend().is_experimental_sgl_trtllm(): try: - from sglang.srt.runtime_context import get_server_args - use_standard_for_lora = bool(get_server_args().enable_lora) + use_standard_for_lora = bool(get_lora().enable_lora) except ValueError: use_standard_for_lora = False output_format = ( @@ -1362,9 +1360,9 @@ def _eplb_remap_enabled() -> bool: # there is no EPLB mapping, so the remap must be skipped. return False return ( - server_args.enable_eplb - or server_args.init_expert_location != "trivial" - or server_args.ep_num_redundant_experts > 0 + get_exec().moe.enable_eplb + or get_exec().moe.init_expert_location != "trivial" + or get_exec().moe.ep_num_redundant_experts > 0 ) diff --git a/python/sglang/srt/layers/moe/utils.py b/python/sglang/srt/layers/moe/utils.py index 09bc8834f..824beff64 100644 --- a/python/sglang/srt/layers/moe/utils.py +++ b/python/sglang/srt/layers/moe/utils.py @@ -12,7 +12,7 @@ from sglang.srt.environ import envs from sglang.srt.layers.dp_attention import ( is_dp_attention_enabled, ) -from sglang.srt.runtime_context import get_flags, get_forward, get_parallel +from sglang.srt.runtime_context import get_exec, get_flags, get_forward, get_parallel from sglang.srt.utils import is_cuda, is_npu _is_npu = is_npu() @@ -239,8 +239,8 @@ def get_deepep_output_dtype(self) -> DispatcherOutputDtype: # 0. Parse server argument. server_args = get_server_args() - if server_args and server_args.deepep_dispatcher_output_dtype != "auto": - return DispatcherOutputDtype(server_args.deepep_dispatcher_output_dtype) + if server_args and get_exec().moe.deepep_dispatcher_output_dtype != "auto": + return DispatcherOutputDtype(get_exec().moe.deepep_dispatcher_output_dtype) # 1. Parse deprecated environment variables. if envs.SGLANG_DEEPEP_BF16_DISPATCH.get(): diff --git a/python/sglang/srt/layers/quantization/fp8_utils.py b/python/sglang/srt/layers/quantization/fp8_utils.py index 84539c9b2..97762fd98 100755 --- a/python/sglang/srt/layers/quantization/fp8_utils.py +++ b/python/sglang/srt/layers/quantization/fp8_utils.py @@ -13,7 +13,7 @@ from sglang.kernels.ops.quantization.fp8_kernel import ( ) from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers.quantization.mxfp4_tensor import MXFP4QuantizeUtil -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.utils.common import torch_release if TYPE_CHECKING: @@ -34,7 +34,6 @@ from sglang.kernels.ops.quantization.fp8_kernel import ( w8a8_block_fp8_matmul_deepgemm, w8a8_block_fp8_matmul_triton, ) -from sglang.srt.runtime_context import get_server_args from sglang.srt.utils import ( ceil_align, ceil_div, @@ -1844,7 +1843,7 @@ def apply_fp8_linear( if ( input_scale is not None and input_scale.numel() == 1 - and get_server_args().cuda_graph_config.prefill.tc_compiler == "inductor" + and get_exec().graph.cuda_graph_config.prefill.tc_compiler == "inductor" ): qinput = ( (input_2d * input_scale.reciprocal()) diff --git a/python/sglang/srt/layers/quantization/mxfp4.py b/python/sglang/srt/layers/quantization/mxfp4.py index 332ad7eda..a5bb043c5 100644 --- a/python/sglang/srt/layers/quantization/mxfp4.py +++ b/python/sglang/srt/layers/quantization/mxfp4.py @@ -48,7 +48,7 @@ from sglang.srt.layers.quantization.base_config import ( QuantizeMethodBase, ) from sglang.srt.layers.quantization.utils import is_layer_skipped -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec from sglang.srt.utils import ( cpu_has_amx_support, is_cpu, @@ -333,7 +333,7 @@ class Mxfp4MoEMethod(FusedMoEMethodBase): self.use_flashinfer = get_moe_runner_backend().is_flashinfer_mxfp4() self.use_marlin = get_moe_runner_backend().is_marlin() self.flashinfer_mxfp4_moe_precision = ( - get_server_args().flashinfer_mxfp4_moe_precision + get_exec().moe.flashinfer_mxfp4_moe_precision ) # When `flashinfer_mxfp4` is enabled, dispatch to one of three FlashInfer # entry points depending on the GPU: diff --git a/python/sglang/srt/layers/quantization/mxfp4_flashinfer_trtllm_moe.py b/python/sglang/srt/layers/quantization/mxfp4_flashinfer_trtllm_moe.py index 0973f05e6..dbdcc713b 100644 --- a/python/sglang/srt/layers/quantization/mxfp4_flashinfer_trtllm_moe.py +++ b/python/sglang/srt/layers/quantization/mxfp4_flashinfer_trtllm_moe.py @@ -14,7 +14,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import ( ) from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.moe.utils import RoutingMethodType -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec from sglang.srt.utils import ( is_flashinfer_available, log_info_on_rank0, @@ -51,7 +51,7 @@ class Mxfp4FlashinferTrtllmMoEMethod: self._fp8 = fp8_method self.prefix = prefix self.flashinfer_mxfp4_moe_precision = ( - get_server_args().flashinfer_mxfp4_moe_precision + get_exec().moe.flashinfer_mxfp4_moe_precision ) def create_moe_runner(self, layer, moe_runner_config): diff --git a/python/sglang/srt/layers/rotary_embedding/base.py b/python/sglang/srt/layers/rotary_embedding/base.py index 8443bb71d..f009f4d1f 100644 --- a/python/sglang/srt/layers/rotary_embedding/base.py +++ b/python/sglang/srt/layers/rotary_embedding/base.py @@ -11,7 +11,7 @@ from sglang.srt.environ import envs from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.platforms import current_platform -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec from sglang.srt.utils import ( cpu_has_amx_support, get_bool_env_var, @@ -129,7 +129,7 @@ class RotaryEmbedding(MultiPlatformOp): self._apply_rotary_emb_wrapped = apply_rotary_emb # XXX (MUSA): Implement sgl_kernel.rotary_embedding support for MUSA backend - if get_server_args().rl_on_policy_target is not None or _is_musa: + if get_exec().deterministic.rl_on_policy_target is not None or _is_musa: self._forward_method = self.forward_native self._apply_rotary_emb_wrapped = torch.compile( dynamic=True, @@ -153,7 +153,7 @@ class RotaryEmbedding(MultiPlatformOp): # create the cache on GPU for faster initialization. This may cause # a slight numerical difference between the HF implementation and ours. init_device = ( - "cpu" if get_server_args().rl_on_policy_target is not None else None + "cpu" if get_exec().deterministic.rl_on_policy_target is not None else None ) inv_freq = 1.0 / ( base @@ -164,7 +164,7 @@ class RotaryEmbedding(MultiPlatformOp): / self.rotary_dim ) ) - if get_server_args().rl_on_policy_target is not None: + if get_exec().deterministic.rl_on_policy_target is not None: inv_freq = inv_freq.cuda() return inv_freq diff --git a/python/sglang/srt/layers/rotary_embedding/mrope.py b/python/sglang/srt/layers/rotary_embedding/mrope.py index 68346e505..de0de81a3 100644 --- a/python/sglang/srt/layers/rotary_embedding/mrope.py +++ b/python/sglang/srt/layers/rotary_embedding/mrope.py @@ -18,7 +18,7 @@ from sglang.srt.layers.rotary_embedding.yarn import ( yarn_get_mscale_simple, yarn_linear_ramp_mask, ) -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec, get_server_args from sglang.srt.utils import ( cpu_has_amx_support, is_cuda, @@ -132,7 +132,7 @@ class MRotaryEmbedding(RotaryEmbedding): self.register_buffer("axis_map", axis_map, persistent=False) else: self.axis_map = None - if get_server_args().rl_on_policy_target is not None: + if get_exec().deterministic.rl_on_policy_target is not None: self._forward_method = self.forward_native def get_cos_sin_with_position(self, positions): diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index e8a1d7161..dc9e5128d 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -15,7 +15,7 @@ from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.logprob_processor import ( OutputLogprobProcessor, ) -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_params import TOP_K_ALL from sglang.srt.utils.async_probe import sanitize_nan_logits @@ -74,12 +74,14 @@ class Sampler(nn.Module): if is_dp_attention_enabled(): self.tp_sync_group = get_parallel().attn_tp_group.device_group - self.rl_on_policy_target = get_server_args().rl_on_policy_target + self.rl_on_policy_target = get_exec().deterministic.rl_on_policy_target # In RL on-policy mode, deterministic inference is automatically enabled. - self.enable_deterministic = get_server_args().enable_deterministic_inference + self.enable_deterministic = ( + get_exec().deterministic.enable_deterministic_inference + ) # In RL on-policy mode, we use log_softmax to compute logprobs to match the trainer. self.use_log_softmax_logprob = self.rl_on_policy_target is not None - self.use_ascend_backend = get_server_args().sampling_backend == "ascend" + self.use_ascend_backend = get_exec().kernel.sampling_backend == "ascend" self.output_logprob_processor = OutputLogprobProcessor() @@ -260,7 +262,7 @@ class Sampler(nn.Module): positions=positions, ) else: - backend = get_server_args().sampling_backend + backend = get_exec().kernel.sampling_backend if backend == "flashinfer": assert ( sampling_info.sampling_seed is None @@ -540,7 +542,7 @@ def create_sampler(backend: Optional[str] = None) -> "Sampler": """Create a sampler honoring custom backend registrations.""" server_args = get_server_args() - backend = backend or (server_args.sampling_backend if server_args else None) + backend = backend or (get_exec().kernel.sampling_backend if server_args else None) if backend in _CUSTOM_SAMPLER_FACTORIES: sampler = _CUSTOM_SAMPLER_FACTORIES[backend]() diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index 7bd856698..81f83f03f 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -48,7 +48,7 @@ from sglang.srt.managers.scheduler import run_scheduler_process from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread from sglang.srt.observability.req_time_stats import DPControllerReqTimeStats from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info -from sglang.srt.runtime_context import publish +from sglang.srt.runtime_context import get_exec, publish from sglang.srt.server_args import ( DP_ATTENTION_HANDSHAKE_PORT_DELTA, PortArgs, @@ -232,7 +232,7 @@ class DataParallelController: sock_send(worker, obj) def update_active_ranks(self, ranks: ActiveRanksOutput): - if self.server_args.elastic_ep_backend is not None: + if get_exec().moe.elastic_ep_backend is not None: if len(ranks.status) != self.max_dp_size: logger.warning( "[Elastic EP][DPC] active rank status len=%d != max_dp_size=%d; " @@ -485,7 +485,7 @@ class DataParallelController: logger.debug("Worker port broadcast completed") return worker_ports finally: - if self.server_args.elastic_ep_backend is None: + if get_exec().moe.elastic_ep_backend is None: rep_socket.close() else: threading.Thread( diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index d8cf47e1f..fe2afd353 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -33,7 +33,12 @@ from sglang.srt.managers.schedule_batch import ( from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.multimodal.evs import EVSEmbeddingResult -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_disagg, + get_parallel, + get_schedule, + get_server_args, +) from sglang.srt.utils import flatten_nested_list, is_hip, is_npu, print_warning_once from sglang.srt.utils.stale_shm_cleanup import make_shm_name from sglang.utils import logger @@ -931,7 +936,7 @@ def _adjust_embedding_length( f"tokens from multimodal embeddings." ) if num_mm_tokens_in_input_ids < num_mm_tokens_in_embedding: - chunked_prefill_size = get_server_args().chunked_prefill_size + chunked_prefill_size = get_schedule().chunked_prefill_size if chunked_prefill_size != -1: logger.warning( "You may want to avoid this issue by raising `chunked_prefill_size`, or disabling chunked prefill" @@ -1295,7 +1300,7 @@ def general_mm_embed_routine( # encoder/ViT execution and multimodal feature placement, while # the language model range below excludes both. with torch.profiler.record_function("sglang.vlm.mm_embedding"): - if server_args and server_args.enable_adaptive_dispatch_to_encoder: + if server_args and get_disagg().enable_adaptive_dispatch_to_encoder: # Split by precomputed vs non-precomputed so get_embedding_and_mask only sees uniform batches input_embeds, other_info = _embed_mm_inputs_with_split( mm_inputs_list=mm_inputs_list, @@ -1340,7 +1345,7 @@ def general_mm_embed_routine( feature = getattr(mm_item, "feature", None) if isinstance(feature, torch.Tensor) and feature.is_cuda: mm_item.feature = feature.to("cpu", non_blocking=True) - if get_server_args().language_only: + if get_disagg().language_only: precomputed_embeddings = getattr( mm_item, "precomputed_embeddings", None ) diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 98b4341cb..d63fb2af6 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -2,6 +2,7 @@ from __future__ import annotations from sglang.srt.dllm.config import DllmConfig from sglang.srt.model_executor.forward_batch_info import ForwardBatch +from sglang.srt.runtime_context import get_exec, get_schedule, get_serving, get_spec from sglang.srt.utils.common import ( Range, ceil_align, @@ -1097,7 +1098,7 @@ class Req(ReqDllmMixin): """Check if this request is prefill-only (no token generation needed).""" # NOTE: when spec is enabled, prefill_only optimizations are disabled - spec_alg = get_server_args().speculative_algorithm + spec_alg = get_spec().speculative_algorithm return self.sampling_params.max_new_tokens == 0 and spec_alg is None @property @@ -1118,7 +1119,7 @@ class Req(ReqDllmMixin): def effective_kv_committed_len(self) -> int: # Report only the prompt prefix so thinking + answer fall into the # overallocated range and are reclaimed by release_kv_cache. #22373. - if get_server_args().strip_thinking_cache and self.reasoning_tokens > 0: + if get_serving().strip_thinking_cache and self.reasoning_tokens > 0: return min(self.kv_committed_len, len(self.origin_input_ids)) return self.kv_committed_len @@ -2922,7 +2923,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): ) if server_args.enable_mamba_extra_buffer(): - mamba_track_interval = server_args.mamba_track_interval + mamba_track_interval = get_exec().mamba.mamba_track_interval if len(self.reqs) == 0: self.mamba_track_indices = torch.empty( @@ -3168,8 +3169,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): continue else: pre_len = ( - pre_len - server_args.chunked_prefill_size - if server_args.chunked_prefill_size > 0 + pre_len - get_schedule().chunked_prefill_size + if get_schedule().chunked_prefill_size > 0 else pre_len ) self._evict_swa(req, pre_len) diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index 3cc7ce842..39bc02b78 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -5,6 +5,7 @@ from array import array from sglang.srt.environ import envs from sglang.srt.managers.prefill_delayer import PrefillDelayerSinglePassExecutor +from sglang.srt.runtime_context import get_disagg from sglang.srt.utils import get_bool_env_var _ROUTING_KEY_POLICY_DEBUG_LOG = get_bool_env_var("SGLANG_ROUTING_KEY_POLICY_DEBUG_LOG") @@ -56,7 +57,6 @@ from sglang.srt.mem_cache.multi_ended_allocator import ( UnifiedMambaTokenToKVPoolAllocator, ) from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode -from sglang.srt.runtime_context import get_server_args from sglang.srt.server_args import ServerArgs if TYPE_CHECKING: @@ -195,7 +195,7 @@ class SchedulePolicy: if ( not isinstance(policy, CacheAwarePolicy) and self.tree_cache.supports_fast_match_prefix() - and get_server_args().disaggregation_mode != "decode" + and get_disagg().disaggregation_mode != "decode" ): for r in waiting_queue: match_prefix_for_req(self.tree_cache, r, include_req=True) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 456ac76e2..afea8748c 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -27,6 +27,20 @@ from functools import partial from http import HTTPStatus from typing import Any, Deque, Dict, List, Optional, Tuple, Union +from sglang.srt.runtime_context import ( + get_device, + get_disagg, + get_exec, + get_lora, + get_memory, + get_mm, + get_model, + get_observability, + get_schedule, + get_serving, + get_spec, +) + from sglang.srt.utils.common import suppress_noisy_warnings # isort: skip suppress_noisy_warnings() @@ -482,9 +496,9 @@ class Scheduler( attn_tp_cpu_group=self.attn_tp_cpu_group, tp_cpu_group=self.tp_cpu_group, attn_cp_cpu_group=self.attn_cp_cpu_group, - enable_metrics=self.server_args.enable_metrics, + enable_metrics=get_observability().enable_metrics, enable_kv_cache_events=bool( - self.server_args.kv_events_config + get_observability().kv_events_config and self.ps.pp_rank == 0 and self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0 @@ -526,8 +540,8 @@ class Scheduler( self.init_hisparse_coordinator() if ( - self.server_args.disaggregation_mode == "decode" - and self.server_args.disaggregation_decode_enable_offload_kvcache + get_disagg().disaggregation_mode == "decode" + and get_disagg().disaggregation_decode_enable_offload_kvcache ): self.decode_offload_manager = DecodeKVCacheOffloadManager( req_to_token_pool=self.req_to_token_pool, @@ -642,7 +656,7 @@ class Scheduler( self.dllm_config = ( # For diffusion LLM DllmConfig.from_server_args(self.server_args) - if self.server_args.dllm_algorithm is not None + if get_exec().dllm.dllm_algorithm is not None else None ) @@ -671,10 +685,10 @@ class Scheduler( port_args=port_args, is_rank_zero=is_rank_zero, skip_tokenizer_init=self.server_args.skip_tokenizer_init, - metrics_enabled=self.server_args.enable_metrics + metrics_enabled=get_observability().enable_metrics and ( self.ps.attn_tp_rank == 0 - or self.server_args.enable_metrics_for_all_schedulers + or get_observability().enable_metrics_for_all_schedulers ), enable_scripted_runtime=envs.SGLANG_TEST_SCRIPTED_RUNTIME.get(), ) @@ -693,7 +707,7 @@ class Scheduler( port_args, self.ps.dp_size, dp_rank, - publish_interval=self.server_args.load_snapshot_publish_interval, + publish_interval=get_observability().load_snapshot_publish_interval, ) except Exception as e: logger.warning("load snapshot writer init failed: %s", e) @@ -703,7 +717,7 @@ class Scheduler( self.ps.pp_rank == 0 and self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0 - and self.server_args.sleep_on_idle + and get_device().sleep_on_idle ): self.idle_sleeper = IdleSleeper( sockets=[ @@ -737,22 +751,22 @@ class Scheduler( else: if self.model_config.is_multimodal: self.processor = get_processor( - server_args.tokenizer_path, - tokenizer_mode=server_args.tokenizer_mode, - trust_remote_code=server_args.trust_remote_code, - revision=server_args.revision, - use_fast=not server_args.disable_fast_image_processor, - tokenizer_backend=server_args.tokenizer_backend, - model_name=server_args.model_path, + get_serving().tokenizer_path, + tokenizer_mode=get_serving().tokenizer_mode, + trust_remote_code=get_model().trust_remote_code, + revision=get_model().revision, + use_fast=not get_mm().disable_fast_image_processor, + tokenizer_backend=get_serving().tokenizer_backend, + model_name=get_model().model_path, ) self.tokenizer = get_tokenizer_from_processor(self.processor) else: self.tokenizer = get_tokenizer( - server_args.tokenizer_path, - tokenizer_mode=server_args.tokenizer_mode, - trust_remote_code=server_args.trust_remote_code, - revision=server_args.revision, - tokenizer_backend=server_args.tokenizer_backend, + get_serving().tokenizer_path, + tokenizer_mode=get_serving().tokenizer_mode, + trust_remote_code=get_model().trust_remote_code, + revision=get_model().revision, + tokenizer_backend=get_serving().tokenizer_backend, ) # Load multimodal processor for M-RoPE fallback computation. @@ -774,9 +788,9 @@ class Scheduler( ) # Set reasoning_parser and think_end_id if --reasoning_parser is enabled - if self.server_args.reasoning_parser and self.tokenizer: + if get_serving().reasoning_parser and self.tokenizer: reasoning_parser = ReasoningParser( - model_type=self.server_args.reasoning_parser, + model_type=get_serving().reasoning_parser, stream_reasoning=False, tokenizer=self.tokenizer, ) @@ -847,7 +861,7 @@ class Scheduler( target_worker=self.tp_worker, ) - if self.server_args.speculative_draft_load_format is not None: + if get_spec().speculative_draft_load_format is not None: # Write the draft load_format onto server_args (not just the bag): # the draft worker is built from a copy of self.server_args and # build_load_config reads server_args.load_format, so a bag-only @@ -855,10 +869,10 @@ class Scheduler( # format. self.server_args.override( "scheduler.draft_load_format", - load_format=self.server_args.speculative_draft_load_format, + load_format=get_spec().speculative_draft_load_format, ) logger.info( - f"Using draft model load_format: '{self.server_args.speculative_draft_load_format}'" + f"Using draft model load_format: '{get_spec().speculative_draft_load_format}'" ) DraftWorkerClass = self.spec_algorithm.create_worker(self.server_args) @@ -925,8 +939,8 @@ class Scheduler( model_runner.post_capture_resize_kv_pool() if ( - self.server_args.elastic_ep_backend is not None - and self.server_args.ep_join_mode == "recover" + get_exec().moe.elastic_ep_backend is not None + and get_exec().moe.ep_join_mode == "recover" ): model_runner.post_capture_elastic_ep_recover() @@ -955,7 +969,7 @@ class Scheduler( # --min-free-slots-delay. Built independently of the prefill delayer. self.min_free_slots_delayer: Optional[MinFreeSlotsDelayer] = None min_free_slots = resolve_min_free_slots( - self.server_args.min_free_slots_delay, + get_schedule().min_free_slots_delay, self.max_running_requests, is_dflash_family=self.spec_algorithm.is_dflash_family(), ) @@ -1001,14 +1015,14 @@ class Scheduler( if self.ps.tp_rank == 0: logger.info( f"max_total_num_tokens={self.max_total_num_tokens}, " - f"chunked_prefill_size={self.server_args.chunked_prefill_size}, " + f"chunked_prefill_size={get_schedule().chunked_prefill_size}, " f"max_prefill_tokens={self.max_prefill_tokens}, " f"max_running_requests={self.max_running_requests}, " f"context_len={self.model_config.context_len}, " f"{'available_cpu_mem' if self.device == 'cpu' else 'available_gpu_mem'}={avail_mem:.2f} GB" ) - if self.server_args.enable_metrics: + if get_observability().enable_metrics: self.metrics_collector.emit_constants( max_total_num_tokens=self.max_total_num_tokens, # TODO: max_running_requests_under_SLO has no setter — dead chain. @@ -1055,7 +1069,7 @@ class Scheduler( self._engine_paused = False def init_chunked_prefill(self): - self.chunked_prefill_size = self.server_args.chunked_prefill_size + self.chunked_prefill_size = get_schedule().chunked_prefill_size uses_transformers_backend = ( get_resolved_model_impl(self.model_config) == ModelImpl.TRANSFORMERS ) @@ -1075,13 +1089,12 @@ class Scheduler( self.chunked_req = None self._pending_chunked_abort_req = None self.is_mixed_chunk = ( - self.chunked_prefill_size is not None - and self.server_args.enable_mixed_chunk + self.chunked_prefill_size is not None and get_schedule().enable_mixed_chunk ) # Init the dynamic chunking predictor for PP self.enable_dynamic_chunking = ( - self.server_args.enable_dynamic_chunking and self.ps.pp_size > 1 + get_schedule().enable_dynamic_chunking and self.ps.pp_size > 1 ) if self.enable_dynamic_chunking: try: @@ -1117,8 +1130,8 @@ class Scheduler( ) self.prefill_delayer: Optional[PrefillDelayer] = None self.max_prefill_bs: int = 0 - if self.server_args.enable_prefill_delayer: - if self.server_args.disaggregation_mode == "decode": + if get_schedule().enable_prefill_delayer: + if get_disagg().disaggregation_mode == "decode": logger.info( "Ignoring --enable-prefill-delayer on decode engine " "(no prefill scheduling path; delayer would be a no-op)." @@ -1135,15 +1148,15 @@ class Scheduler( if self.metrics_reporter.enable_metrics else None ), - max_delay_passes=self.server_args.prefill_delayer_max_delay_passes, - token_usage_low_watermark=self.server_args.prefill_delayer_token_usage_low_watermark, + max_delay_passes=get_schedule().prefill_delayer_max_delay_passes, + token_usage_low_watermark=get_schedule().prefill_delayer_token_usage_low_watermark, device=self.tp_group.device, ) # NOTE: preemption is enabled by default for priority scheduling. self.enable_priority_preemption = ( self.enable_priority_scheduling - and not self.server_args.disable_priority_preemption + and not get_schedule().disable_priority_preemption ) self.new_token_ratio_tracker = NewTokenRatioTracker.from_server_args( @@ -1159,12 +1172,12 @@ class Scheduler( def init_watch_dog_memory_saver_input_blocker(self): # Start watchdog thread self.watchdog = create_scheduler_watchdog( - self, watchdog_timeout=self.server_args.watchdog_timeout + self, watchdog_timeout=get_device().watchdog_timeout ) # Init memory saver, profiler and metric stats self.memory_saver_adapter = TorchMemorySaverAdapter.create( - enable=self.server_args.enable_memory_saver + enable=get_exec().features.enable_memory_saver ) # Init recv skipper and input blocker @@ -1186,11 +1199,9 @@ class Scheduler( self.disagg_decode_prealloc_queue = None self.disagg_decode_transfer_queue = None - self.disaggregation_mode = DisaggregationMode( - self.server_args.disaggregation_mode - ) + self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode) self.transfer_backend = TransferBackend( - self.server_args.disaggregation_transfer_backend + get_disagg().disaggregation_transfer_backend ) # todo: should we fix this when enabling mtp or it doesn't matter since we only enable mtp in decode node thus we don't transfer draft kvs between P and D? @@ -1260,10 +1271,10 @@ class Scheduler( tp_size=self.ps.tp_size, dp_size=self.server_args.dp_size, gpu_id=self.ps.gpu_id, - bootstrap_port=self.server_args.disaggregation_bootstrap_port, + bootstrap_port=get_disagg().disaggregation_bootstrap_port, max_total_num_tokens=self.max_total_num_tokens, pp_rank=self.ps.pp_rank, - num_reserved_decode_tokens=self.server_args.num_reserved_decode_tokens, + num_reserved_decode_tokens=get_disagg().num_reserved_decode_tokens, transfer_backend=self.transfer_backend, ) @@ -1289,7 +1300,7 @@ class Scheduler( tp_rank=self.ps.tp_rank, tp_size=self.ps.tp_size, gpu_id=self.ps.gpu_id, - bootstrap_port=self.server_args.disaggregation_bootstrap_port, + bootstrap_port=get_disagg().disaggregation_bootstrap_port, gloo_group=self.attn_tp_cpu_group, max_total_num_tokens=self.max_total_num_tokens, scheduler=self, @@ -1303,11 +1314,10 @@ class Scheduler( self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get() # Init mm receiver for EPD disaggregation mode - if ( - self.server_args.language_only - and self.server_args.encoder_transfer_backend - in ["zmq_to_scheduler", "mooncake"] - ): + if get_disagg().language_only and get_disagg().encoder_transfer_backend in [ + "zmq_to_scheduler", + "mooncake", + ]: self.mm_receiver = create_mm_receiver( self.server_args, dtype=self.model_config.dtype, @@ -1388,7 +1398,7 @@ class Scheduler( def init_deterministic_inference_config(self): """Initialize deterministic inference configuration for different attention backends.""" - if not self.server_args.enable_deterministic_inference: + if not get_exec().deterministic.enable_deterministic_inference: self.truncation_align_size = None return @@ -1794,10 +1804,10 @@ class Scheduler( ) def init_lora_drainer(self) -> None: - if self.server_args.lora_drain_wait_threshold > 0.0: + if get_lora().lora_drain_wait_threshold > 0.0: self.lora_drainer = LoRADrainer( - self.server_args.max_loras_per_batch, - self.server_args.lora_drain_wait_threshold, + get_lora().max_loras_per_batch, + get_lora().lora_drain_wait_threshold, ) else: self.lora_drainer = None @@ -1923,7 +1933,7 @@ class Scheduler( def init_kv_events_publisher(self) -> None: self.kv_events_publisher = SchedulerKvEventsPublisher( - kv_events_config=self.server_args.kv_events_config, + kv_events_config=get_observability().kv_events_config, ps=self.ps, attn_tp_rank=self.ps.attn_tp_rank, attn_cp_rank=self.ps.attn_cp_rank, @@ -2107,7 +2117,7 @@ class Scheduler( return image_inputs def _get_multimodal_inputs(self, mm_inputs_dict): - if self.server_args.enable_broadcast_mm_inputs_process: + if get_mm().enable_broadcast_mm_inputs_process: return self._process_and_broadcast_mm_inputs(mm_inputs_dict) else: return MultimodalInputs.from_processor_output(mm_inputs_dict) @@ -2154,7 +2164,7 @@ class Scheduler( def _maybe_namespace_elastic_radix_cache(self, req: Req) -> None: if ( - self.server_args.elastic_ep_backend is None + get_exec().moe.elastic_ep_backend is None or self.disable_radix_cache or not self.tree_cache.is_tree_cache() ): @@ -2200,8 +2210,7 @@ class Scheduler( ) # Radix-native sessions use only the top-level session_id. radix_native_session = ( - recv_req.session_id is not None - and self.server_args.enable_session_radix_cache + recv_req.session_id is not None and get_memory().enable_session_radix_cache ) if session_id is None or radix_native_session: @@ -2213,7 +2222,7 @@ class Scheduler( if recv_req.bootstrap_port is None: # Use default bootstrap port - recv_req.bootstrap_port = self.server_args.disaggregation_bootstrap_port + recv_req.bootstrap_port = get_disagg().disaggregation_bootstrap_port req = Req( recv_req.rid, @@ -2366,7 +2375,7 @@ class Scheduler( self._add_request_to_queue(req) return - if req.return_sampling_mask and self.server_args.sampling_backend == "ascend": + if req.return_sampling_mask and get_exec().kernel.sampling_backend == "ascend": # The ascend backend samples from logits directly and never builds the # top-k/top-p support, so it cannot produce a sampling mask. error_msg = ( @@ -2415,7 +2424,7 @@ class Scheduler( error_msg = validate_input_length( req, self.max_req_input_len, - self.server_args.allow_auto_truncate, + get_serving().allow_auto_truncate, ) if error_msg: req.set_finish_with_abort(error_msg) @@ -2693,7 +2702,7 @@ class Scheduler( error_msg = validate_input_length( req, self.max_req_input_len, - self.server_args.allow_auto_truncate, + get_serving().allow_auto_truncate, ) if error_msg: self._add_request_to_queue(req) @@ -2905,7 +2914,7 @@ class Scheduler( if ( need_mlp_sync and not self.spec_algorithm.is_none() - and not self.server_args.speculative_skip_dp_mlp_sync + and not get_spec().speculative_skip_dp_mlp_sync ): # NOTE: This branch makes sure prefill and decode batches will not be mixed when spec and dp-attn is enabled. # Before merging the new batch into running batch: @@ -2979,7 +2988,7 @@ class Scheduler( for req in ready_grammar_requests: self._add_request_to_queue(req) - if self.enable_hierarchical_cache or self.server_args.enable_flexkv: + if self.enable_hierarchical_cache or get_memory().enable_flexkv: self.tree_cache.check_hicache_events() if self.enable_priority_preemption or self.is_hybrid_swa: @@ -3046,7 +3055,7 @@ class Scheduler( self.priority_scheduling_preemption_threshold, max_prefill_bs=self.max_prefill_bs, max_running_requests=self.max_running_requests, - prefill_max_requests=self.server_args.prefill_max_requests, + prefill_max_requests=get_schedule().prefill_max_requests, prefill_delayer_single_pass=prefill_delayer_single_pass, dllm_config=self.dllm_config, waiting_queue_len=len(self.waiting_queue), @@ -3619,7 +3628,7 @@ class Scheduler( def _maybe_report_active_ranks(self) -> None: if not ( - self.enable_dp_attention and self.server_args.elastic_ep_backend is not None + self.enable_dp_attention and get_exec().moe.elastic_ep_backend is not None ): return from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager @@ -3924,7 +3933,7 @@ class Scheduler( ok, msg = self.tree_cache.attach_storage_backend( storage_backend=recv_req.hicache_storage_backend, storage_backend_extra_config_json=recv_req.hicache_storage_backend_extra_config_json, - served_model_name=self.server_args.served_model_name, + served_model_name=get_serving().served_model_name, hicache_storage_prefetch_policy=recv_req.hicache_storage_prefetch_policy, hicache_write_policy=recv_req.hicache_write_policy, ) @@ -4044,7 +4053,7 @@ class Scheduler( } ret["effective_max_running_requests_per_dp"] = self.max_running_requests - if self.server_args.elastic_ep_backend is not None: + if get_exec().moe.elastic_ep_backend is not None: from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager ret["is_scaling_elastic_ep"] = ElasticEPStateManager.is_scaling() @@ -4583,10 +4592,10 @@ class Scheduler( return None def close_session(self, recv_req: CloseSessionReqInput): - if self.server_args.enable_session_radix_cache: + if get_memory().enable_session_radix_cache: self.tree_cache.release_radix_session(recv_req.session_id) if recv_req.session_id in self.session_controller or not ( - self.server_args.enable_session_radix_cache + get_memory().enable_session_radix_cache ): self.session_controller.close(recv_req) 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 248a92939..5729f1abd 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -27,7 +27,13 @@ from sglang.srt.mem_cache.common import ( maybe_cache_unfinished_req, release_kv_cache, ) -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import ( + get_disagg, + get_exec, + get_memory, + get_observability, + get_server_args, +) from sglang.srt.speculative.base_spec_worker import BaseSpecWorker from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer @@ -84,7 +90,7 @@ class SchedulerBatchResultProcessor: def process_batch_result_prebuilt(self, batch: ScheduleBatch): assert self.disaggregation_mode == DisaggregationMode.DECODE - use_free_group = self.server_args.disaggregation_decode_enable_radix_cache + use_free_group = get_disagg().disaggregation_decode_enable_radix_cache if use_free_group: self.token_to_kv_pool_allocator.free_group_begin() for req in batch.reqs: @@ -92,7 +98,7 @@ class SchedulerBatchResultProcessor: req.update_finish_state() if req.finished(): req.time_stats.set_quick_finish_time() - if self.server_args.enable_hisparse: + if get_memory().enable_hisparse: self.hisparse_coordinator.request_finished(req) release_kv_cache(req, self.tree_cache) @@ -243,7 +249,7 @@ class SchedulerBatchResultProcessor: req.time_stats.set_completion_time() elif not batch.decoding_reqs or req not in batch.decoding_reqs: maybe_cache_unfinished_req(req, self.tree_cache) - if self.server_args.enable_hisparse: + if get_memory().enable_hisparse: self.hisparse_coordinator.admit_request_into_staging(req) self._maybe_collect_customized_info(i, req, logits_output) @@ -756,7 +762,7 @@ class SchedulerBatchResultProcessor: num_block_accept_tokens=result.num_block_accept_tokens, num_cap_tokens=result.num_cap_tokens, ) - if self.server_args.enable_metrics: + if get_observability().enable_metrics: self.metrics_collector.increment_decode_cuda_graph_pass( value=can_run_cuda_graph ) @@ -939,7 +945,7 @@ class SchedulerBatchResultProcessor: self._mamba_prefix_cache_update(req, batch, result, i) if ( - self.server_args.disaggregation_decode_enable_offload_kvcache + get_disagg().disaggregation_decode_enable_offload_kvcache and not req.finished() ): self.decode_offload_manager.offload_kv_cache(req) @@ -959,12 +965,12 @@ class SchedulerBatchResultProcessor: self._maybe_collect_routed_experts(req) self._maybe_collect_indexer_topk(req) - if self.server_args.disaggregation_decode_enable_offload_kvcache: + if get_disagg().disaggregation_decode_enable_offload_kvcache: # Asynchronously offload KV cache; release_kv_cache will be called after Device->Host transfer completes if not self.decode_offload_manager.offload_kv_cache(req): self.decode_offload_manager.finalize_release_on_finish(req) else: - if self.server_args.enable_hisparse: + if get_memory().enable_hisparse: self.hisparse_coordinator.request_finished(req) prepare_release = getattr( self.model_worker, "prepare_for_kv_cache_release", None @@ -1063,7 +1069,7 @@ class SchedulerBatchResultProcessor: other_idx ].item() == -1 and mamba_lazy_spec_in_window( req, - server_args.mamba_track_interval, + get_exec().mamba.mamba_track_interval, server_args.max_speculative_num_draft_tokens, ) if ( @@ -1102,7 +1108,7 @@ class SchedulerBatchResultProcessor: For spec decode, the boundary is detected by comparing the accepted seq_len range against interval boundaries. """ - interval = get_server_args().mamba_track_interval + interval = get_exec().mamba.mamba_track_interval if batch.spec_algorithm.is_none(): if req.kv_committed_len % interval == 0: diff --git a/python/sglang/srt/managers/scheduler_components/dp_attn.py b/python/sglang/srt/managers/scheduler_components/dp_attn.py index 01a1d2adb..4ea7f2355 100644 --- a/python/sglang/srt/managers/scheduler_components/dp_attn.py +++ b/python/sglang/srt/managers/scheduler_components/dp_attn.py @@ -26,6 +26,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( ) from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.observability.metrics_collector import DPCooperationInfo +from sglang.srt.runtime_context import get_schedule from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils.common import require_mlp_tp_gather @@ -385,7 +386,7 @@ class SchedulerDPAttnAdapter: 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), - disable_overlap_schedule=self.server_args.disable_overlap_schedule, + disable_overlap_schedule=get_schedule().disable_overlap_schedule, offload_tags=self.offload_tags, dwdp=self.server_args.dwdp_size > 1, ) diff --git a/python/sglang/srt/managers/scheduler_components/load_inquirer.py b/python/sglang/srt/managers/scheduler_components/load_inquirer.py index 3edcf86a1..f26a98e7e 100644 --- a/python/sglang/srt/managers/scheduler_components/load_inquirer.py +++ b/python/sglang/srt/managers/scheduler_components/load_inquirer.py @@ -14,6 +14,7 @@ from sglang.srt.managers.load_snapshot import ( QueueMetrics, SpeculativeMetrics, ) +from sglang.srt.runtime_context import get_lora if TYPE_CHECKING: from sglang.srt.distributed.parallel_state_wrapper import ParallelState @@ -155,7 +156,7 @@ class SchedulerLoadInquirer: ) lora = None - if self.server_args.enable_lora: + if get_lora().enable_lora: lora = LoRAMetrics( slots_used=stats.lora_pool_slots_used, slots_total=stats.lora_pool_slots_total, diff --git a/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py b/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py index d97d9ae80..8c72deffd 100644 --- a/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py @@ -11,6 +11,7 @@ import torch from sglang.srt.configs.model_config import ModelConfig from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.managers.schedule_batch import Req +from sglang.srt.runtime_context import get_exec from sglang.srt.server_args import ( MIS_DELIMITER_TOKEN_ID, ServerArgs, @@ -164,7 +165,7 @@ class SchedulerLogprobResultProcessor: delimiter token receive logprobs. """ return ( - self.server_args.enable_mis + get_exec().features.enable_mis and req.is_prefill_only and req.multi_item_delimiter_indices is not None ) diff --git a/python/sglang/srt/managers/scheduler_components/metrics_reporter.py b/python/sglang/srt/managers/scheduler_components/metrics_reporter.py index 15e1b814d..0e5973627 100644 --- a/python/sglang/srt/managers/scheduler_components/metrics_reporter.py +++ b/python/sglang/srt/managers/scheduler_components/metrics_reporter.py @@ -26,6 +26,7 @@ from sglang.srt.observability.metrics_collector import ( SchedulerStats, compute_routing_key_stats, ) +from sglang.srt.runtime_context import get_spec from sglang.srt.utils.device_timer import DeviceTimer from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger @@ -764,12 +765,10 @@ class SchedulerMetricsReporter: else: spec_accept_length = self.spec_num_accept_tokens / self.spec_num_forward_ct num_correct_drafts = self.spec_num_accept_tokens - self.spec_num_forward_ct - if self.scheduler.server_args.speculative_num_draft_tokens: - draft_per_round = ( - self.scheduler.server_args.speculative_num_draft_tokens - 1 - ) + if get_spec().speculative_num_draft_tokens: + draft_per_round = get_spec().speculative_num_draft_tokens - 1 else: - draft_per_round = self.scheduler.server_args.speculative_num_steps or 0 + draft_per_round = get_spec().speculative_num_steps or 0 total_draft_tokens = self.spec_num_forward_ct * draft_per_round spec_accept_rate = ( num_correct_drafts / total_draft_tokens if total_draft_tokens > 0 else 0 diff --git a/python/sglang/srt/managers/scheduler_components/output_streamer.py b/python/sglang/srt/managers/scheduler_components/output_streamer.py index b06c66e33..1a61e7555 100644 --- a/python/sglang/srt/managers/scheduler_components/output_streamer.py +++ b/python/sglang/srt/managers/scheduler_components/output_streamer.py @@ -27,6 +27,7 @@ from sglang.srt.managers.schedule_batch import ( Req, ) from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache +from sglang.srt.runtime_context import get_observability, get_serving from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.spec_info import SpeculativeAlgorithm @@ -153,7 +154,7 @@ class SchedulerOutputStreamer: return_sampling_mask=return_sampling_mask, spec_algorithm=self.spec_algorithm, disaggregation_mode=self.disaggregation_mode, - default_stream_interval=self.server_args.stream_interval, + default_stream_interval=get_serving().stream_interval, default_force_stream_interval=DEFAULT_FORCE_STREAM_INTERVAL, get_cached_tokens_details=self.get_cached_tokens_details, rust_server_mode=self.rust_server is not None, @@ -184,7 +185,7 @@ class SchedulerOutputStreamer: if ( req.finished() and self.ps.attn_tp_rank == 0 - and self.server_args.enable_request_time_stats_logging + and get_observability().enable_request_time_stats_logging ): req.log_time_stats() diff --git a/python/sglang/srt/managers/scheduler_components/profiler_manager.py b/python/sglang/srt/managers/scheduler_components/profiler_manager.py index f60d53675..39eeec2e7 100644 --- a/python/sglang/srt/managers/scheduler_components/profiler_manager.py +++ b/python/sglang/srt/managers/scheduler_components/profiler_manager.py @@ -19,7 +19,7 @@ from sglang.srt.environ import envs from sglang.srt.managers.io_struct import ProfileReq, ProfileReqOutput, ProfileReqType from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.platforms import current_platform -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_device from sglang.srt.utils import is_mps, is_npu from sglang.srt.utils.profile_merger import ProfileMerger from sglang.srt.utils.profile_utils import ProfileManager @@ -257,7 +257,7 @@ class SchedulerProfilerManager: self.profile_in_progress = True if "CUDA_PROFILER" in activities: - if self.ps.gpu_id == get_server_args().base_gpu_id: + if self.ps.gpu_id == get_device().base_gpu_id: torch.cuda.cudart().cudaProfilerStart() self.profile_in_progress = True @@ -368,7 +368,7 @@ class SchedulerProfilerManager: torch.cuda.memory._record_memory_history(enabled=None) if "CUDA_PROFILER" in self.profiler_activities: - if self.ps.gpu_id == get_server_args().base_gpu_id: + if self.ps.gpu_id == get_device().base_gpu_id: torch.cuda.cudart().cudaProfilerStop() merge_message = self._merge_profile_traces() diff --git a/python/sglang/srt/managers/scheduler_components/request_receiver.py b/python/sglang/srt/managers/scheduler_components/request_receiver.py index d5a93fc1f..bbda80ccd 100644 --- a/python/sglang/srt/managers/scheduler_components/request_receiver.py +++ b/python/sglang/srt/managers/scheduler_components/request_receiver.py @@ -27,6 +27,7 @@ from sglang.srt.managers.mm_utils import ( has_shm_features, unwrap_shm_features, ) +from sglang.srt.runtime_context import get_disagg from sglang.srt.utils import ( broadcast_pyobj, point_to_point_pyobj, @@ -231,8 +232,8 @@ class SchedulerRequestReceiver: # Process MM requests under EPD-disaggregation mode if ( self.ps.pp_rank == 0 - and self.server_args.language_only - and self.server_args.encoder_transfer_backend + and get_disagg().language_only + and get_disagg().encoder_transfer_backend in ["zmq_to_scheduler", "mooncake"] ): recv_reqs, abort_reqs = self.mm_receiver.process_waiting_requests(recv_reqs) diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index 4e475f8a7..8bb69f489 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -36,6 +36,7 @@ from sglang.srt.model_executor.forward_batch_info import ( PPProxyTensors, ) from sglang.srt.observability.req_time_stats import set_time_batch +from sglang.srt.runtime_context import get_disagg from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.utils import DynamicGradMode, broadcast_pyobj, point_to_point_pyobj from sglang.srt.utils.common import get_device_module, is_xpu @@ -479,7 +480,7 @@ class SchedulerPPMixin: ) ) - if self.server_args.disaggregation_decode_enable_offload_kvcache: + if get_disagg().disaggregation_decode_enable_offload_kvcache: self.decode_offload_manager.check_offload_progress() if rmbs[next_mb_id] is not None: @@ -549,7 +550,7 @@ class SchedulerPPMixin: + len(self.disagg_decode_transfer_queue.queue) + len(self.disagg_decode_prealloc_queue.queue) ) - if self.server_args.disaggregation_decode_enable_offload_kvcache: + if get_disagg().disaggregation_decode_enable_offload_kvcache: queue_size += len(self.decode_offload_manager.ongoing_offload) if server_is_idle and queue_size == 0: diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 19ff16881..07f4c8df8 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -47,6 +47,7 @@ from sglang.srt.model_executor.forward_batch_info import ( PPProxyTensors, ) from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig +from sglang.srt.runtime_context import get_exec, get_model, get_schedule, get_spec from sglang.srt.server_args import ServerArgs from sglang.srt.utils import MultiprocessingSerializer, broadcast_pyobj, set_random_seed from sglang.srt.utils.hf_transformers_utils import ( @@ -408,14 +409,14 @@ class TpModelWorker(BaseTpWorker): self.model_config = ModelConfig.from_server_args( self.server_args, model_path=( - self.server_args.model_path + get_model().model_path if not self.is_draft_worker - else self.server_args.speculative_draft_model_path + else get_spec().speculative_draft_model_path ), model_revision=( - self.server_args.revision + get_model().revision if not self.is_draft_worker - else self.server_args.speculative_draft_model_revision + else get_spec().speculative_draft_model_revision ), is_draft_model=self.is_draft_worker, context_length=self.context_length, @@ -426,7 +427,7 @@ class TpModelWorker(BaseTpWorker): self._model_runner = ModelRunner( model_config=self.model_config, - mem_fraction_static=self.server_args.mem_fraction_static, + mem_fraction_static=get_schedule().mem_fraction_static, gpu_id=self.gpu_id, ps=self.ps, nccl_port=self.nccl_port, @@ -442,11 +443,11 @@ class TpModelWorker(BaseTpWorker): from sglang.srt.model_executor.model_runner import ModelRunner self.model_runner_list.append(self.model_runner) - for i in range(1, self.server_args.speculative_num_steps): + for i in range(1, get_spec().speculative_num_steps): self.model_runner_list.append( ModelRunner( model_config=self.model_config, - mem_fraction_static=self.server_args.mem_fraction_static, + mem_fraction_static=get_schedule().mem_fraction_static, gpu_id=self.gpu_id, ps=self.ps, nccl_port=self.nccl_port, @@ -462,7 +463,7 @@ class TpModelWorker(BaseTpWorker): def _init_dllm_algorithm(self): from sglang.srt.dllm.algorithm.base import DllmAlgorithm - if self.server_args.dllm_algorithm is not None: + if get_exec().dllm.dllm_algorithm is not None: self.dllm_algorithm = DllmAlgorithm.from_server_args(self.server_args) else: self.dllm_algorithm = None @@ -488,9 +489,9 @@ class TpModelWorker(BaseTpWorker): ) return ( self.model_runner.max_total_num_tokens, - self.server_args.max_prefill_tokens, + get_schedule().max_prefill_tokens, self.model_runner.max_running_requests, - self.server_args.max_queued_requests, + get_schedule().max_queued_requests, max_req_len, max_req_len - 5, self.random_seed, diff --git a/python/sglang/srt/mem_cache/base_prefix_cache.py b/python/sglang/srt/mem_cache/base_prefix_cache.py index fe208c986..1f40409a4 100644 --- a/python/sglang/srt/mem_cache/base_prefix_cache.py +++ b/python/sglang/srt/mem_cache/base_prefix_cache.py @@ -23,6 +23,7 @@ from sglang.srt.observability.metrics_collector import ( RadixCacheMetricsCollector, resolve_collector_class, ) +from sglang.srt.runtime_context import get_observability if TYPE_CHECKING: from sglang.srt.managers.schedule_batch import Req @@ -238,8 +239,8 @@ class BasePrefixCache(ABC, PrefixCacheTrait): server_args = get_server_args() labels = {"cache_type": self.__class__.__name__} - if server_args.extra_metric_labels: - labels.update(server_args.extra_metric_labels) + if get_observability().extra_metric_labels: + labels.update(get_observability().extra_metric_labels) radix_cache_cls = resolve_collector_class( server_args, STAT_LOGGER_ROLE_RADIX_CACHE, diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index d50c77677..11b2c33c8 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -16,7 +16,12 @@ from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import ( from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import ( + get_schedule, + get_server_args, + get_serving, + get_spec, +) from sglang.srt.utils.common import ceil_align if TYPE_CHECKING: @@ -179,12 +184,12 @@ def _release_overallocated_kv_indices( req: Req, start_p: int, end_p: int, tree_cache: BasePrefixCache ) -> None: global_server_args = get_server_args() - page_size = global_server_args.page_size - spec_algo = global_server_args.speculative_algorithm + page_size = get_schedule().page_size + spec_algo = get_spec().speculative_algorithm # strip_thinking_cache intentionally reports output tokens as overallocated # so they fall into the free path below (#22373). - if spec_algo is None and not global_server_args.strip_thinking_cache: + if spec_algo is None and not get_serving().strip_thinking_cache: assert ( start_p == end_p ), f"Unexpected overallocated KV cache, {req.kv_committed_len=}, {req.kv.kv_allocated_len=}" diff --git a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py index d1efe6db8..062e60bc9 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py @@ -21,7 +21,7 @@ from sglang.srt.environ import envs from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool from sglang.srt.mem_cache.memory_pool import KVCache -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec, get_server_args, get_spec from sglang.srt.utils import ceil_div, is_hip logger = logging.getLogger(__name__) @@ -276,7 +276,7 @@ class DeepSeekV4IndexerPool(KVCache): end_layer, ) self.index_head_dim = index_head_dim - self.use_fp4_indexer = get_server_args().enable_deepseek_v4_fp4_indexer + self.use_fp4_indexer = get_exec().kernel.enable_deepseek_v4_fp4_indexer self._create_buffer() @@ -577,8 +577,8 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): self.c128_kv_pool = None server_args = get_server_args() spec_extra = ( - (server_args.speculative_num_draft_tokens - 1) - if server_args.speculative_algorithm is not None + (get_spec().speculative_num_draft_tokens - 1) + if get_spec().speculative_algorithm is not None else 0 ) self.unified_kv_pool = DeepSeekV4UnifiedKVPool( @@ -659,7 +659,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): def get_ring_size(self, compress_ratio: int) -> int: server_args = get_server_args() - is_speculative = server_args.speculative_algorithm is not None + is_speculative = get_spec().speculative_algorithm is not None return get_compress_state_ring_size(compress_ratio, is_speculative) def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor): diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index 5b9e38d02..16d624586 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -58,7 +58,15 @@ from sglang.srt.mem_cache.memory_pool import ( ) from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.platforms import current_platform -from sglang.srt.runtime_context import get_model, get_parallel +from sglang.srt.runtime_context import ( + get_context, + get_disagg, + get_exec, + get_memory, + get_parallel, + get_schedule, + get_spec, +) from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils.common import ( @@ -184,6 +192,7 @@ class KVCacheConfigurator: token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] memory_pool_config: Optional[MemoryPoolConfig] draft_model_idx: Optional[int] = None + kv_cache_dtype_str: Optional[str] = None mambaish_config: Optional[Any] = field(init=False) hybrid_gdn_config: Optional[Any] = field(init=False) is_inkling_mtp_draft: bool = field(init=False) @@ -211,7 +220,7 @@ class KVCacheConfigurator: def _build_fp4_quant_method(self, *, num_layers: int): if not is_float4_e2m1fn_x2(self.kv_cache_dtype): return None - quant_name = resolve_kv_cache_quant(get_model().kv_cache_dtype) + quant_name = resolve_kv_cache_quant(self.kv_cache_dtype_str) if quant_name is None: return None quant_method = get_kv_cache_quant_method( @@ -314,8 +323,8 @@ class KVCacheConfigurator: # from one byte buffer, then return. Gated to the target worker # (req_to_token_pool is None); supports hybrid Mamba and hybrid SWA (not DSV4). if ( - self.server_args.enable_unified_memory - and self.server_args.disaggregation_mode == "null" + get_memory().enable_unified_memory + and get_disagg().disaggregation_mode == "null" and req_to_token_pool is None ): if self.mambaish_config is not None: @@ -364,13 +373,13 @@ class KVCacheConfigurator: # TARGET_VERIFY, so their pools skip the per-step intermediate # (SpeculativeState) buffers only the target pool consumes. req_to_token_pool = req_to_token_pool.clone_with_new_mamba( - mamba_size=self.server_args.max_mamba_cache_size, + mamba_size=get_schedule().max_mamba_cache_size, mamba_spec_state_size=sizes.max_running_requests, cache_params=self.mambaish_config.mamba2_cache_params, device=self.device, enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(), draft_model_idx=self.draft_model_idx, - speculative_eagle_topk=self.server_args.speculative_eagle_topk, + speculative_eagle_topk=get_spec().speculative_eagle_topk, ) # Initialize token_to_kv_pool @@ -400,7 +409,7 @@ class KVCacheConfigurator: # unsupported pool families before allocation. Keep this guard here so # future pool-selection refactors fail at boot instead of on first use. if ( - self.server_args.prefill_only_disable_kv_cache + get_schedule().prefill_only_disable_kv_cache and not self.is_draft_worker and not isinstance(token_to_kv_pool, NoOpMHATokenToKVPool) ): @@ -435,8 +444,8 @@ class KVCacheConfigurator: assert self.page_size >= 1, f"page_size must be >= 1, got {self.page_size}" # Mirror the non-shared path's extra_max_context_len computation. extra_max_context_len = 4 - if self.server_args.speculative_num_draft_tokens is not None: - extra_max_context_len += self.server_args.speculative_num_draft_tokens + if get_spec().speculative_num_draft_tokens is not None: + extra_max_context_len += get_spec().speculative_num_draft_tokens mamba_layer_ids = [ i @@ -471,14 +480,14 @@ class KVCacheConfigurator: model_context_len=self.model_config.context_len, extra_max_context_len=extra_max_context_len, max_total_num_tokens=max_total_num_tokens, - max_mamba_cache_size=self.server_args.max_mamba_cache_size, + max_mamba_cache_size=get_schedule().max_mamba_cache_size, max_num_reqs=max_num_reqs, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(), - speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens, - disable_overlap_schedule=self.server_args.disable_overlap_schedule, - need_sort=self.server_args.disaggregation_mode in ("decode", "prefill"), - mamba_full_memory_ratio=self.server_args.mamba_full_memory_ratio, + speculative_num_draft_tokens=get_spec().speculative_num_draft_tokens, + disable_overlap_schedule=get_schedule().disable_overlap_schedule, + need_sort=get_disagg().disaggregation_mode in ("decode", "prefill"), + mamba_full_memory_ratio=get_schedule().mamba_full_memory_ratio, # Overlap mode: the allocator's `free` drops a wait_stream(forward_stream) # barrier so eager compaction serializes after the in-flight forward's # v2p/KV reads. Near-no-op in normal mode. @@ -511,13 +520,13 @@ class KVCacheConfigurator: ), "unified memory pool does not support MLA-SWA hybrid yet" # Mirror the non-shared path's extra_max_context_len computation. extra_max_context_len = 4 - if self.server_args.speculative_num_draft_tokens is not None: - extra_max_context_len += self.server_args.speculative_num_draft_tokens + if get_spec().speculative_num_draft_tokens is not None: + extra_max_context_len += get_spec().speculative_num_draft_tokens req_to_token_pool = ReqToTokenPool( size=max_num_reqs, max_context_len=self.model_config.context_len + extra_max_context_len, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, ) head_num = self.model_config.get_num_kv_heads(get_parallel().attn_tp_size) @@ -567,8 +576,8 @@ class KVCacheConfigurator: full_attention_layer_ids=full_attention_layer_ids, full_max_total_num_tokens=full_max_total_num_tokens, swa_max_total_num_tokens=swa_max_total_num_tokens, - enable_memory_saver=self.server_args.enable_memory_saver, - need_sort=self.server_args.disaggregation_mode in ("decode", "prefill"), + enable_memory_saver=get_exec().features.enable_memory_saver, + need_sort=get_disagg().disaggregation_mode in ("decode", "prefill"), # Overlap mode: same wait_stream(forward_stream) rationale as # `_init_unified_mamba_pools`. forward_stream=self.forward_stream, @@ -588,7 +597,7 @@ class KVCacheConfigurator: is_dsv4_model: bool, current_platform, ): - if not self.server_args.prefill_only_disable_kv_cache or self.is_draft_worker: + if not get_schedule().prefill_only_disable_kv_cache or self.is_draft_worker: return unsupported_pool_family = None @@ -623,9 +632,9 @@ class KVCacheConfigurator: def _build_req_to_token_pool(self, *, max_num_reqs: int) -> ReqToTokenPool: extra_max_context_len = get_req_to_token_extra_context_len(self.server_args) - if self.server_args.disaggregation_mode == "decode": + if get_disagg().disaggregation_mode == "decode": # Extra slots for pre-allocated requests - pre_alloc_size = self.server_args.disaggregation_decode_extra_slots + pre_alloc_size = get_disagg().disaggregation_decode_extra_slots if self.mambaish_config: req_to_token_pool = self._build_hybrid_mamba_decode_req_pool( max_num_reqs=max_num_reqs, @@ -665,7 +674,7 @@ class KVCacheConfigurator: size=max_num_reqs, max_context_len=self.model_config.context_len + extra_max_context_len, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, cache_params=self.mambaish_config.mamba2_cache_params, mamba_layer_ids=( [ @@ -675,16 +684,16 @@ class KVCacheConfigurator: ] ), speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens, - speculative_eagle_topk=self.server_args.speculative_eagle_topk, + speculative_eagle_topk=get_spec().speculative_eagle_topk, enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(), pre_alloc_size=pre_alloc_size, - enable_overlap_schedule=not self.server_args.disable_overlap_schedule, - mamba_size=self.server_args.max_mamba_cache_size, + enable_overlap_schedule=not get_schedule().disable_overlap_schedule, + mamba_size=get_schedule().max_mamba_cache_size, start_layer=self.layer_info.start_layer, - linear_replayssm_cache_len=self.server_args.linear_replayssm_cache_len, - mamba_envelope_layout=self.server_args.enable_page_major_kv_layout, + linear_replayssm_cache_len=get_exec().mamba.linear_replayssm_cache_len, + mamba_envelope_layout=get_memory().enable_page_major_kv_layout, enable_gdn_replayssm_spec=( - self.server_args.enable_gdn_replayssm_spec + get_exec().mamba.enable_gdn_replayssm_spec and self.hybrid_gdn_config is not None ), ) @@ -708,7 +717,7 @@ class KVCacheConfigurator: size=max_num_reqs, max_context_len=self.model_config.context_len + extra_max_context_len, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, pre_alloc_size=pre_alloc_size, ) return req_to_token_pool @@ -721,11 +730,11 @@ class KVCacheConfigurator: ) -> ReqToTokenPool: req_to_token_pool = HybridReqToTokenPool( size=max_num_reqs, - mamba_size=self.server_args.max_mamba_cache_size, + mamba_size=get_schedule().max_mamba_cache_size, mamba_spec_state_size=max_num_reqs, max_context_len=self.model_config.context_len + extra_max_context_len, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, cache_params=self.mambaish_config.mamba2_cache_params, mamba_layer_ids=( [ @@ -737,14 +746,14 @@ class KVCacheConfigurator: enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(), enable_mamba_extra_buffer_lazy=self.server_args.enable_mamba_extra_buffer_lazy(), speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens, - speculative_eagle_topk=self.server_args.speculative_eagle_topk, - enable_overlap_schedule=not self.server_args.disable_overlap_schedule, + speculative_eagle_topk=get_spec().speculative_eagle_topk, + enable_overlap_schedule=not get_schedule().disable_overlap_schedule, start_layer=self.layer_info.start_layer, - enable_linear_replayssm=self.server_args.enable_linear_replayssm, - linear_replayssm_cache_len=self.server_args.linear_replayssm_cache_len, - mamba_envelope_layout=self.server_args.enable_page_major_kv_layout, + enable_linear_replayssm=get_exec().mamba.enable_linear_replayssm, + linear_replayssm_cache_len=get_exec().mamba.linear_replayssm_cache_len, + mamba_envelope_layout=get_memory().enable_page_major_kv_layout, enable_gdn_replayssm_spec=( - self.server_args.enable_gdn_replayssm_spec + get_exec().mamba.enable_gdn_replayssm_spec and self.hybrid_gdn_config is not None ), ) @@ -770,7 +779,7 @@ class KVCacheConfigurator: size=max_num_reqs, max_context_len=self.model_config.context_len + extra_max_context_len, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, ) return req_to_token_pool @@ -786,7 +795,7 @@ class KVCacheConfigurator: # selected by swapping in the PageMajorMHATokenToKVPool subclass. The # default keeps upstream's per-layer layout. The Mamba state pool is routed # separately via `mamba_envelope_layout` on the req-to-token pool above. - enable_page_major = self.server_args.enable_page_major_kv_layout + enable_page_major = get_memory().enable_page_major_kv_layout mha_pool_class = ( PageMajorMHATokenToKVPool if enable_page_major else MHATokenToKVPool ) @@ -894,7 +903,7 @@ class KVCacheConfigurator: c128_state_dtype: Optional[torch.dtype], req_to_token_pool: ReqToTokenPool, ) -> KVCache: - swa_page_size = self.server_args.page_size + swa_page_size = get_schedule().page_size if not _is_npu: assert swa_page_size == 256, "In paged swa mode, page_size must be 256." @@ -928,12 +937,12 @@ class KVCacheConfigurator: # sliding eviction in ``ScheduleBatch._evict_swa``. c4_state_pool_size = npu_state_pool_size( ratio=4, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, max_num_reqs=max_running_requests, ) c128_state_pool_size = npu_state_pool_size( ratio=128, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, max_num_reqs=max_running_requests, ) else: @@ -951,7 +960,7 @@ class KVCacheConfigurator: c128_size=c128_max_total_num_tokens, c4_state_pool_size=c4_state_pool_size, c128_state_pool_size=c128_state_pool_size, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, swa_page_size=swa_page_size, sliding_window=self.model_config.window_size, dtype=self.kv_cache_dtype, @@ -962,11 +971,11 @@ class KVCacheConfigurator: indexer_head_dim=self.model_config.index_head_dim, layer_num=self.layer_info.num_effective_layers, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, compression_ratios=compression_ratios, start_layer=self.layer_info.start_layer, end_layer=self.layer_info.end_layer, - enable_hisparse=self.server_args.enable_hisparse, + enable_hisparse=get_memory().enable_hisparse, online_mtp_max_draft_tokens=( self.server_args.max_speculative_num_draft_tokens or 0 ), @@ -977,7 +986,7 @@ class KVCacheConfigurator: PoolCls = current_platform.get_dsa_kv_pool_cls() token_to_kv_pool = PoolCls( max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, @@ -988,7 +997,7 @@ class KVCacheConfigurator: kv_cache_dtype=self.kv_cache_dtype, server_args=self.server_args, ), - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, start_layer=self.layer_info.start_layer, end_layer=self.layer_info.end_layer, index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config), @@ -1001,14 +1010,14 @@ class KVCacheConfigurator: PoolCls = current_platform.get_mla_kv_pool_cls() token_to_kv_pool = PoolCls( max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, index_head_dim=(self.model_config.index_head_dim if is_dsa_model else None), layer_num=self.layer_info.num_effective_layers, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, start_layer=self.layer_info.start_layer, end_layer=self.layer_info.end_layer, ) @@ -1018,13 +1027,13 @@ class KVCacheConfigurator: PoolCls = current_platform.get_mha_kv_pool_cls() token_to_kv_pool = PoolCls( max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_dim=self.model_config.head_dim, layer_num=self.layer_info.num_effective_layers, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, start_layer=self.layer_info.start_layer, end_layer=self.layer_info.end_layer, ) @@ -1055,7 +1064,7 @@ class KVCacheConfigurator: token_to_kv_pool = SWAKVPool( size=full_max_total_num_tokens, size_swa=swa_max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, post_capture_active=self.post_capture_kv_active, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), @@ -1077,14 +1086,14 @@ class KVCacheConfigurator: token_to_kv_pool = NPUMLATokenToKVPool( max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, index_head_dim=(self.model_config.index_head_dim if is_dsa_model else None), layer_num=self.layer_info.num_effective_layers, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, start_layer=self.layer_info.start_layer, end_layer=self.layer_info.end_layer, ) @@ -1097,13 +1106,13 @@ class KVCacheConfigurator: token_to_kv_pool = NPUMHATokenToKVPool( max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_dim=self.model_config.head_dim, layer_num=self.layer_info.num_effective_layers, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, start_layer=self.layer_info.start_layer, end_layer=self.layer_info.end_layer, ) @@ -1117,7 +1126,7 @@ class KVCacheConfigurator: dsa_cp_layer_shard_size, ) = get_glm_dsa_cp_layer_shard_info(self) pool_kwargs = {} - if self.server_args.enable_hisparse: + if get_memory().enable_hisparse: PoolCls = HiSparseDSATokenToKVPool from sglang.srt.mem_cache.sparsity import parse_hisparse_config @@ -1137,7 +1146,7 @@ class KVCacheConfigurator: PoolCls = DSATokenToKVPool token_to_kv_pool = PoolCls( max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, @@ -1148,7 +1157,7 @@ class KVCacheConfigurator: kv_cache_dtype=self.kv_cache_dtype, server_args=self.server_args, ), - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, start_layer=self.layer_info.start_layer, end_layer=self.layer_info.end_layer, index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config), @@ -1159,13 +1168,13 @@ class KVCacheConfigurator: def _build_mla_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: token_to_kv_pool = MLATokenToKVPoolFP4( max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, layer_num=self.layer_info.num_effective_layers, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, start_layer=self.layer_info.start_layer, end_layer=self.layer_info.end_layer, ) @@ -1174,13 +1183,13 @@ class KVCacheConfigurator: def _build_mla_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: token_to_kv_pool = MLATokenToKVPool( max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, kv_lora_rank=self.model_config.kv_lora_rank, qk_rope_head_dim=self.model_config.qk_rope_head_dim, layer_num=self.layer_info.num_effective_layers, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, start_layer=self.layer_info.start_layer, end_layer=self.layer_info.end_layer, ) @@ -1207,7 +1216,7 @@ class KVCacheConfigurator: } swa_pool_class = ( MHATokenToKVPoolMXFP8 - if get_model().kv_cache_dtype == "mxfp8" + if self.kv_cache_dtype_str == "mxfp8" else mha_pool_class ) swa_attention_layer_ids = self.model_config.swa_attention_layer_ids @@ -1237,7 +1246,7 @@ class KVCacheConfigurator: token_to_kv_pool = SWAKVPool( size=full_max_total_num_tokens, size_swa=size_swa, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, post_capture_active=self.post_capture_kv_active, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), @@ -1245,7 +1254,7 @@ class KVCacheConfigurator: swa_attention_layer_ids=swa_attention_layer_ids, full_attention_layer_ids=full_attention_layer_ids, device=self.device, - enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None), + enable_kv_cache_copy=(get_spec().speculative_algorithm is not None), token_to_kv_pool_class=swa_pool_class, **kwargs, ) @@ -1260,7 +1269,7 @@ class KVCacheConfigurator: ) token_to_kv_pool = MiniMaxSparseKVPool( size=max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, index_dtype=self.model_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), @@ -1270,7 +1279,7 @@ class KVCacheConfigurator: sparse_layer_ids=sparse_layer_ids, disable_value_sparse_layer_ids=disable_value_sparse_layer_ids, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, start_layer=self.layer_info.start_layer, end_layer=self.layer_info.end_layer, ) @@ -1305,11 +1314,11 @@ class KVCacheConfigurator: # buffers) for the full-attention layers, same as the SWA branch. full_pool_class = ( MHATokenToKVPoolMXFP8 - if get_model().kv_cache_dtype == "mxfp8" and not self.use_mla_backend + if self.kv_cache_dtype_str == "mxfp8" and not self.use_mla_backend else mha_pool_class ) token_to_kv_pool = HybridLinearKVPool( - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, size=max_total_num_tokens, dtype=self.kv_cache_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), @@ -1318,8 +1327,8 @@ class KVCacheConfigurator: full_attention_layer_ids=full_attention_layer_ids, device=self.device, mamba_pool=req_to_token_pool.mamba_pool, - enable_memory_saver=self.server_args.enable_memory_saver, - enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None), + enable_memory_saver=get_exec().features.enable_memory_saver, + enable_kv_cache_copy=(get_spec().speculative_algorithm is not None), use_mla=self.use_mla_backend, start_layer=self.layer_info.start_layer, full_kv_pool_class=full_pool_class, @@ -1332,30 +1341,30 @@ class KVCacheConfigurator: def _build_mha_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: token_to_kv_pool = MHATokenToKVPoolFP4( max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_dim=self.model_config.head_dim, v_head_dim=self.model_config.v_head_dim, layer_num=self.layer_info.num_effective_layers, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, start_layer=self.layer_info.start_layer, end_layer=self.layer_info.end_layer, - enable_alt_stream=not self.server_args.enable_pdmux, - enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None), + enable_alt_stream=not get_disagg().enable_pdmux, + enable_kv_cache_copy=(get_spec().speculative_algorithm is not None), ) return token_to_kv_pool def _build_mha_kv_pool( self, *, max_total_num_tokens: int, mha_pool_class: type, quant_method=None ) -> KVCache: - if get_model().kv_cache_dtype == "mxfp8": + if self.kv_cache_dtype_str == "mxfp8": pool_cls = MHATokenToKVPoolMXFP8 else: pool_cls = ( NoOpMHATokenToKVPool - if self.server_args.prefill_only_disable_kv_cache + if get_schedule().prefill_only_disable_kv_cache else mha_pool_class ) pool_kwargs = {} @@ -1365,18 +1374,18 @@ class KVCacheConfigurator: pool_kwargs["post_capture_active"] = self.post_capture_kv_active token_to_kv_pool = pool_cls( max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size), head_dim=self.model_config.head_dim, v_head_dim=self.model_config.v_head_dim, layer_num=self.layer_info.num_effective_layers, device=self.device, - enable_memory_saver=self.server_args.enable_memory_saver, + enable_memory_saver=get_exec().features.enable_memory_saver, start_layer=self.layer_info.start_layer, end_layer=self.layer_info.end_layer, - enable_alt_stream=not self.server_args.enable_pdmux, - enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None), + enable_alt_stream=not get_disagg().enable_pdmux, + enable_kv_cache_copy=(get_spec().speculative_algorithm is not None), **pool_kwargs, ) return token_to_kv_pool @@ -1391,13 +1400,13 @@ class KVCacheConfigurator: token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator], ) -> BaseTokenToKVPoolAllocator: # Initialize token_to_kv_pool_allocator - need_sort = self.server_args.disaggregation_mode in ("decode", "prefill") + need_sort = get_disagg().disaggregation_mode in ("decode", "prefill") if token_to_kv_pool_allocator is None: if current_platform.is_out_of_tree(): AllocatorCls = current_platform.get_paged_allocator_cls() token_to_kv_pool_allocator = AllocatorCls( sizes.max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, device=self.device, kvcache=token_to_kv_pool, @@ -1422,7 +1431,7 @@ class KVCacheConfigurator: token_to_kv_pool_allocator = swa_allocator_cls( sizes.full_max_total_num_tokens, sizes.swa_max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, device=self.device, kvcache=token_to_kv_pool, @@ -1435,7 +1444,7 @@ class KVCacheConfigurator: token_to_kv_pool_allocator = NPUPagedTokenToKVPoolAllocator( sizes.max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, device=self.device, kvcache=token_to_kv_pool, @@ -1445,7 +1454,7 @@ class KVCacheConfigurator: if self.is_hybrid_swa and sizes.full_max_total_num_tokens == 0: token_to_kv_pool_allocator = PureSWATokenToKVPoolAllocator( sizes.swa_max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, device=self.device, kvcache=token_to_kv_pool, @@ -1455,14 +1464,14 @@ class KVCacheConfigurator: token_to_kv_pool_allocator = SWATokenToKVPoolAllocator( sizes.full_max_total_num_tokens, sizes.swa_max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, device=self.device, kvcache=token_to_kv_pool, need_sort=need_sort, ) else: - if self.server_args.enable_hisparse: + if get_memory().enable_hisparse: from sglang.srt.mem_cache.sparsity import ( parse_hisparse_config, ) @@ -1470,7 +1479,7 @@ class KVCacheConfigurator: hisparse_cfg = parse_hisparse_config(self.server_args) token_to_kv_pool_allocator = HiSparseTokenToKVPoolAllocator( sizes.max_total_num_tokens, - page_size=self.server_args.page_size, + page_size=get_schedule().page_size, dtype=self.kv_cache_dtype, device=self.device, kvcache=token_to_kv_pool, @@ -1478,8 +1487,7 @@ class KVCacheConfigurator: host_to_device_ratio=hisparse_cfg.host_to_device_ratio, ) elif ( - self.server_args.page_size == 1 - and self.server_args.dcp_size == 1 + get_schedule().page_size == 1 and self.server_args.dcp_size == 1 ): token_to_kv_pool_allocator = TokenToKVPoolAllocator( sizes.max_total_num_tokens, @@ -1491,7 +1499,7 @@ class KVCacheConfigurator: else: token_to_kv_pool_allocator = PagedTokenToKVPoolAllocator( sizes.max_total_num_tokens * self.server_args.dcp_size, - page_size=self.server_args.page_size + page_size=get_schedule().page_size * self.server_args.dcp_size, dtype=self.kv_cache_dtype, device=self.device, @@ -1499,7 +1507,7 @@ class KVCacheConfigurator: need_sort=need_sort, ) - if self.server_args.enable_hisparse and is_dsv4_model: + if get_memory().enable_hisparse and is_dsv4_model: assert self.is_hybrid_swa, "DeepSeek V4 HiSparse requires SWA mode." token_to_kv_pool_allocator = DeepSeekV4HiSparseTokenToKVPoolAllocator( token_to_kv_pool_allocator @@ -1551,7 +1559,7 @@ class KVCacheConfigurator: cpu_group=get_world_group().cpu_group, ) - slack_gb = pre_model_load_memory * (1 - self.server_args.mem_fraction_static) + slack_gb = pre_model_load_memory * (1 - get_schedule().mem_fraction_static) if self.mambaish_config is not None and self.post_capture_kv_active: # Mamba state is a fixed pre-capture allocation, so it can't ride the ~0 post-capture slack. slack_gb = max( @@ -1575,7 +1583,7 @@ class KVCacheConfigurator: ) raise ValueError( f"Loaded weights leave no GPU memory for the KV cache under " - f"--mem-fraction-static={self.server_args.mem_fraction_static}. " + f"--mem-fraction-static={get_schedule().mem_fraction_static}. " f"Raise --mem-fraction-static above " f"{suggested_mem_fraction_static:.3f} " f"(minimum viable = 1 - available/pre = " @@ -1586,7 +1594,7 @@ class KVCacheConfigurator: return int(rest_memory * (1 << 30)) # return in bytes def _calculate_mamba_ratio(self) -> int: - if self.server_args.disable_radix_cache: + if get_memory().disable_radix_cache: return 1 skip_decode_lock = envs.SGLANG_OPT_MAMBA_SKIP_DECODE_LOCK.get() @@ -1598,7 +1606,7 @@ class KVCacheConfigurator: if self.server_args.enable_mamba_extra_buffer(): # ping-pong buffer size is 2 when overlap schedule is on, 1 otherwise. # Lazy mode saves 1 slot (2 → 1) for overlap; non-overlap already uses 1. - if not self.server_args.disable_overlap_schedule: + if not get_schedule().disable_overlap_schedule: if self.server_args.enable_mamba_extra_buffer_lazy(): additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP_LAZY else: @@ -1622,7 +1630,7 @@ class KVCacheConfigurator: Page alignment is handled by the configurator, not here. If constraints change the value, the configurator re-runs and re-aligns. """ - user_limit = self.server_args.max_total_tokens + user_limit = get_schedule().max_total_tokens # Apply user-specified upper bound if user_limit is not None: @@ -1652,7 +1660,7 @@ class KVCacheConfigurator: estimated = int(token_capacity / self.model_config.context_len * 512) estimated = max(min(estimated, 4096), 2048) - max_num_reqs = self.server_args.max_running_requests + max_num_reqs = get_schedule().max_running_requests if max_num_reqs is not None: requested_per_worker = max_num_reqs // self.ps.attn_dp_size max_num_reqs = min(requested_per_worker, token_capacity // 2) @@ -1663,13 +1671,13 @@ class KVCacheConfigurator: if self.mambaish_config is not None: ratio = self._calculate_mamba_ratio() max_num_reqs = min( - max_num_reqs, self.server_args.max_mamba_cache_size // ratio + max_num_reqs, get_schedule().max_mamba_cache_size // ratio ) if max_num_reqs <= 0: raise RuntimeError( f"Hybrid (mamba/linear-attention) state cache is too small to serve " - f"any requests. max_mamba_cache_size={self.server_args.max_mamba_cache_size}, " + f"any requests. max_mamba_cache_size={get_schedule().max_mamba_cache_size}, " f"mamba_ratio={ratio}, resulting max_num_reqs={max_num_reqs}. " f"Try: (1) reduce --max-running-requests, " f"(2) increase --mem-fraction-static, or " @@ -1699,7 +1707,7 @@ class KVCacheConfigurator: ) configurator = create_memory_pool_configurator(self) config = configurator.finalize_with_max_running_requests(config) - config.mem_fraction_static = self.server_args.mem_fraction_static + config.mem_fraction_static = get_schedule().mem_fraction_static return config def config_from_budget( @@ -1715,14 +1723,14 @@ class KVCacheConfigurator: configurator = create_memory_pool_configurator(self) config = configurator.calculate_pool_sizes( - budget_bytes, self.server_args.page_size + budget_bytes, get_schedule().page_size ) max_tokens = self._apply_token_constraints(config.max_total_num_tokens) if cap_tokens is not None: max_tokens = min(max_tokens, cap_tokens) if max_tokens != config.max_total_num_tokens: config = configurator.calculate_pool_sizes_from_max_tokens( - max_tokens, self.server_args.page_size + max_tokens, get_schedule().page_size ) return config @@ -1735,13 +1743,14 @@ class KVCacheConfigurator: # The ring is allocated per slot but is not part of mamba_cache_per_req; # the solve must charge it too or num_slots is over-provisioned. replayssm_active = ( - server_args.enable_gdn_replayssm_spec and self.hybrid_gdn_config is not None + get_exec().mamba.enable_gdn_replayssm_spec + and self.hybrid_gdn_config is not None ) if replayssm_active: record_len = ( server_args.max_speculative_num_draft_tokens if server_args.max_speculative_num_draft_tokens is not None - else server_args.linear_replayssm_cache_len + else get_exec().mamba.linear_replayssm_cache_len ) replayssm_ring_per_req = ( config.mamba2_cache_params.replayssm_ring_bytes_per_req( @@ -1751,45 +1760,45 @@ class KVCacheConfigurator: else: replayssm_ring_per_req = 0 if has_spec_dec: - assert server_args.speculative_num_draft_tokens is not None - assert server_args.max_running_requests is not None + assert get_spec().speculative_num_draft_tokens is not None + assert get_schedule().max_running_requests is not None - if server_args.max_mamba_cache_size is not None: + if get_schedule().max_mamba_cache_size is not None: # Use explicitly set max_mamba_cache_size - server_args.override( + get_context().override( "mamba_pool.per_dp_shard", - max_mamba_cache_size=server_args.max_mamba_cache_size + max_mamba_cache_size=get_schedule().max_mamba_cache_size // self.ps.attn_dp_size, ) # Reserve intermediate memory based on capped max_num_reqs (+1 padding slot) if has_spec_dec and not replayssm_active: ratio = self._calculate_mamba_ratio() capped_reqs = min( - server_args.max_running_requests // self.ps.attn_dp_size, - server_args.max_mamba_cache_size // ratio, + get_schedule().max_running_requests // self.ps.attn_dp_size, + get_schedule().max_mamba_cache_size // ratio, ) intermediate_size = ( config.mamba2_cache_params.mamba_cache_per_req * (capped_reqs + 1) - * server_args.speculative_num_draft_tokens + * get_spec().speculative_num_draft_tokens ) total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) elif ( - server_args.disable_radix_cache - and server_args.max_running_requests is not None + get_memory().disable_radix_cache + and get_schedule().max_running_requests is not None ): # Use explicitly set max_running_requests when radix cache is disabled - server_args.override( + get_context().override( "mamba_pool.from_max_running_requests", - max_mamba_cache_size=server_args.max_running_requests + max_mamba_cache_size=get_schedule().max_running_requests // self.ps.attn_dp_size, ) # Reserve intermediate memory based on capped max_num_reqs (+1 padding slot) if has_spec_dec and not replayssm_active: intermediate_size = ( config.mamba2_cache_params.mamba_cache_per_req - * (server_args.max_mamba_cache_size + 1) - * server_args.speculative_num_draft_tokens + * (get_schedule().max_mamba_cache_size + 1) + * get_spec().speculative_num_draft_tokens ) total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) else: @@ -1802,16 +1811,16 @@ class KVCacheConfigurator: # (K + 1) * per_req + (K / ratio + 1) * D * per_req = mamba_budget_bytes mamba_budget = ( total_rest_memory - * server_args.mamba_full_memory_ratio - / (1 + server_args.mamba_full_memory_ratio) + * get_schedule().mamba_full_memory_ratio + / (1 + get_schedule().mamba_full_memory_ratio) ) mamba_budget_bytes = mamba_budget * (1 << 30) if has_spec_dec and not replayssm_active: ratio = self._calculate_mamba_ratio() - D = server_args.speculative_num_draft_tokens + D = get_spec().speculative_num_draft_tokens # Joint solve: main_state + intermediate = mamba_budget - server_args.override( + get_context().override( "mamba_pool.memory_budget_spec", max_mamba_cache_size=int( (mamba_budget_bytes - per_req * (1 + D)) @@ -1821,14 +1830,14 @@ class KVCacheConfigurator: # Intermediate memory is included in mamba_budget, subtract it # so the return value only has main_state subtracted from total capped_reqs = min( - server_args.max_running_requests // self.ps.attn_dp_size, - server_args.max_mamba_cache_size // ratio, + get_schedule().max_running_requests // self.ps.attn_dp_size, + get_schedule().max_mamba_cache_size // ratio, ) intermediate_size = per_req * (capped_reqs + 1) * D total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) else: per_slot = per_req + replayssm_ring_per_req - server_args.override( + get_context().override( "mamba_pool.memory_budget", max_mamba_cache_size=int( (mamba_budget_bytes - per_slot) // per_slot @@ -1839,10 +1848,10 @@ class KVCacheConfigurator: # A non-positive value means GPU memory is insufficient for the requested # configuration. Fail fast with actionable advice instead of silently # producing garbled output at runtime. - if server_args.max_mamba_cache_size <= 0: + if get_schedule().max_mamba_cache_size <= 0: raise RuntimeError( f"Not enough GPU memory for hybrid (mamba/linear-attention) state cache. " - f"Computed max_mamba_cache_size={server_args.max_mamba_cache_size} " + f"Computed max_mamba_cache_size={get_schedule().max_mamba_cache_size} " f"(total_rest_memory={total_rest_memory:.2f} GB, " f"mamba_cache_per_req={config.mamba2_cache_params.mamba_cache_per_req / (1 << 20):.2f} MB). " f"Try: (1) reduce --max-running-requests, " @@ -1853,7 +1862,7 @@ class KVCacheConfigurator: # +1: the pool's padding slot mamba_state_memory = ( - (server_args.max_mamba_cache_size + 1) + (get_schedule().max_mamba_cache_size + 1) * (config.mamba2_cache_params.mamba_cache_per_req + replayssm_ring_per_req) / (1 << 30) ) diff --git a/python/sglang/srt/mem_cache/storage/flexkv/flexkv_radix_cache.py b/python/sglang/srt/mem_cache/storage/flexkv/flexkv_radix_cache.py index 3c3f8e562..a77d82cfc 100644 --- a/python/sglang/srt/mem_cache/storage/flexkv/flexkv_radix_cache.py +++ b/python/sglang/srt/mem_cache/storage/flexkv/flexkv_radix_cache.py @@ -42,6 +42,7 @@ from sglang.srt.mem_cache.base_prefix_cache import ( ) from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode from sglang.srt.mem_cache.storage.flexkv.flexkv_connector import FlexKVConnector +from sglang.srt.runtime_context import get_spec if TYPE_CHECKING: from sglang.srt.configs.model_config import ModelConfig @@ -393,7 +394,7 @@ class FlexKVRadixCache(RadixCache): from sglang.srt.runtime_context import get_server_args global_server_args = get_server_args() - topk = global_server_args.speculative_eagle_topk + topk = get_spec().speculative_eagle_topk enable_kv_committed_len = topk is None or topk == 1 if enable_kv_committed_len: kv_committed_len = req.kv_committed_len diff --git a/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py b/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py index caf6907b1..010ebc65e 100644 --- a/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py +++ b/python/sglang/srt/mem_cache/storage/lmcache/lmc_radix_cache.py @@ -16,7 +16,7 @@ from sglang.srt.mem_cache.base_prefix_cache import ( MatchResult, ) from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_memory, get_server_args, get_spec from sglang.srt.utils import create_device_stream, device_stream_context try: @@ -109,7 +109,7 @@ class LMCRadixCache(RadixCache): ): super().__init__(params) - cli_lmc_cfg = get_server_args().lmcache_config_file or "" + cli_lmc_cfg = get_memory().lmcache_config_file or "" kvcache = self.token_to_kv_pool_allocator.get_kvcache() connector_kwargs = dict( @@ -448,7 +448,7 @@ class LMCRadixCache(RadixCache): return global_server_args = get_server_args() - topk = global_server_args.speculative_eagle_topk + topk = get_spec().speculative_eagle_topk enable_kv_committed_len = topk is None or topk == 1 if enable_kv_committed_len: kv_committed_len = req.kv_committed_len diff --git a/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py b/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py index 9dcb921ba..c42b33615 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py @@ -35,7 +35,7 @@ from sglang.srt.mem_cache.unified_cache.components.tree_component import ( TreeComponent, get_and_increase_time_counter, ) -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec, get_server_args if TYPE_CHECKING: from sglang.srt.managers.schedule_batch import Req @@ -66,7 +66,7 @@ class MambaComponent(TreeComponent): ), f"MambaComponent requires page_size=1 when mamba_extra_buffer is disabled, got {params.page_size}" super().__init__(cache, params) self.mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size - self.mamba_max_states_per_path = get_server_args().mamba_max_states_per_path + self.mamba_max_states_per_path = get_exec().mamba.mamba_max_states_per_path # HiCache state self._mamba_pool_host = None # set to host mamba pool when HiCache enabled diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index 28d91fb1c..b55397550 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -37,7 +37,7 @@ from sglang.srt.model_executor.forward_batch_info import ( ) from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.runner_utils.capture_mode import model_capture_mode -from sglang.srt.runtime_context import get_flags, get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel, get_spec from sglang.srt.utils import ( empty_context, log_info_on_rank0, @@ -1027,9 +1027,9 @@ class CPUGraphRunner: retrieve_next_token=None, retrieve_next_sibling=None, retrieve_cum_len=None, - spec_steps=self.model_runner.server_args.speculative_num_steps, + spec_steps=get_spec().speculative_num_steps, topk=self.model_runner.server_args.speculative_eagle_topk, - draft_token_num=self.model_runner.server_args.speculative_num_draft_tokens, + draft_token_num=get_spec().speculative_num_draft_tokens, capture_hidden_mode=CaptureHiddenMode.FULL, seq_lens_sum=None, seq_lens_cpu=None, diff --git a/python/sglang/srt/model_executor/cuda_graph_config.py b/python/sglang/srt/model_executor/cuda_graph_config.py index 2599dc59d..37c07a4a4 100644 --- a/python/sglang/srt/model_executor/cuda_graph_config.py +++ b/python/sglang/srt/model_executor/cuda_graph_config.py @@ -16,7 +16,7 @@ cuda_graph_config, and the --cuda-graph-config JSON CLI parser. Module-level imports are pure stdlib — no torch / sglang.srt deps — so ServerArgs can import everything here without pulling in backend -classes. check_cuda_graph_backend lazy-imports get_server_args +classes. check_cuda_graph_backend lazy-imports the config accessor inside the function body to preserve that invariant. """ @@ -179,15 +179,14 @@ def _diff_phase(actual: PhaseConfig, baseline: PhaseConfig) -> Dict[str, Any]: def check_cuda_graph_backend(phase: str, backend: str) -> bool: """True if cuda_graph_config[phase].backend == backend on the - global server args. Returns False if the global server args have not - been initialized yet (e.g. unit tests, early startup).""" - from sglang.srt.runtime_context import get_server_args + published config. Returns False if the config has not been published + yet (e.g. unit tests, early startup).""" + from sglang.srt.runtime_context import get_exec try: - server_args = get_server_args() + cfg = get_exec().graph.cuda_graph_config except ValueError: return False - cfg = server_args.cuda_graph_config if cfg is None or phase not in Phase.ALL: return False return getattr(cfg, phase).backend == backend diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 4a3e2e215..9447362b9 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -51,7 +51,7 @@ from sglang.srt.layers.dp_attention import ( from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import ( ForwardBatchDeepSeekMHAMixin, ) -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.utils import ( is_cuda, is_hip, @@ -965,7 +965,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): # --enable-mis: every request must carry delimiter indices (the score # endpoint always produces MIS-structured requests; consumers index # without None-checking). - if get_server_args().enable_mis and any( + if get_exec().features.enable_mis and any( r.multi_item_delimiter_indices is not None for r in batch.reqs ): assert all( @@ -1134,7 +1134,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): # batch_size * [3 * seq_len] batch_size = self.seq_lens_cpu.shape[0] mrope_positions_list = [[]] * batch_size - rl_on_policy_target = get_server_args().rl_on_policy_target + rl_on_policy_target = get_exec().deterministic.rl_on_policy_target for batch_idx in range(batch_size): mm_input = batch.multimodal_inputs[batch_idx] if self.forward_mode.is_decode(): diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 03bffbf4f..319ce8db4 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -162,9 +162,13 @@ from sglang.srt.model_executor.runner import ( ) from sglang.srt.platforms import current_platform from sglang.srt.runtime_context import ( + get_context, + get_exec, get_global_dwdp_manager, + get_lora, + get_model, get_parallel, - get_server_args, + get_schedule, set_global_dwdp_manager, ) from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo @@ -322,7 +326,7 @@ class ModelRunner: self.init_threads_binding() # Set float32 matmul precision - if get_server_args().enable_tf32_matmul: + if get_exec().features.enable_tf32_matmul: torch.set_float32_matmul_precision("high") # Set device early so that TransferEngine init (e.g. Ascend NPU) @@ -399,7 +403,7 @@ class ModelRunner: def _initialize_elastic_ep_joiner(self) -> None: if not ( - self.server_args.elastic_ep_backend is not None + get_exec().moe.elastic_ep_backend is not None and self.server_args.is_ep_scale_joiner ): return @@ -473,7 +477,7 @@ class ModelRunner: device=self.device, gpu_id=self.gpu_id, model_config=self.model_config, - custom_weight_loaders=self.server_args.custom_weight_loader, + custom_weight_loaders=get_model().custom_weight_loader, get_model=lambda: self.model, update_model_fields=self.update_model_fields, recapture_cuda_graph=self.init_decode_cuda_graph, @@ -527,6 +531,7 @@ class ModelRunner: model_config=self.model_config, server_args=self.server_args, kv_cache_dtype=self.kv_cache_dtype, + kv_cache_dtype_str=self.kv_cache_dtype_str, model_dtype=self.dtype, page_size=self.page_size, sliding_window_size=self.sliding_window_size, @@ -550,7 +555,7 @@ class ModelRunner: def init_mindspore_runner(self): # Init the mindspore runner # for now, there is only some communication initialization work - if self.server_args.model_impl.lower() == ModelImpl.MINDSPORE and _is_npu: + if get_model().model_impl.lower() == ModelImpl.MINDSPORE and _is_npu: from sglang.srt.model_executor.mindspore_runner import init_ms_distributed init_ms_distributed( @@ -607,7 +612,7 @@ class ModelRunner: def init_memory_saver_adapter(self): self.memory_saver_adapter = TorchMemorySaverAdapter.create( - enable=self.server_args.enable_memory_saver + enable=get_exec().features.enable_memory_saver ) def maybe_init_remote_instance_transfer_engine(self): @@ -643,7 +648,7 @@ class ModelRunner: ) def maybe_init_lplb_solvers(self): - if self.server_args.ep_dispatch_algorithm == "lp" and not self.is_draft_worker: + if get_exec().moe.ep_dispatch_algorithm == "lp" and not self.is_draft_worker: init_lplb_solvers(model_config=self.model_config) def maybe_init_eplb_manager(self): @@ -657,12 +662,12 @@ class ModelRunner: get_expert_backup_client=lambda: self.expert_backup_client, get_weight_updater=lambda: self.weight_updater, ) - if self.server_args.enable_eplb and (not self.is_draft_worker) + if get_exec().moe.enable_eplb and (not self.is_draft_worker) else None ) def maybe_init_elastic_ep(self): - if self.server_args.elastic_ep_backend: + if get_exec().moe.elastic_ep_backend: ElasticEPStateManager.init(self.server_args) def init_token_oracle(self): @@ -681,8 +686,8 @@ class ModelRunner: get_model=lambda: self.model, ) if ( - self.server_args.enable_elastic_expert_backup - and self.server_args.elastic_ep_backend is not None + get_exec().moe.enable_elastic_expert_backup + and get_exec().moe.elastic_ep_backend is not None ) else None ) @@ -691,17 +696,17 @@ class ModelRunner: # In layered loading, torchao may have been applied torchao_applied = getattr(self.model, "torchao_applied", False) if not torchao_applied: - apply_torchao_config_to_model(self.model, get_server_args().torchao_config) + apply_torchao_config_to_model(self.model, get_exec().graph.torchao_config) supports_torch_tp = getattr(self.model, "supports_torch_tp", False) if self.ps.tp_size > 1 and supports_torch_tp: self.apply_torch_tp() def maybe_init_lora_manager(self): - if self.server_args.enable_lora: + if get_lora().enable_lora: self.init_lora_manager() def maybe_enable_batch_invariant_mode(self): - if self.server_args.enable_deterministic_inference: + if get_exec().deterministic.enable_deterministic_inference: from sglang.srt.batch_invariant_ops import enable_batch_invariant_mode enable_batch_invariant_mode() @@ -996,8 +1001,8 @@ class ModelRunner: remote_instance_weight_transporter_engine=self.remote_instance_weight_transporter.engine, remote_instance_weight_transporter_session_id=self.remote_instance_weight_transporter.session_id, draft_model_idx=self.draft_model_idx, - weight_cache_mode=self.server_args.weight_cache_mode, - weight_cache_socket=self.server_args.weight_cache_socket, + weight_cache_mode=get_model().weight_cache_mode, + weight_cache_socket=get_model().weight_cache_socket, ) # If the weight cache is enabled, override the load format to IPC_CACHE @@ -1038,7 +1043,7 @@ class ModelRunner: get_offloader().post_init() # Register model for layerwise NVTX profiling if enabled - if self.server_args.enable_layerwise_nvtx_marker: + if get_exec().comm.enable_layerwise_nvtx_marker: pyt_hooks = PytHooks() pyt_hooks.register_hooks(self.model, module_prefix="model") @@ -1095,7 +1100,7 @@ class ModelRunner: ) dist_barrier_after_load( - elastic_ep_backend=self.server_args.elastic_ep_backend, + elastic_ep_backend=get_exec().moe.elastic_ep_backend, tp_rank=self.ps.tp_rank, is_ep_joiner=self.server_args.is_ep_joiner, ) @@ -1115,16 +1120,16 @@ class ModelRunner: self.lora_manager = LoRAManager( base_model=self.model, base_hf_config=self.model_config.hf_config, - max_loras_per_batch=self.server_args.max_loras_per_batch, + max_loras_per_batch=get_lora().max_loras_per_batch, load_config=self.load_config, dtype=self.dtype, server_args=self.server_args, - lora_backend=self.server_args.lora_backend, + lora_backend=get_lora().lora_backend, tp_size=self.ps.tp_size, tp_rank=self.ps.tp_rank, - max_lora_rank=self.server_args.max_lora_rank, - target_modules=self.server_args.lora_target_modules, - lora_paths=self.server_args.lora_paths, + max_lora_rank=get_lora().max_lora_rank, + target_modules=get_lora().lora_target_modules, + lora_paths=get_lora().lora_paths, ) if not cuda_graph_fully_disabled(): init_lora_cuda_graph_moe_buffers( @@ -1157,29 +1162,10 @@ class ModelRunner: else: return self.max_total_num_tokens - def _record_kv_cache_dtype(self, resolved: str) -> None: - # the weight-resolved kv-cache dtype is written to the config - # bags via get_context().override, so get_model().kv_cache_dtype readers - # see it. server_args stays the pristine RAW record -- configure_kv_cache - # _dtype reads it as the resolver INPUT. A draft / mock runner whose - # server_args is not the published object keeps the private-bag write. - from sglang.srt.runtime_context import get_context - - if get_context()._server_args is self.server_args: - get_context().override( - "ModelRunner.configure_kv_cache_dtype", kv_cache_dtype=resolved - ) - else: - self.server_args.override( - "ModelRunner.configure_kv_cache_dtype", kv_cache_dtype=resolved - ) - def configure_kv_cache_dtype(self): spec_algorithm = getattr(self, "spec_algorithm", None) resolved_kv_cache_dtype, self.kv_cache_dtype = ( kv_cache_dtype.configure_kv_cache_dtype( - # RAW user intent = resolver INPUT; server_args stays pristine - # so read it here -- not the resolved get_model() bag. server_args_kv_cache_dtype=self.server_args.kv_cache_dtype, model=getattr(self, "model", None), model_dtype=getattr(self, "dtype", torch.bfloat16), @@ -1201,8 +1187,6 @@ class ModelRunner: if resolved_kv_cache_dtype is not None else self.server_args.kv_cache_dtype ) - if resolved_kv_cache_dtype is not None: - self._record_kv_cache_dtype(resolved_kv_cache_dtype) def _get_attention_backend(self, init_new_workspace: bool = False): return get_attention_backend( @@ -1391,7 +1375,7 @@ class ModelRunner: ) output.expert_distribution_metrics = recorder_outputs.get("metrics") - no_copy_to_cpu = not self.server_args.disable_overlap_schedule + no_copy_to_cpu = not get_schedule().disable_overlap_schedule if ( not self.is_draft_worker and (experts_capturer := get_global_experts_capturer()) is not None @@ -1421,7 +1405,7 @@ class ModelRunner: self.msprobe_debugger.stop() self.msprobe_debugger.step() - if self.server_args.elastic_ep_backend is not None: + if get_exec().moe.elastic_ep_backend is not None: self.maybe_join_ep_ranks() return output @@ -1852,7 +1836,7 @@ class ModelRunner: local_timeout = ( state.pending_since is not None and time.monotonic() - state.pending_since - > self.server_args.elastic_ep_scale_timeout + > get_exec().moe.elastic_ep_scale_timeout ) timeout = state.active_ranks.new_tensor(int(local_timeout)) dist.all_reduce(timeout, op=dist.ReduceOp.MAX, group=dist.group.WORLD) @@ -1922,7 +1906,7 @@ class ModelRunner: load_config: LoadConfig, ) -> None: self.model = new_model - self.server_args.override( + get_context().override( "model_runner.update_model_fields", model_path=model_path, load_format=load_format, diff --git a/python/sglang/srt/model_executor/model_runner_components/remote_instance_weight_transporter.py b/python/sglang/srt/model_executor/model_runner_components/remote_instance_weight_transporter.py index bf56487ad..1f4bb50ff 100644 --- a/python/sglang/srt/model_executor/model_runner_components/remote_instance_weight_transporter.py +++ b/python/sglang/srt/model_executor/model_runner_components/remote_instance_weight_transporter.py @@ -11,6 +11,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import ( RemoteInstanceWeightLoaderBackend, register_memory_region, ) +from sglang.srt.runtime_context import get_model from sglang.srt.server_args import ServerArgs from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto @@ -58,7 +59,7 @@ class RemoteInstanceWeightTransporter: # ModelExpress owns TransferEngine memory registration and metadata # publishing for backend=modelexpress. Re-registering here would # overlap the same weight buffers. - and self.server_args.remote_instance_weight_loader_backend + and get_model().remote_instance_weight_loader_backend != RemoteInstanceWeightLoaderBackend.MODELEXPRESS and self.engine is not None and self.weight_info is None @@ -84,7 +85,7 @@ class RemoteInstanceWeightTransporter: else: bootstrap_host = "127.0.0.1" - bootstrap_port = self.server_args.engine_info_bootstrap_port + bootstrap_port = get_model().engine_info_bootstrap_port bootstrap_na = NetworkAddress(bootstrap_host, bootstrap_port) url = f"{bootstrap_na.to_url()}/register_transfer_engine_info" diff --git a/python/sglang/srt/model_executor/pool_configurator.py b/python/sglang/srt/model_executor/pool_configurator.py index 194dd9fef..b05f7358a 100644 --- a/python/sglang/srt/model_executor/pool_configurator.py +++ b/python/sglang/srt/model_executor/pool_configurator.py @@ -33,7 +33,7 @@ from sglang.srt.environ import envs from sglang.srt.mem_cache.allocation_sizing import get_alloc_len_per_decode from sglang.srt.mem_cache.deepseek_v4_memory_pool import get_compress_state_ring_size from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool -from sglang.srt.runtime_context import get_model, get_parallel +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils.common import ( ceil_align, ceil_div, @@ -119,6 +119,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator): """ def __init__(self, kvc: KVCacheConfigurator): + self.kv_cache_dtype_str = kvc.kv_cache_dtype_str # Determine effective number of layers for KV cache if mambaish := mambaish_config(kvc.model_config): effective_layer_ids = [ @@ -304,7 +305,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator): ) # FP4 prefill uses one shared FP8 dequant workspace across layers. cell_size += n * k * 2 * kv_size - elif get_model().kv_cache_dtype == "mxfp8": + elif self.kv_cache_dtype_str == "mxfp8": scale_block_size = 32 n = model_config.get_num_kv_heads(tp_size) cell_size += ( @@ -339,6 +340,7 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator): """ def __init__(self, kvc: KVCacheConfigurator): + self.kv_cache_dtype_str = kvc.kv_cache_dtype_str model_config = kvc.model_config kv_cache_dtype = kvc.kv_cache_dtype kv_size = torch._utils._element_size(kv_cache_dtype) @@ -368,7 +370,7 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator): * kv_size ) - if get_model().kv_cache_dtype == "mxfp8": + if self.kv_cache_dtype_str == "mxfp8": scale_block_size = 32 self._full_per_token += ( model_config.get_num_kv_heads(tp_size) @@ -501,6 +503,7 @@ class SWAChunkCapPoolConfigurator(HybridSWAPoolConfigurator): """ def __init__(self, kvc: KVCacheConfigurator): + self.kv_cache_dtype_str = kvc.kv_cache_dtype_str super().__init__(kvc) assert self._full_layers_num > 0 @@ -613,6 +616,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator): """ def __init__(self, kvc: KVCacheConfigurator): + self.kv_cache_dtype_str = kvc.kv_cache_dtype_str cfg = kvc.model_config self.qk_nope_head_dim = cfg.qk_nope_head_dim self.qk_rope_head_dim = cfg.qk_rope_head_dim 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 9da4ba737..6be24d982 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 @@ -91,7 +91,7 @@ from sglang.srt.model_executor.runner_utils.deepep_adapter import ( DeepEPCudaGraphRunnerAdapter, ) from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups -from sglang.srt.runtime_context import get_flags, get_parallel +from sglang.srt.runtime_context import get_flags, get_parallel, get_spec from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout from sglang.srt.utils import ( empty_context, @@ -246,12 +246,12 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): self.is_dllm = self.dllm_config is not None self.attn_backend = attn_backend or model_runner.attn_backend self.speculative_num_steps = ( - model_runner.server_args.speculative_num_steps + get_spec().speculative_num_steps if speculative_num_steps is None else speculative_num_steps ) self.speculative_num_draft_tokens = ( - model_runner.server_args.speculative_num_draft_tokens + get_spec().speculative_num_draft_tokens if speculative_num_draft_tokens is None else speculative_num_draft_tokens ) diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index db18b4823..73f899add 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -47,7 +47,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import ( get_remote_instance_transfer_engine_info_per_rank, register_memory_region, ) -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec, get_model, get_server_args from sglang.srt.utils import get_available_gpu_memory # Try to import accelerate (optional dependency) @@ -495,10 +495,10 @@ class DefaultModelLoader(BaseModelLoader): hf_folder = model_name_or_path server_args = get_server_args() - if server_args and server_args.model_checksum is not None: + if server_args and get_model().model_checksum is not None: from sglang.srt.utils.model_file_verifier import verify - checksums_source = server_args.model_checksum or model_name_or_path + checksums_source = get_model().model_checksum or model_name_or_path verify(model_path=hf_folder, checksums_source=checksums_source) hf_weights_files: List[str] = [] @@ -581,11 +581,11 @@ class DefaultModelLoader(BaseModelLoader): ) elif use_safetensors: server_args = get_server_args() - weight_loader_disable_mmap = server_args.weight_loader_disable_mmap - weight_loader_prefetch = server_args.weight_loader_prefetch_checkpoints - prefetch_num_threads = server_args.weight_loader_prefetch_num_threads + weight_loader_disable_mmap = get_model().weight_loader_disable_mmap + weight_loader_prefetch = get_model().weight_loader_prefetch_checkpoints + prefetch_num_threads = get_model().weight_loader_prefetch_num_threads weight_loader_drop_cache_after_load = ( - server_args.weight_loader_drop_cache_after_load + get_model().weight_loader_drop_cache_after_load ) # Prefetch and multi-threaded loading both read the same shards, @@ -879,9 +879,8 @@ class LayeredModelLoader(DefaultModelLoader): device_config: DeviceConfig, ) -> nn.Module: from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model - from sglang.srt.runtime_context import get_server_args - torchao_config = get_server_args().torchao_config + torchao_config = get_exec().graph.torchao_config target_device = torch.device(device_config.device) quant_config = _get_quantization_config(model_config, self.load_config) @@ -1751,13 +1750,13 @@ class PreshardedModelLoader(DefaultModelLoader): "moe_dense_tp_size": server_args.moe_dense_tp_size, "moe_dp_size": server_args.moe_dp_size, "enable_dp_lm_head": server_args.enable_dp_lm_head, - "enable_fp32_lm_head": server_args.enable_fp32_lm_head, + "enable_fp32_lm_head": get_exec().features.enable_fp32_lm_head, "quantization": model_config.quantization, "model_dtype": str(model_config.dtype), - "ep_num_redundant_experts": server_args.ep_num_redundant_experts, - "enable_eplb": server_args.enable_eplb, + "ep_num_redundant_experts": get_exec().moe.ep_num_redundant_experts, + "enable_eplb": get_exec().moe.enable_eplb, "init_expert_location": self._normalize_init_expert_location( - server_args.init_expert_location + get_exec().moe.init_expert_location ), "structural_signature": self._compute_structural_signature(model_config), } @@ -3934,10 +3933,10 @@ class RunaiModelStreamerLoader(BaseModelLoader): ) server_args = get_server_args() - if server_args and server_args.model_checksum is not None: + if server_args and get_model().model_checksum is not None: from sglang.srt.utils.model_file_verifier import verify - checksums_source = server_args.model_checksum or model_name_or_path + checksums_source = get_model().model_checksum or model_name_or_path verify(model_path=hf_folder, checksums_source=checksums_source) hf_weights_files = list_safetensors(path=hf_folder) diff --git a/python/sglang/srt/models/bailing_moe.py b/python/sglang/srt/models/bailing_moe.py index 9abccdbe0..1f1297c57 100644 --- a/python/sglang/srt/models/bailing_moe.py +++ b/python/sglang/srt/models/bailing_moe.py @@ -78,6 +78,7 @@ from sglang.srt.models.utils import ( enable_fused_set_kv_buffer, ) from sglang.srt.runtime_context import ( + get_exec, get_forward, get_parallel, get_server_args, @@ -209,7 +210,7 @@ class BailingMoESparseMoeBlock(nn.Module): self.router_dtype = torch.bfloat16 # TODO global_server_args.ep_num_redundant_experts is used for eplb, not supported now - assert get_server_args().ep_num_redundant_experts == 0 + assert get_exec().moe.ep_num_redundant_experts == 0 # check group topk self.num_expert_group = getattr(config, "n_group", 0) self.topk_group = getattr(config, "topk_group", 0) @@ -223,9 +224,7 @@ class BailingMoESparseMoeBlock(nn.Module): self.num_expert_group = self.topk_group = None self.use_grouped_topk = False - self.num_experts = ( - config.num_experts + get_server_args().ep_num_redundant_experts - ) + self.num_experts = config.num_experts + get_exec().moe.ep_num_redundant_experts self.gate = BailingMoEGate( config=config, diff --git a/python/sglang/srt/models/bailing_moe_linear.py b/python/sglang/srt/models/bailing_moe_linear.py index 27b265814..bc62cc009 100644 --- a/python/sglang/srt/models/bailing_moe_linear.py +++ b/python/sglang/srt/models/bailing_moe_linear.py @@ -59,6 +59,7 @@ from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip from sglang.srt.models.utils import WeightsMapper from sglang.srt.runtime_context import ( + get_device, get_forward, get_parallel, get_server_args, @@ -529,7 +530,7 @@ class BailingMoELinearAttention(nn.Module): base=self.rope_theta, rope_scaling=config.rope_scaling, is_neox_style=True, - device=get_server_args().device, + device=get_device().device, dtype=torch.float32, ) @@ -690,7 +691,7 @@ class BailingMoEAttention(nn.Module): max_position=self.max_position_embeddings, base=self.rope_theta, rope_scaling=config.rope_scaling, - device=get_server_args().device, + device=get_device().device, ) self.attn = RadixAttention( self.num_heads, diff --git a/python/sglang/srt/models/bert.py b/python/sglang/srt/models/bert.py index 82881395f..154900912 100644 --- a/python/sglang/srt/models/bert.py +++ b/python/sglang/srt/models/bert.py @@ -16,7 +16,7 @@ from sglang.srt.layers.radix_attention import AttentionType, RadixAttention from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_loader.weight_utils import default_weight_loader -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_model, get_parallel from sglang.srt.utils import add_prefix BertConfig = None @@ -365,9 +365,7 @@ class BertModel(nn.Module): quant_config=quant_config, prefix=add_prefix("encoder", prefix), ) - pooling_type = ( - PoolingType.CLS if get_server_args().is_embedding else PoolingType.LAST - ) + pooling_type = PoolingType.CLS if get_model().is_embedding else PoolingType.LAST self.pooler = ( BertPooler(config) if self.use_bert_pooler diff --git a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py index f55634729..9445213bb 100644 --- a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py +++ b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py @@ -11,7 +11,7 @@ from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods AttnForwardMethod, ) from sglang.srt.models.deepseek_common.utils import _is_hip -from sglang.srt.runtime_context import get_exec, get_server_args +from sglang.srt.runtime_context import get_exec from sglang.srt.utils import is_sm100_or_sm110_supported, use_intel_amx_backend MHA_ONE_SHOT_SUPPORTED_BACKENDS = ["fa3", "flashinfer", "flashmla"] @@ -118,7 +118,7 @@ def handle_attention_flashinfer(attn, forward_batch): def handle_attention_fa3(attn, forward_batch): # when deterministic inference is enabled, use MLA - if get_server_args().enable_deterministic_inference: + if get_exec().deterministic.enable_deterministic_inference: return _dispatch_mla_subtype(attn, forward_batch) else: return _handle_attention_backend(attn, forward_batch, "fa3") @@ -194,7 +194,7 @@ def handle_attention_triton(attn, forward_batch): return AttnForwardMethod.MLA # when deterministic inference is enabled, use MLA - if get_server_args().enable_deterministic_inference: + if get_exec().deterministic.enable_deterministic_inference: return _dispatch_mla_subtype(attn, forward_batch) if ( diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py index c661322ac..19b4805c1 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha.py @@ -30,7 +30,11 @@ from sglang.srt.models.deepseek_common.utils import ( _use_aiter_bpreshuffle_gfx95, _use_aiter_gfx95, ) -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_exec, + get_parallel, + get_schedule, +) from sglang.srt.utils import BumpAllocator, get_bool_env_var, next_power_of_2 _use_fp8_prefill_attn = ( @@ -142,9 +146,7 @@ def _forward_dsa_indexer_for_mha( class DeepseekMHAForwardMixin: def init_mha_forward(self: DeepseekV2AttentionMLA): - self.disable_chunked_prefix_cache = ( - get_server_args().disable_chunked_prefix_cache - ) + self.disable_chunked_prefix_cache = get_schedule().disable_chunked_prefix_cache # TODO: Design a finer way to determine the threshold self.chunked_prefix_cache_threshold = ( @@ -305,8 +307,8 @@ class DeepseekMHAForwardMixin: self.use_dsa and self.kv_cache_dtype == "fp8_e4m3" and ( - not get_server_args().dsa_decode_backend == "trtllm" - or not get_server_args().dsa_prefill_backend == "trtllm" + not get_exec().kernel.dsa_decode_backend == "trtllm" + or not get_exec().kernel.dsa_prefill_backend == "trtllm" ) ): # FP8 path: dequantize DSA-specific FP8 format to BF16 diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py index 604bbcdd2..d92b18c03 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -68,7 +68,7 @@ from sglang.srt.models.deepseek_common.utils import ( _use_aiter_bpreshuffle_gfx95, _use_aiter_gfx95, ) -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args, get_spec from sglang.srt.state_capturer.indexer_topk import ( maybe_capture_indexer_topk, ) @@ -105,8 +105,8 @@ def _is_dcp_mla_decode_phase(forward_batch: ForwardBatch) -> bool: server_args.decode_attention_backend or server_args.attention_backend ) return ( - server_args.speculative_algorithm == "DSPARK" - and server_args.speculative_attention_mode == "decode" + get_spec().speculative_algorithm == "DSPARK" + and get_spec().speculative_attention_mode == "decode" and decode_backend in ("tokenspeed_mla", "cutedsl_mla") ) @@ -177,7 +177,7 @@ def _should_defer_dsa_cp_kv_gather( class DeepseekMLAForwardMixin: def init_mla_forward(self: DeepseekV2AttentionMLA): self.flashinfer_mla_disable_ragged = ( - get_server_args().flashinfer_mla_disable_ragged + get_exec().kernel.flashinfer_mla_disable_ragged ) def should_run_indexer( @@ -1239,8 +1239,8 @@ class DeepseekMLAForwardMixin: """ if self.current_attention_backend in ("dsa", "nsa"): return ( - get_server_args().dsa_decode_backend == "trtllm" - or get_server_args().dsa_prefill_backend == "trtllm" + get_exec().kernel.dsa_decode_backend == "trtllm" + or get_exec().kernel.dsa_prefill_backend == "trtllm" ) and get_attn_backend().kv_cache_dtype == torch.float8_e4m3fn return ( @@ -1263,8 +1263,8 @@ class DeepseekMLAForwardMixin: _use_aiter_gfx95 and self.current_attention_backend in ("dsa", "nsa") and ( - server_args.dsa_decode_backend == "tilelang" - or server_args.dsa_prefill_backend == "tilelang" + get_exec().kernel.dsa_decode_backend == "tilelang" + or get_exec().kernel.dsa_prefill_backend == "tilelang" ) ) diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py index 741487484..be937318c 100644 --- a/python/sglang/srt/models/deepseek_nextn.py +++ b/python/sglang/srt/models/deepseek_nextn.py @@ -59,7 +59,12 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.models.deepseek_common.utils import enable_nextn_moe_bf16_cast_to_fp8 from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM from sglang.srt.models.utils import WeightsMapper -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_model, + get_parallel, + get_server_args, + get_spec, +) from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu @@ -148,7 +153,7 @@ class DeepseekModelNextN(nn.Module): self.rot_weight = None if _is_npu: - rot_weight_path = get_server_args().model_path + "/rot.safetensors" + rot_weight_path = get_model().model_path + "/rot.safetensors" if os.path.isfile(rot_weight_path): self.rot_weight = load_file(rot_weight_path) self.rot_weight = self.rot_weight["rot.weight"].npu() @@ -161,8 +166,7 @@ class DeepseekModelNextN(nn.Module): layer_name = "decoder" if _is_npu and ( - get_server_args().speculative_draft_model_path - == get_server_args().model_path + get_spec().speculative_draft_model_path == get_model().model_path ): layer_name = "layers." + str(config.num_hidden_layers) @@ -201,7 +205,7 @@ class DeepseekModelNextN(nn.Module): if ( _is_npu and self.quant_config is None - and get_server_args().quantization is not None + and get_model().quantization is not None ): # ascend mtp unquant exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True)) diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index e6ffbc034..e63313f85 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -188,12 +188,14 @@ from sglang.srt.models.deepseek_common.utils import ( is_wint4afp8_or_wint4a16_config, ) from sglang.srt.runtime_context import ( + get_device, get_exec, get_flags, get_forward, get_model, get_parallel, get_server_args, + get_spec, ) from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.utils import ( @@ -503,7 +505,7 @@ class MoEGate(nn.Module): True, # is_vnni ) - if get_server_args().enable_deterministic_inference: + if get_exec().deterministic.enable_deterministic_inference: return F.linear(hidden_states, self.weight, None) if ( @@ -569,7 +571,7 @@ class DeepseekV2MoE(nn.Module): n_shared_experts = ( 0 if config.n_shared_experts is None else int(config.n_shared_experts) ) - _fusion_disabled = get_server_args().disable_shared_experts_fusion + _fusion_disabled = get_exec().moe.disable_shared_experts_fusion # num_fused_shared_experts drives weight remapping in deepseek_weight_loader: # mlp.shared_experts → mlp.experts.256 when > 0. @@ -639,8 +641,7 @@ class DeepseekV2MoE(nn.Module): fused_shared_experts_scaling_factor = 1.0 / float(self.moe_ep_size) self.experts = get_moe_impl_class(quant_config)( - num_experts=num_experts_for_moe - + get_server_args().ep_num_redundant_experts, + num_experts=num_experts_for_moe + get_exec().moe.ep_num_redundant_experts, num_fused_shared_experts=self.num_fused_shared_experts, top_k=top_k_for_moe, hidden_size=config.hidden_size, @@ -814,7 +815,7 @@ class DeepseekV2MoE(nn.Module): # TODO: we will support tp < ep in the future self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.n_routed_experts + get_server_args().ep_num_redundant_experts + config.n_routed_experts + get_exec().moe.ep_num_redundant_experts ) self.renormalize = config.norm_topk_prob self.topk_group = config.topk_group @@ -911,8 +912,8 @@ class DeepseekV2MoE(nn.Module): and not ( get_flags().capture.enable_torch_compile and hidden_states.shape[0] - <= server_args.torch_compile_max_bs - * (server_args.speculative_num_draft_tokens or 1) + <= get_exec().graph.torch_compile_max_bs + * (get_spec().speculative_num_draft_tokens or 1) ) ): return self.forward_normal_dual_stream( @@ -962,7 +963,7 @@ class DeepseekV2MoE(nn.Module): server_args = get_server_args() dispatch_info = ( ExpertLocationDispatchInfo.init_new(layer_id=self.layer_id) - if server_args.enable_eplb and not self.is_nextn + if get_exec().moe.enable_eplb and not self.is_nextn else None ) # router_logits: (num_tokens, n_experts) @@ -1063,7 +1064,7 @@ class DeepseekV2MoE(nn.Module): server_args = get_server_args() dispatch_info = ( ExpertLocationDispatchInfo.init_new(layer_id=self.layer_id) - if server_args.enable_eplb and not self.is_nextn + if get_exec().moe.enable_eplb and not self.is_nextn else None ) defer_shared = not self.experts.moe_runner_config.inplace @@ -1879,7 +1880,7 @@ class DeepseekV2AttentionMLA( base=rope_theta, rope_scaling=rope_scaling, is_neox_style=is_neox_style, - device=get_server_args().device, + device=get_device().device, ) if rope_scaling and rope_scaling.get("apply_yarn_scaling", True): @@ -2011,7 +2012,7 @@ class DeepseekV2AttentionMLA( or forward_batch.forward_mode.is_draft_extend_v2() ): # Use the specified backend for speculative operations (both verify and draft extend) - if server_args.speculative_attention_mode == "decode": + if get_spec().speculative_attention_mode == "decode": attention_backend = decode_backend_str else: # default to prefill attention_backend = prefill_backend_str @@ -2270,7 +2271,7 @@ class DeepseekV2DecoderLayer(nn.Module): rope_scaling = config.rope_scaling max_position_embeddings = config.max_position_embeddings self.speculative_algorithm = SpeculativeAlgorithm.from_string( - get_server_args().speculative_algorithm + get_spec().speculative_algorithm ) self.dsa_enable_prefill_cp = dsa_enable_prefill_cp self.mla_enable_prefill_cp = mla_enable_prefill_cp @@ -2968,11 +2969,11 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): self.num_fused_shared_experts = 0 server_args = get_server_args() - if get_server_args().disable_shared_experts_fusion: + if get_exec().moe.disable_shared_experts_fusion: return disable_reason = None - if server_args.enforce_shared_experts_fusion: + if get_exec().moe.enforce_shared_experts_fusion: pass elif is_sbo_enabled() or is_tbo_enabled(): disable_reason = "SBO/TBO enabled: incompatible with fusing shared expert into MoE kernel." diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index a81b5de4d..e64cf2c1a 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -138,7 +138,13 @@ from sglang.srt.models.deepseek_v2 import ( _is_npu, _is_xpu, ) -from sglang.srt.runtime_context import get_forward, get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_device, + get_exec, + get_forward, + get_parallel, + get_server_args, +) if not _is_hip: from sglang.srt.layers.utils.cp_utils import ( @@ -620,7 +626,7 @@ class MQALayer(MqaAttentionBase): base=self.rope_base, rope_scaling=self.rope_scaling, is_neox_style=False, - device=get_server_args().device, + device=get_device().device, ) if _is_npu: @@ -2546,11 +2552,11 @@ class DeepseekV4ForCausalLM(nn.Module): def determine_num_fused_shared_experts(self): self.num_fused_shared_experts = 0 - if get_server_args().disable_shared_experts_fusion: + if get_exec().moe.disable_shared_experts_fusion: return disable_reason = None - if get_server_args().enforce_shared_experts_fusion: + if get_exec().moe.enforce_shared_experts_fusion: if self.config.n_shared_experts != 1: raise ValueError( "DeepSeek V4 shared-experts fusion expects exactly one shared " diff --git a/python/sglang/srt/models/exaone_moe.py b/python/sglang/srt/models/exaone_moe.py index b90c75efe..88e707642 100755 --- a/python/sglang/srt/models/exaone_moe.py +++ b/python/sglang/srt/models/exaone_moe.py @@ -62,7 +62,12 @@ from sglang.srt.layers.vocab_parallel_embedding import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_loader.weight_utils import default_weight_loader -from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream +from sglang.srt.runtime_context import ( + get_exec, + get_parallel, + get_server_args, + get_stream, +) from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers logger = logging.getLogger(__name__) @@ -165,7 +170,7 @@ class ExaoneMoESparseMoEBlock(nn.Module): ) self.experts = get_moe_impl_class(quant_config)( - num_experts=config.num_experts + get_server_args().ep_num_redundant_experts, + num_experts=config.num_experts + get_exec().moe.ep_num_redundant_experts, top_k=config.num_experts_per_tok, hidden_size=config.hidden_size, intermediate_size=config.moe_intermediate_size, @@ -206,7 +211,7 @@ class ExaoneMoESparseMoEBlock(nn.Module): if get_moe_a2a_backend().is_deepep(): self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.num_experts + get_server_args().ep_num_redundant_experts + config.num_experts + get_exec().moe.ep_num_redundant_experts ) self.top_k = config.num_experts_per_tok diff --git a/python/sglang/srt/models/gemma4_causal.py b/python/sglang/srt/models/gemma4_causal.py index f2d9a645e..b37cc7dd4 100644 --- a/python/sglang/srt/models/gemma4_causal.py +++ b/python/sglang/srt/models/gemma4_causal.py @@ -58,7 +58,7 @@ from sglang.srt.models.gemma3_causal import Gemma3MLP, Gemma3TextScaledWordEmbed from sglang.srt.models.utils import ( create_fused_set_kv_buffer_arg, ) -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args from sglang.srt.utils import add_prefix, make_layers logger = logging.getLogger(__name__) @@ -254,7 +254,7 @@ class Gemma4MoE(nn.Module): experts_type = get_moe_impl_class(quant_config) self.experts = experts_type( - num_experts=config.num_experts + get_server_args().ep_num_redundant_experts, + num_experts=config.num_experts + get_exec().moe.ep_num_redundant_experts, hidden_size=config.hidden_size, intermediate_size=config.moe_intermediate_size, layer_id=layer_id, @@ -788,7 +788,7 @@ class Gemma4TextModel(PreTrainedModel): # PP + PLE eagerly with --disable-cuda-graph. if self.pp_group.world_size > 1 and self.hidden_size_per_layer_input > 0: sa = get_server_args() - if sa is not None and not sa.disable_cuda_graph: + if sa is not None and not get_exec().graph.disable_cuda_graph: raise ValueError( "Pipeline parallelism is currently incompatible with " "per-layer-input (PLE) embeddings under CUDA graph: " diff --git a/python/sglang/srt/models/gemma4_vision.py b/python/sglang/srt/models/gemma4_vision.py index 7e440555c..63fa2d064 100644 --- a/python/sglang/srt/models/gemma4_vision.py +++ b/python/sglang/srt/models/gemma4_vision.py @@ -29,7 +29,7 @@ from sglang.srt.layers.clippable_linear import ( ) from sglang.srt.layers.layernorm import Gemma4RMSNorm from sglang.srt.layers.quantization.base_config import QuantizationConfig -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_mm, get_parallel from sglang.srt.utils import add_prefix, get_device_capability, is_cuda, is_hip # --------------------------------------------------------------------------- @@ -181,9 +181,8 @@ class Gemma4VisionAttention(nn.Module): @staticmethod def _select_backend() -> str: """Mirror VisionAttention._determine_attention_backend for consistency.""" - from sglang.srt.runtime_context import get_server_args - override = get_server_args().mm_attention_backend + override = get_mm().mm_attention_backend if override is not None: return override if is_cuda(): diff --git a/python/sglang/srt/models/glm4_moe.py b/python/sglang/srt/models/glm4_moe.py index 758aeaa52..1e48115c3 100644 --- a/python/sglang/srt/models/glm4_moe.py +++ b/python/sglang/srt/models/glm4_moe.py @@ -84,6 +84,7 @@ from sglang.srt.models.deepseek_nextn import DeepseekV3ForCausalLMNextN from sglang.srt.models.deepseek_v2 import DeepseekV2ForCausalLM from sglang.srt.models.utils import WeightsMapper, apply_qk_norm from sglang.srt.runtime_context import ( + get_exec, get_forward, get_parallel, get_server_args, @@ -406,7 +407,7 @@ class Glm4MoeSparseMoeBlock(nn.Module): self.n_shared_experts = config.n_shared_experts self.num_fused_shared_experts = ( 0 - if get_server_args().disable_shared_experts_fusion + if get_exec().moe.disable_shared_experts_fusion else config.n_shared_experts ) @@ -526,7 +527,7 @@ class Glm4MoeSparseMoeBlock(nn.Module): # TODO: we will support tp < ep in the future self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.n_routed_experts + get_server_args().ep_num_redundant_experts + config.n_routed_experts + get_exec().moe.ep_num_redundant_experts ) self.renormalize = config.norm_topk_prob self.topk_group = config.topk_group @@ -1178,7 +1179,7 @@ class Glm4MoeForCausalLM(nn.Module): self.capture_aux_hidden_states = False def determine_num_fused_shared_experts(self): - if get_server_args().disable_shared_experts_fusion: + if get_exec().moe.disable_shared_experts_fusion: return disable_reason = None diff --git a/python/sglang/srt/models/glm4_moe_lite.py b/python/sglang/srt/models/glm4_moe_lite.py index 68f9dea7c..30e2c7f6a 100644 --- a/python/sglang/srt/models/glm4_moe_lite.py +++ b/python/sglang/srt/models/glm4_moe_lite.py @@ -75,6 +75,7 @@ from sglang.srt.models.deepseek_common.deepseek_weight_loader import ( from sglang.srt.models.deepseek_common.utils import _is_cuda, _use_aiter from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA from sglang.srt.runtime_context import ( + get_exec, get_forward, get_parallel, get_server_args, @@ -191,7 +192,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module): self.n_shared_experts = config.n_shared_experts self.num_fused_shared_experts = ( 0 - if get_server_args().disable_shared_experts_fusion + if get_exec().moe.disable_shared_experts_fusion else config.n_shared_experts ) self.config = config @@ -218,7 +219,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module): self.experts = get_moe_impl_class(quant_config)( num_experts=config.n_routed_experts + self.num_fused_shared_experts - + get_server_args().ep_num_redundant_experts, + + get_exec().moe.ep_num_redundant_experts, num_fused_shared_experts=self.num_fused_shared_experts, top_k=config.num_experts_per_tok + self.num_fused_shared_experts, hidden_size=config.hidden_size, @@ -287,7 +288,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module): # TODO: we will support tp < ep in the future self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.n_routed_experts + get_server_args().ep_num_redundant_experts + config.n_routed_experts + get_exec().moe.ep_num_redundant_experts ) self.renormalize = config.norm_topk_prob self.topk_group = config.topk_group @@ -931,7 +932,7 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin): self, architecture: str = "Glm4MoeLiteForCausalLM" ): self.num_fused_shared_experts = 0 - if get_server_args().disable_shared_experts_fusion: + if get_exec().moe.disable_shared_experts_fusion: return disable_reason = None diff --git a/python/sglang/srt/models/glm4_moe_lite_nextn.py b/python/sglang/srt/models/glm4_moe_lite_nextn.py index a8ad68b18..9c91bbae2 100644 --- a/python/sglang/srt/models/glm4_moe_lite_nextn.py +++ b/python/sglang/srt/models/glm4_moe_lite_nextn.py @@ -35,7 +35,7 @@ from sglang.srt.models.glm4_moe_lite import ( Glm4MoeLiteDecoderLayer, Glm4MoeLiteForCausalLM, ) -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args, get_spec from sglang.srt.utils import BumpAllocator, add_prefix, is_npu logger = logging.getLogger(__name__) @@ -139,7 +139,7 @@ class Glm4MoeLiteForCausalLMNextN(Glm4MoeLiteForCausalLM): nn.Module.__init__(self) self.config = config self.tp_size = get_parallel().tp_size - if is_npu() and get_server_args().speculative_draft_model_quantization is None: + if is_npu() and get_spec().speculative_draft_model_quantization is None: quant_config = None self.quant_config = quant_config @@ -156,7 +156,7 @@ class Glm4MoeLiteForCausalLMNextN(Glm4MoeLiteForCausalLM): self.logits_processor = LogitsProcessor(config) self.num_fused_shared_experts = ( - 0 if get_server_args().disable_shared_experts_fusion else 1 + 0 if get_exec().moe.disable_shared_experts_fusion else 1 ) @torch.no_grad() diff --git a/python/sglang/srt/models/glm4_moe_nextn.py b/python/sglang/srt/models/glm4_moe_nextn.py index 3126fd026..5804bf241 100644 --- a/python/sglang/srt/models/glm4_moe_nextn.py +++ b/python/sglang/srt/models/glm4_moe_nextn.py @@ -32,7 +32,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.models.glm4_moe import Glm4MoeDecoderLayer, Glm4MoeForCausalLM -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args, get_spec from sglang.srt.utils import add_prefix, is_npu logger = logging.getLogger(__name__) @@ -125,7 +125,7 @@ class Glm4MoeForCausalLMNextN(Glm4MoeForCausalLM): nn.Module.__init__(self) self.config = config self.tp_size = get_parallel().tp_size - if is_npu() and get_server_args().speculative_draft_model_quantization is None: + if is_npu() and get_spec().speculative_draft_model_quantization is None: quant_config = None self.quant_config = quant_config @@ -142,7 +142,7 @@ class Glm4MoeForCausalLMNextN(Glm4MoeForCausalLM): self.logits_processor = LogitsProcessor(config) self.num_fused_shared_experts = ( - 0 if get_server_args().disable_shared_experts_fusion else 1 + 0 if get_exec().moe.disable_shared_experts_fusion else 1 ) @torch.no_grad() diff --git a/python/sglang/srt/models/glm4v.py b/python/sglang/srt/models/glm4v.py index 598c6a22f..43a74f4ad 100644 --- a/python/sglang/srt/models/glm4v.py +++ b/python/sglang/srt/models/glm4v.py @@ -57,7 +57,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTe from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.glm4 import Glm4Model from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_mm, get_parallel from sglang.srt.utils import add_prefix, is_npu from sglang.srt.utils.hf_transformers_utils import get_processor @@ -558,7 +558,7 @@ class Glm4vForConditionalGeneration(nn.Module): self.pp_group = get_pp_group() self.config = config - self.use_data_parallel = get_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_mm().mm_enable_dp_encoder vision_utils.update_vit_attn_dummy_heads_config(self.config) self.visual = Glm4vVisionModel( config.vision_config, diff --git a/python/sglang/srt/models/glm4v_moe.py b/python/sglang/srt/models/glm4v_moe.py index c69899003..38fdc0a64 100644 --- a/python/sglang/srt/models/glm4v_moe.py +++ b/python/sglang/srt/models/glm4v_moe.py @@ -18,7 +18,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.glm4_moe import Glm4MoeModel from sglang.srt.models.glm4v import Glm4vForConditionalGeneration, Glm4vVisionModel -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_mm, get_parallel, get_server_args from sglang.srt.utils import add_prefix, get_device_sm, is_cuda, log_info_on_rank0 from sglang.srt.utils.hf_transformers_utils import get_processor @@ -41,7 +41,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): self.pp_group = get_pp_group() self.config = config - self.use_data_parallel = get_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_mm().mm_enable_dp_encoder vision_utils.update_vit_attn_dummy_heads_config(self.config) self.tp_size = get_parallel().tp_size self.quant_config = quant_config @@ -83,7 +83,7 @@ class Glm4vMoeForConditionalGeneration(Glm4vForConditionalGeneration): self.capture_aux_hidden_states = False def determine_num_fused_shared_experts(self): - if get_server_args().disable_shared_experts_fusion: + if get_exec().moe.disable_shared_experts_fusion: return disable_reason = None diff --git a/python/sglang/srt/models/glm_image_vl.py b/python/sglang/srt/models/glm_image_vl.py index 7402aa80b..6c13290d2 100644 --- a/python/sglang/srt/models/glm_image_vl.py +++ b/python/sglang/srt/models/glm_image_vl.py @@ -57,7 +57,7 @@ from sglang.srt.models.qwen2 import Qwen2MLP as GlmImageTextMLP from sglang.srt.models.qwen3_vl import Qwen3_VisionMLP as GlmImageVisionMLP from sglang.srt.models.utils import compute_cu_seqlens_from_grid_numpy from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_mm, get_parallel from sglang.srt.utils import add_prefix, is_npu logger = logging.getLogger(__name__) @@ -1018,7 +1018,7 @@ class GlmImageForConditionalGeneration(nn.Module): self.vision_config = config.vision_config self.vq_config = config.vq_config self.text_config = config.text_config - self.use_data_parallel = get_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_mm().mm_enable_dp_encoder # Bridge rope_parameters -> rope_scaling so Glm4Model can pick it up if hasattr(self.text_config, "rope_parameters") and not getattr( diff --git a/python/sglang/srt/models/glm_ocr.py b/python/sglang/srt/models/glm_ocr.py index e696b7c01..dfd00f31b 100644 --- a/python/sglang/srt/models/glm_ocr.py +++ b/python/sglang/srt/models/glm_ocr.py @@ -54,7 +54,7 @@ from sglang.srt.models.glm4v import ( Glm4vVisionModel, Glm4vVisionPatchEmbed, ) -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_mm from sglang.srt.utils import add_prefix from sglang.srt.utils.hf_transformers_utils import get_processor @@ -282,7 +282,7 @@ class GlmOcrForConditionalGeneration(Glm4vForConditionalGeneration): self.pp_group = get_pp_group() self.config = config - self.use_data_parallel = get_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_mm().mm_enable_dp_encoder self.visual = GlmOcrVisionModel( vision_config=config.vision_config, text_config=config.text_config, diff --git a/python/sglang/srt/models/glm_ocr_nextn.py b/python/sglang/srt/models/glm_ocr_nextn.py index 07a2bb245..a4d0566b4 100644 --- a/python/sglang/srt/models/glm_ocr_nextn.py +++ b/python/sglang/srt/models/glm_ocr_nextn.py @@ -33,7 +33,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.models.glm4 import Glm4DecoderLayer from sglang.srt.models.glm_ocr import GlmOcrForConditionalGeneration -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args from sglang.srt.utils import add_prefix logger = logging.getLogger(__name__) @@ -139,7 +139,7 @@ class GlmOcrForConditionalGenerationNextN(GlmOcrForConditionalGeneration): self.logits_processor = LogitsProcessor(config) self.num_fused_shared_experts = ( - 0 if get_server_args().disable_shared_experts_fusion else 1 + 0 if get_exec().moe.disable_shared_experts_fusion else 1 ) @torch.no_grad() diff --git a/python/sglang/srt/models/gpt_oss.py b/python/sglang/srt/models/gpt_oss.py index 1fe767cdb..7d99e74cf 100644 --- a/python/sglang/srt/models/gpt_oss.py +++ b/python/sglang/srt/models/gpt_oss.py @@ -69,6 +69,7 @@ from sglang.srt.models.utils import ( enable_fused_set_kv_buffer, ) from sglang.srt.runtime_context import ( + get_exec, get_forward, get_parallel, get_server_args, @@ -230,7 +231,7 @@ class GptOssSparseMoeBlock(nn.Module): self.experts = experts_type( num_experts=config.num_local_experts - + get_server_args().ep_num_redundant_experts, + + get_exec().moe.ep_num_redundant_experts, top_k=config.num_experts_per_tok, layer_id=layer_id, hidden_size=config.hidden_size, diff --git a/python/sglang/srt/models/inkling.py b/python/sglang/srt/models/inkling.py index bbaa4ca22..f3f5a610c 100644 --- a/python/sglang/srt/models/inkling.py +++ b/python/sglang/srt/models/inkling.py @@ -72,7 +72,16 @@ from sglang.srt.models.inkling_common.util import ( trtllm_bf16_weight_prep_enabled, use_inkling_shared_fused_moe, ) -from sglang.srt.runtime_context import get_model, get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_disagg, + get_exec, + get_memory, + get_mm, + get_model, + get_parallel, + get_schedule, + get_server_args, +) from sglang.srt.utils import add_prefix, is_cuda, make_layers logger = logging.getLogger(__name__) @@ -218,7 +227,7 @@ class InklingDecoderLayer(nn.Module): # cache (configs/inkling.py stream_dim) shard with them. The layer # all-gathers back to [T, H] after each sconv, before the residual add. self.attn_tp_group = get_parallel().attn_tp_group - self.scattered_sconv = get_server_args().enable_scattered_sconv + self.scattered_sconv = get_exec().comm.enable_scattered_sconv sconv_hidden = config.hidden_size if self.scattered_sconv: assert config.use_sconv, "--enable-scattered-sconv requires use_sconv" @@ -1010,9 +1019,9 @@ class InklingForConditionalGeneration(nn.Module): server_args = get_server_args() assert envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.get() - if server_args.disaggregation_mode != "decode": - assert not server_args.disable_radix_cache - assert not server_args.disable_hybrid_swa_memory + if get_disagg().disaggregation_mode != "decode": + assert not get_memory().disable_radix_cache + assert not get_schedule().disable_hybrid_swa_memory assert server_args.enable_mamba_extra_buffer() from types import SimpleNamespace @@ -1022,7 +1031,7 @@ class InklingForConditionalGeneration(nn.Module): ) inkling_quant_config = get_quantization_config( - SimpleNamespace(hf_config=self.config, model_path=server_args.model_path) + SimpleNamespace(hf_config=self.config, model_path=get_model().model_path) ) if inkling_quant_config is not None: quant_config = inkling_quant_config @@ -1039,7 +1048,7 @@ class InklingForConditionalGeneration(nn.Module): # checkpoint served text-only must not allocate/load the towers (wasted # GPU memory / avoidable startup OOM). The mm dispatch (forward) and the # weight loader already skip audio./visual. when these are None. - build_multimodal = bool(server_args.enable_multimodal) + build_multimodal = bool(get_mm().enable_multimodal) self.audio = ( InklingAudio(self.config.audio_config) if build_multimodal and self.config.audio_config.decoder_dmodel is not None diff --git a/python/sglang/srt/models/inkling_common/attn.py b/python/sglang/srt/models/inkling_common/attn.py index 8f4657e60..4ad23d0f8 100644 --- a/python/sglang/srt/models/inkling_common/attn.py +++ b/python/sglang/srt/models/inkling_common/attn.py @@ -29,7 +29,7 @@ from sglang.srt.models.inkling_common.kernels.comm import ( from sglang.srt.models.inkling_common.norm import RMSNorm from sglang.srt.models.inkling_common.sconv import SconvType, ShortConvolution from sglang.srt.models.utils import apply_qk_norm -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args from sglang.srt.utils import add_prefix, get_current_device_stream_fast try: @@ -296,7 +296,7 @@ class InklingAttention(nn.Module): ) # --enable-scattered-sconv: the output reduction becomes a hidden-dim # reduce-scatter (the consumer attn_sconv runs on the [T, H/P] shard). - self.scattered_sconv = get_server_args().enable_scattered_sconv + self.scattered_sconv = get_exec().comm.enable_scattered_sconv if is_local: self.rel_extent = local_extent diff --git a/python/sglang/srt/models/inkling_common/dense_mlp.py b/python/sglang/srt/models/inkling_common/dense_mlp.py index 6d775a333..82954319d 100644 --- a/python/sglang/srt/models/inkling_common/dense_mlp.py +++ b/python/sglang/srt/models/inkling_common/dense_mlp.py @@ -18,7 +18,7 @@ from sglang.srt.models.inkling_common.util import ( lora_compatible_layout_enabled, ) from sglang.srt.models.llama import LlamaMLP -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec, get_server_args logger = logging.getLogger(__name__) @@ -123,7 +123,7 @@ class InklingDenseMLP(LlamaMLP): fused = fused and not lora_compatible_layout_enabled() self.layer_id = layer_id self.act_fn = InklingSwiglu(interleaved=fused) - self.scattered_sconv = get_server_args().enable_scattered_sconv + self.scattered_sconv = get_exec().comm.enable_scattered_sconv def forward( self, diff --git a/python/sglang/srt/models/inkling_common/kernels/comm.py b/python/sglang/srt/models/inkling_common/kernels/comm.py index 0ddd38f0b..876baca84 100644 --- a/python/sglang/srt/models/inkling_common/kernels/comm.py +++ b/python/sglang/srt/models/inkling_common/kernels/comm.py @@ -7,7 +7,7 @@ import msgspec import torch from sglang.srt.environ import envs -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec from sglang.srt.utils import is_cuda if TYPE_CHECKING: @@ -251,7 +251,7 @@ def ar_sconv_norm_fusable( and envs.SGLANG_OPT_USE_INKLING_FUSED_AR_SCONV_NORM.get() ): return False - if get_server_args().enable_scattered_sconv: + if get_exec().comm.enable_scattered_sconv: # The decode {AR -> sconv -> norm} fusion is full-width; under scattered # sconv the output sconvs are hidden-sharded, so it does not apply. return False @@ -396,7 +396,7 @@ def get_ar_buffer( envs.SGLANG_OPT_USE_INKLING_CUSTOM_AR.get() # Scattered sconv replaces the AR with reduce_scatter_hidden, which # stages from comm.buffer[:n] -- never hand out the v4 region there. - and not get_server_args().enable_scattered_sconv + and not get_exec().comm.enable_scattered_sconv ): res = _get_inkling_ar_resources(comm) if ( @@ -696,7 +696,7 @@ def scattered_ar_sconv_fusable( if not is_cuda(): return False if not ( - get_server_args().enable_scattered_sconv + get_exec().comm.enable_scattered_sconv and envs.SGLANG_OPT_USE_INKLING_CUSTOM_AR.get() and envs.SGLANG_OPT_USE_INKLING_FUSED_AR_SCONV.get() ): @@ -1033,7 +1033,7 @@ def fullwidth_ar_sconv_fusable( if not is_cuda(): return False if not ( - not get_server_args().enable_scattered_sconv + not get_exec().comm.enable_scattered_sconv and envs.SGLANG_OPT_USE_INKLING_CUSTOM_AR.get() and envs.SGLANG_OPT_USE_INKLING_FUSED_AR_SCONV.get() ): diff --git a/python/sglang/srt/models/inkling_common/moe.py b/python/sglang/srt/models/inkling_common/moe.py index 244476357..0740f7fd8 100644 --- a/python/sglang/srt/models/inkling_common/moe.py +++ b/python/sglang/srt/models/inkling_common/moe.py @@ -59,7 +59,7 @@ from sglang.srt.models.inkling_common.util import ( lora_compatible_layout_enabled, use_inkling_shared_fused_moe, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer from sglang.srt.utils import add_prefix, is_cuda, is_hip @@ -891,9 +891,8 @@ class InklingMoE(nn.Module): ) # --enable-scattered-sconv: the output reduction becomes a hidden-dim # reduce-scatter (the consumer mlp_sconv runs on the [T, H/P] shard). - from sglang.srt.runtime_context import get_server_args - self.scattered_sconv = get_server_args().enable_scattered_sconv + self.scattered_sconv = get_exec().comm.enable_scattered_sconv # Fold the shared-expert partials into the custom AR kernels (or their # stage-in copies) instead of a separate torch.add per MoE layer. self._fused_ar_shared = envs.SGLANG_OPT_USE_INKLING_FUSED_AR_SHARED.get() diff --git a/python/sglang/srt/models/inkling_common/sconv.py b/python/sglang/srt/models/inkling_common/sconv.py index a6877fc0c..ae9a2cc62 100644 --- a/python/sglang/srt/models/inkling_common/sconv.py +++ b/python/sglang/srt/models/inkling_common/sconv.py @@ -18,7 +18,7 @@ from sglang.srt.models.inkling_common.kernels.sconv import ( save_intermediate_conv_windows, update_sconv_cache, ) -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args from sglang.srt.utils import is_cuda, set_weight_attrs @@ -243,7 +243,7 @@ class ShortConvolution(nn.Module): crossed = track_step = None if do_tracking: - mamba_track_interval = get_server_args().mamba_track_interval + mamba_track_interval = get_exec().mamba.mamba_track_interval pre_seqlen = forward_batch.seq_lens[:batch_size] - draft_token_num post_seqlen = pre_seqlen + num_accept_tokens crossed = (pre_seqlen // mamba_track_interval) != ( diff --git a/python/sglang/srt/models/inkling_common/util.py b/python/sglang/srt/models/inkling_common/util.py index 1d8bbea68..19c27da0e 100644 --- a/python/sglang/srt/models/inkling_common/util.py +++ b/python/sglang/srt/models/inkling_common/util.py @@ -9,12 +9,12 @@ from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.unquant import UnquantizedFusedMoEMethod -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_lora def lora_compatible_layout_enabled() -> bool: """Use the contiguous ``[gate || up]`` layout required by LoRA slicing.""" - return get_server_args().enable_lora + return get_lora().enable_lora def use_inkling_shared_fused_moe( diff --git a/python/sglang/srt/models/internvl.py b/python/sglang/srt/models/internvl.py index 952d21b39..e04616689 100644 --- a/python/sglang/srt/models/internvl.py +++ b/python/sglang/srt/models/internvl.py @@ -46,7 +46,7 @@ from sglang.srt.multimodal.internvl_vit_cuda_graph_runner import ( InternViTCudaGraphRunner, ) from sglang.srt.multimodal.mm_utils import run_dp_sharded_vision_model -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_mm, get_parallel from sglang.srt.utils import is_cuda from sglang.utils import logger @@ -520,7 +520,7 @@ class InternVLChatModel(nn.Module): ) -> None: super().__init__() self.config = config - self.use_data_parallel = get_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_mm().mm_enable_dp_encoder self.quant_config = quant_config vision_utils.update_vit_attn_dummy_heads_config(self.config) image_size = config.force_image_size or config.vision_config.image_size diff --git a/python/sglang/srt/models/kimi_k25.py b/python/sglang/srt/models/kimi_k25.py index 47e49c8d3..1e920ace0 100644 --- a/python/sglang/srt/models/kimi_k25.py +++ b/python/sglang/srt/models/kimi_k25.py @@ -39,7 +39,7 @@ from sglang.srt.multimodal.mm_utils import ( materialize_multimodal_features, run_dp_sharded_mrope_vision_model, ) -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_mm, get_parallel, get_server_args from sglang.srt.utils import add_prefix, is_npu logger = logging.getLogger(__name__) @@ -659,7 +659,7 @@ class KimiK25ForConditionalGeneration(nn.Module): super().__init__() self.config = config self.quant_config = quant_config - self.use_data_parallel = get_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_mm().mm_enable_dp_encoder # Create vision tower self.vision_tower = MoonViT3dPretrainedModel( config.vision_config, diff --git a/python/sglang/srt/models/kimi_vl.py b/python/sglang/srt/models/kimi_vl.py index 12c1171b4..7d9f70ea7 100644 --- a/python/sglang/srt/models/kimi_vl.py +++ b/python/sglang/srt/models/kimi_vl.py @@ -74,7 +74,7 @@ from sglang.srt.model_loader.weight_utils import ( from sglang.srt.models.deepseek_v2 import DeepseekV2ForCausalLM from sglang.srt.models.kimi_vl_moonvit import MoonVitPretrainedModel from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_mm from sglang.srt.utils import add_prefix logger = logging.getLogger(__name__) @@ -126,7 +126,7 @@ class KimiVLForConditionalGeneration(nn.Module): self.config = config assert isinstance(config.vision_config, MoonViTConfig) - self.use_data_parallel = get_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_mm().mm_enable_dp_encoder self.vision_tower = MoonVitPretrainedModel( config.vision_config, prefix=add_prefix("vision_tower", prefix), diff --git a/python/sglang/srt/models/laguna.py b/python/sglang/srt/models/laguna.py index 0373fd0c1..03377b4b4 100644 --- a/python/sglang/srt/models/laguna.py +++ b/python/sglang/srt/models/laguna.py @@ -53,7 +53,12 @@ from sglang.srt.layers.vocab_parallel_embedding import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.utils import apply_qk_norm -from sglang.srt.runtime_context import get_forward, get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_exec, + get_forward, + get_parallel, + get_server_args, +) from sglang.srt.utils import LazyValue, add_prefix, make_layers logger = logging.getLogger(__name__) @@ -155,7 +160,7 @@ class LagunaMoE(nn.Module): self.gate = LagunaMoEGate(config, prefix=add_prefix("gate", prefix)) self.experts = get_moe_impl_class(quant_config)( - num_experts=config.num_experts + get_server_args().ep_num_redundant_experts, + num_experts=config.num_experts + get_exec().moe.ep_num_redundant_experts, top_k=config.num_experts_per_tok, layer_id=layer_id, hidden_size=config.hidden_size, diff --git a/python/sglang/srt/models/llada2.py b/python/sglang/srt/models/llada2.py index a6471de6b..8a184db53 100644 --- a/python/sglang/srt/models/llada2.py +++ b/python/sglang/srt/models/llada2.py @@ -77,6 +77,7 @@ from sglang.srt.models.utils import ( enable_fused_set_kv_buffer, ) from sglang.srt.runtime_context import ( + get_exec, get_forward, get_parallel, get_server_args, @@ -231,7 +232,7 @@ class LLaDA2MoeSparseMoeBlock(nn.Module): self.router_dtype = torch.bfloat16 # TODO global_server_args.ep_num_redundant_experts is used for eplb, not supported now - assert get_server_args().ep_num_redundant_experts == 0 + assert get_exec().moe.ep_num_redundant_experts == 0 # check group topk self.num_expert_group = getattr(config, "n_group", 0) self.topk_group = getattr(config, "topk_group", 0) @@ -245,9 +246,7 @@ class LLaDA2MoeSparseMoeBlock(nn.Module): self.num_expert_group = self.topk_group = None self.use_grouped_topk = False - self.num_experts = ( - config.num_experts + get_server_args().ep_num_redundant_experts - ) + self.num_experts = config.num_experts + get_exec().moe.ep_num_redundant_experts self.gate = LLaDA2MoeGate( config=config, diff --git a/python/sglang/srt/models/llama_eagle3.py b/python/sglang/srt/models/llama_eagle3.py index be8e8a7c3..89f2ac4d9 100644 --- a/python/sglang/srt/models/llama_eagle3.py +++ b/python/sglang/srt/models/llama_eagle3.py @@ -13,6 +13,7 @@ See the License for the specific language governing permissions and limitations under the License. """ +from sglang.srt.runtime_context import get_spec from sglang.srt.utils import add_prefix # Adapted from @@ -38,7 +39,6 @@ from sglang.srt.layers.vocab_parallel_embedding import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.llama import LlamaDecoderLayer, LlamaForCausalLM, LlamaMLP -from sglang.srt.runtime_context import get_server_args class LlamaDecoderLayer(LlamaDecoderLayer): @@ -275,7 +275,7 @@ class LlamaForCausalLMEagle3(LlamaForCausalLM): # Cache draft SWA size from server args once; consumed both by the post-init # attention patch below and by `get_attention_sliding_window_size` later. self._draft_window_size: Optional[int] = ( - get_server_args().speculative_draft_window_size + get_spec().speculative_draft_window_size ) self.model = LlamaModel( diff --git a/python/sglang/srt/models/mellum.py b/python/sglang/srt/models/mellum.py index ff62c558d..c792e37f8 100644 --- a/python/sglang/srt/models/mellum.py +++ b/python/sglang/srt/models/mellum.py @@ -51,7 +51,7 @@ from sglang.srt.models.utils import ( create_fused_set_kv_buffer_arg, enable_fused_set_kv_buffer, ) -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args from sglang.srt.utils import add_prefix, is_cuda _is_cuda = is_cuda() @@ -231,7 +231,7 @@ class MellumAttention(Qwen3MoeAttention): _yarn_factor = self._yarn_params["factor"] self.use_fused_qk_norm_rope = ( - get_server_args().enable_fused_qk_norm_rope + get_exec().kernel.enable_fused_qk_norm_rope and self.compatible_with_fused_qk_norm_rope and _is_cuda and can_use_fused_qk_norm_rope( diff --git a/python/sglang/srt/models/mimo_audio.py b/python/sglang/srt/models/mimo_audio.py index 650c3309f..6e738b3ad 100644 --- a/python/sglang/srt/models/mimo_audio.py +++ b/python/sglang/srt/models/mimo_audio.py @@ -22,7 +22,7 @@ from transformers.models.qwen2.modeling_qwen2 import Qwen2Model from sglang.srt.layers.attention.vision import VisionAttention from sglang.srt.layers.quantization.base_config import QuantizationConfig -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_model logger = logging.getLogger(__name__) @@ -1255,7 +1255,7 @@ class AudioEncoderMixin: else: raise ValueError(f"Invalid projection layers: {config.projection_layers}") - model_path = get_server_args().model_path + model_path = get_model().model_path if not os.path.isdir(model_path): from huggingface_hub import snapshot_download diff --git a/python/sglang/srt/models/mimo_v2.py b/python/sglang/srt/models/mimo_v2.py index 781f3093e..a16afd50f 100644 --- a/python/sglang/srt/models/mimo_v2.py +++ b/python/sglang/srt/models/mimo_v2.py @@ -79,6 +79,7 @@ from sglang.srt.model_loader.weight_utils import ( from sglang.srt.models.mimo_audio import AudioEncoderMixin, MiMoAudioEncoderConfig from sglang.srt.models.mimo_vl import MiMoVisionTransformer, MiMoVLVisionConfig from sglang.srt.runtime_context import ( + get_exec, get_forward, get_parallel, get_server_args, @@ -413,7 +414,7 @@ class MiMoV2MoE(nn.Module): experts_type = get_moe_impl_class(quant_config) self.experts = experts_type( num_experts=config.n_routed_experts - + get_server_args().ep_num_redundant_experts, + + get_exec().moe.ep_num_redundant_experts, top_k=config.num_experts_per_tok, hidden_size=config.hidden_size, intermediate_size=config.moe_intermediate_size, @@ -448,7 +449,7 @@ class MiMoV2MoE(nn.Module): # TODO: we will support tp < ep in the future self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.n_routed_experts + get_server_args().ep_num_redundant_experts + config.n_routed_experts + get_exec().moe.ep_num_redundant_experts ) self.renormalize = config.norm_topk_prob self.topk_group = config.topk_group diff --git a/python/sglang/srt/models/mimo_vl.py b/python/sglang/srt/models/mimo_vl.py index b36bc9616..3f4c7267b 100644 --- a/python/sglang/srt/models/mimo_vl.py +++ b/python/sglang/srt/models/mimo_vl.py @@ -22,7 +22,7 @@ from sglang.srt.layers.attention.vision import ( from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.models.qwen2_5_vl import Qwen2_5_VisionPatchMerger, Qwen2_5_VLMLP -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_mm, get_server_args from sglang.srt.utils import add_prefix @@ -258,7 +258,7 @@ class MiMoVisionTransformer(nn.Module): self.fullatt_block_indexes = vision_config.fullatt_block_indexes self.window_size = vision_config.window_size self.patch_size = vision_config.patch_size - self.use_data_parallel = self.server_args.mm_enable_dp_encoder + self.use_data_parallel = get_mm().mm_enable_dp_encoder mlp_hidden_size: int = vision_config.intermediate_size self.patch_embed = MiMoVisionPatchEmbed( patch_size=patch_size, diff --git a/python/sglang/srt/models/minimax_m2.py b/python/sglang/srt/models/minimax_m2.py index ea4a6b5be..3dd4e3e9f 100644 --- a/python/sglang/srt/models/minimax_m2.py +++ b/python/sglang/srt/models/minimax_m2.py @@ -80,7 +80,13 @@ from sglang.srt.model_loader.weight_utils import ( maybe_remap_kv_scale_name, narrow_padded_param_and_loaded_weight, ) -from sglang.srt.runtime_context import get_forward, get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_exec, + get_forward, + get_parallel, + get_schedule, + get_server_args, +) # get_bool_env_var is defined in sglang.srt.utils.common, not sglang.srt.distributed. # Importing from the wrong module causes this file to fail import, which prevents the @@ -426,10 +432,10 @@ class MiniMaxM2QKRMSNorm: props = torch.cuda.get_device_properties(device) # probe the maximum tokens for one prefill server_args = get_server_args() - max_tokens = server_args.chunked_prefill_size + max_tokens = get_schedule().chunked_prefill_size if max_tokens is None: max_tokens = server_args.model_config.context_len - max_tokens = max(max_tokens, server_args.max_prefill_tokens) + max_tokens = max(max_tokens, get_schedule().max_prefill_tokens) logger.info(f"[AR] Using CustomAllReduceV2 for MiniMaxM2 with {max_tokens = }") ALIGN = 512 # typically, this should not exceed 1M, since max_tokens is usually less than 16384 @@ -513,7 +519,7 @@ class MiniMaxM2MoE(nn.Module): self.experts = get_moe_impl_class(quant_config)( num_experts=config.num_local_experts - + get_server_args().ep_num_redundant_experts, + + get_exec().moe.ep_num_redundant_experts, top_k=config.num_experts_per_tok, hidden_size=config.hidden_size, intermediate_size=config.intermediate_size, diff --git a/python/sglang/srt/models/minimax_m3.py b/python/sglang/srt/models/minimax_m3.py index a32eace30..4b9017b48 100644 --- a/python/sglang/srt/models/minimax_m3.py +++ b/python/sglang/srt/models/minimax_m3.py @@ -80,7 +80,7 @@ from sglang.srt.model_loader.weight_utils import ( ) from sglang.srt.models.minimax_m2 import MiniMaxM2RMSNormTP from sglang.srt.models.utils import WeightsMapper -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args from sglang.srt.utils import ( add_prefix, get_device_sm, @@ -288,7 +288,7 @@ class MiniMaxM3MoE(nn.Module): self.n_shared_experts = getattr(config, "n_shared_experts", None) self.num_fused_shared_experts = ( 0 - if get_server_args().disable_shared_experts_fusion + if get_exec().moe.disable_shared_experts_fusion else config.n_shared_experts ) @@ -312,7 +312,7 @@ class MiniMaxM3MoE(nn.Module): self.experts = get_moe_impl_class(quant_config)( num_experts=config.num_local_experts + self.num_fused_shared_experts - + get_server_args().ep_num_redundant_experts, + + get_exec().moe.ep_num_redundant_experts, num_fused_shared_experts=self.num_fused_shared_experts, top_k=config.num_experts_per_tok + self.num_fused_shared_experts, hidden_size=config.hidden_size, @@ -1466,7 +1466,7 @@ class MiniMaxM3SparseForCausalLM(nn.Module): return self.model.get_input_embeddings() def determine_num_fused_shared_experts(self): - if get_server_args().disable_shared_experts_fusion: + if get_exec().moe.disable_shared_experts_fusion: return disable_reason = None diff --git a/python/sglang/srt/models/minimax_m3_vl.py b/python/sglang/srt/models/minimax_m3_vl.py index c2a7beaf5..bddc034be 100644 --- a/python/sglang/srt/models/minimax_m3_vl.py +++ b/python/sglang/srt/models/minimax_m3_vl.py @@ -43,7 +43,7 @@ from sglang.srt.models.minimax_vl_common import ( merge_vit_qkv_weights, ) from sglang.srt.models.utils import WeightsMapper -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_mm, get_parallel, get_server_args from sglang.srt.utils import add_prefix, get_device_sm, is_cuda, log_info_on_rank0 from sglang.srt.utils.hf_transformers_utils import get_rope_config @@ -75,7 +75,7 @@ class MiniMaxM3SparseForConditionalGeneration(nn.Module): self.quant_config = quant_config self.pp_group = get_pp_group() - self.use_data_parallel = get_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_mm().mm_enable_dp_encoder self.num_fused_shared_experts = 0 self._determine_num_fused_shared_experts() @@ -135,7 +135,7 @@ class MiniMaxM3SparseForConditionalGeneration(nn.Module): def _determine_num_fused_shared_experts(self) -> None: text_config = self.config.text_config server_args = get_server_args() - if server_args.disable_shared_experts_fusion: + if get_exec().moe.disable_shared_experts_fusion: return disable_reason = None diff --git a/python/sglang/srt/models/minimax_vl_common.py b/python/sglang/srt/models/minimax_vl_common.py index 0987165b7..e5a24013f 100644 --- a/python/sglang/srt/models/minimax_vl_common.py +++ b/python/sglang/srt/models/minimax_vl_common.py @@ -27,7 +27,7 @@ from sglang.srt.layers.rotary_embedding.utils import rotate_half from sglang.srt.managers.schedule_batch import MultimodalDataItem from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_mm, get_parallel from sglang.srt.utils import add_prefix, get_compiler_backend, round_up logger = logging.getLogger(__name__) @@ -413,7 +413,7 @@ class MiniMaxVLVisionTransformer(nn.Module): workspace_buffer: Optional[torch.Tensor] = None if ( - get_server_args().mm_attention_backend == "flashinfer_cudnn" + get_mm().mm_attention_backend == "flashinfer_cudnn" and torch.cuda.is_available() ): workspace_buffer = torch.empty( @@ -679,7 +679,7 @@ class MiniMaxVLVisionTransformer(nn.Module): max_seqlen: Optional[int] = None sequence_lengths: Optional[torch.Tensor] = None encoder_cu_seq_len = cu_seq_len - if get_server_args().mm_attention_backend == "flashinfer_cudnn": + if get_mm().mm_attention_backend == "flashinfer_cudnn": ( encoder_cu_seq_len, sequence_lengths, @@ -691,7 +691,7 @@ class MiniMaxVLVisionTransformer(nn.Module): device=hidden_states.device, packed_indptrs=( encoder_cu_seq_len - if get_server_args().mm_attention_backend == "flashinfer_cudnn" + if get_mm().mm_attention_backend == "flashinfer_cudnn" else None ), sequence_lengths=sequence_lengths, @@ -723,7 +723,7 @@ class MiniMaxVLVisionModel(nn.Module): self.config = config self.quant_config = quant_config - self.use_data_parallel = get_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_mm().mm_enable_dp_encoder self.vision_config = config self.vision_model = MiniMaxVLVisionTransformer( diff --git a/python/sglang/srt/models/mllama4.py b/python/sglang/srt/models/mllama4.py index 68fc8f6db..28ddcdc87 100644 --- a/python/sglang/srt/models/mllama4.py +++ b/python/sglang/srt/models/mllama4.py @@ -33,7 +33,7 @@ from sglang.srt.managers.schedule_batch import ( MultimodalInputs, ) from sglang.srt.model_executor.forward_batch_info import ForwardBatch -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_mm from sglang.srt.utils import is_cpu _is_cpu = is_cpu() @@ -476,9 +476,7 @@ class Llama4ForConditionalGeneration(nn.Module): "Please not that this warning might be inaccurate if the weights haven't been fully downloaded" ) - self.has_vision = ( - self.has_vision_weights and get_server_args().enable_multimodal - ) + self.has_vision = self.has_vision_weights and get_mm().enable_multimodal if self.has_vision: # TODO: make this more general diff --git a/python/sglang/srt/models/moss_vl.py b/python/sglang/srt/models/moss_vl.py index 9b04905f5..55a57111e 100644 --- a/python/sglang/srt/models/moss_vl.py +++ b/python/sglang/srt/models/moss_vl.py @@ -48,7 +48,7 @@ from sglang.srt.managers.schedule_batch import MultimodalInputs from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_loader.weight_utils import default_weight_loader -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.utils import add_prefix logger = logging.getLogger(__name__) @@ -1002,7 +1002,7 @@ class MossVLSelfAttentionDecoderLayer(nn.Module): override_orig_dtype=torch.float32, fp32_residual=True, ) - if get_server_args().rl_on_policy_target is not None + if get_exec().deterministic.rl_on_policy_target is not None else {} ) self.input_layernorm = RMSNorm( diff --git a/python/sglang/srt/models/nemotron_h.py b/python/sglang/srt/models/nemotron_h.py index 71dc6005c..bcb36a496 100644 --- a/python/sglang/srt/models/nemotron_h.py +++ b/python/sglang/srt/models/nemotron_h.py @@ -89,7 +89,12 @@ from sglang.srt.models.nemotron_h_utils import ( pad_to_original_num_tokens, ) from sglang.srt.models.utils import WeightsMapper -from sglang.srt.runtime_context import get_forward, get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_exec, + get_forward, + get_parallel, + get_server_args, +) from sglang.srt.utils import ( add_prefix, get_current_device_stream_fast, @@ -200,7 +205,7 @@ class NemotronHMoE(nn.Module): self.experts = get_moe_impl_class(quant_config)( num_experts=config.n_routed_experts - + get_server_args().ep_num_redundant_experts, + + get_exec().moe.ep_num_redundant_experts, top_k=config.num_experts_per_tok, hidden_size=self.moe_hidden_size, intermediate_size=config.moe_intermediate_size, diff --git a/python/sglang/srt/models/qwen2.py b/python/sglang/srt/models/qwen2.py index 14e2601a7..0fb1592ce 100644 --- a/python/sglang/srt/models/qwen2.py +++ b/python/sglang/srt/models/qwen2.py @@ -50,7 +50,7 @@ from sglang.srt.model_loader.weight_utils import ( kv_cache_scales_loader, ) from sglang.srt.platforms import current_platform -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.utils import add_prefix, make_layers from sglang.srt.utils.hf_transformers_utils import get_rope_config @@ -96,7 +96,7 @@ class Qwen2MLP(nn.Module): x: torch.Tensor, forward_batch: ForwardBatch = None, ) -> torch.Tensor: - if get_server_args().rl_on_policy_target is not None: + if get_exec().deterministic.rl_on_policy_target is not None: x = x.bfloat16() gate_up, _ = self.gate_up_proj(x) @@ -330,7 +330,7 @@ class Qwen2Model(nn.Module): prefix=add_prefix("embed_tokens", prefix), params_dtype=( torch.float32 - if get_server_args().rl_on_policy_target is not None + if get_exec().deterministic.rl_on_policy_target is not None else None ), ) @@ -366,7 +366,7 @@ class Qwen2Model(nn.Module): override_orig_dtype=torch.float32, fp32_residual=True, ) - if get_server_args().rl_on_policy_target is not None + if get_exec().deterministic.rl_on_policy_target is not None else {} ) self.norm = RMSNorm( diff --git a/python/sglang/srt/models/qwen2_5_vl.py b/python/sglang/srt/models/qwen2_5_vl.py index bba029f69..34fca18ce 100644 --- a/python/sglang/srt/models/qwen2_5_vl.py +++ b/python/sglang/srt/models/qwen2_5_vl.py @@ -76,7 +76,7 @@ from sglang.srt.models.qwen2 import Qwen2Model from sglang.srt.models.utils import RotaryPosMixin, WeightsMapper, permute_inv from sglang.srt.multimodal.mm_utils import run_dp_sharded_mrope_vision_model from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_mm, get_parallel from sglang.srt.utils import add_prefix, is_cpu, is_cuda, is_npu _is_cuda = is_cuda() @@ -619,7 +619,7 @@ class Qwen2_5_VLForConditionalGeneration(nn.Module): self.pp_group = get_pp_group() self.config = config - self.use_data_parallel = get_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_mm().mm_enable_dp_encoder if not self.config.encoder_only: self.model = Qwen2Model( diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index fb879a689..942b38d30 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -92,7 +92,12 @@ from sglang.srt.model_executor.cuda_graph_config import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_loader.weight_utils import default_weight_loader -from sglang.srt.runtime_context import get_forward, get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_exec, + get_forward, + get_parallel, + get_server_args, +) from sglang.srt.utils import ( add_prefix, cpu_has_amx_support, @@ -147,7 +152,7 @@ def can_fuse_shared_expert( Caller must still gate on the model/backend support flag. """ if ( - get_server_args().disable_shared_experts_fusion is True + get_exec().moe.disable_shared_experts_fusion is True or getattr(config, "shared_expert_intermediate_size", 0) <= 0 or config.shared_expert_intermediate_size != config.moe_intermediate_size or get_moe_a2a_backend().is_deepep() @@ -286,10 +291,10 @@ class Qwen2MoeSparseMoeBlock(nn.Module): else config.num_experts_per_tok + self.num_fused_shared_experts ), num_experts=( - config.num_experts + get_server_args().ep_num_redundant_experts + config.num_experts + get_exec().moe.ep_num_redundant_experts if not self.enable_shared_expert_fusion else config.num_experts - + get_server_args().ep_num_redundant_experts + + get_exec().moe.ep_num_redundant_experts + self.num_fused_shared_experts ), hidden_size=config.hidden_size, @@ -348,7 +353,7 @@ class Qwen2MoeSparseMoeBlock(nn.Module): # TODO: we will support tp < ep in the future self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.num_experts + get_server_args().ep_num_redundant_experts + config.num_experts + get_exec().moe.ep_num_redundant_experts ) self.top_k = config.num_experts_per_tok self.is_nextn = is_nextn diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py index daaca2276..2f3000d29 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -33,7 +33,12 @@ from sglang.srt.model_loader.weight_utils import ( from sglang.srt.models.qwen2 import Qwen2MLP as Qwen3MLP from sglang.srt.models.qwen2 import Qwen2Model from sglang.srt.models.utils import apply_qk_norm -from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream +from sglang.srt.runtime_context import ( + get_exec, + get_parallel, + get_server_args, + get_stream, +) from sglang.srt.utils import add_prefix, get_bool_env_var, is_cuda, is_hip, is_npu Qwen3Config = None @@ -111,7 +116,7 @@ class Qwen3Attention(nn.Module): weight_dtype=torch.float32, cast_x_before_out_mul=True, ) - if get_server_args().rl_on_policy_target is not None + if get_exec().deterministic.rl_on_policy_target is not None else {} ) self.q_norm = RMSNorm(self.head_dim, eps=rms_norm_eps, **norm_kwargs) @@ -271,14 +276,14 @@ class Qwen3Attention(nn.Module): hidden_states: torch.Tensor, forward_batch: ForwardBatch, ) -> torch.Tensor: - if get_server_args().rl_on_policy_target is not None: + if get_exec().deterministic.rl_on_policy_target is not None: hidden_states = hidden_states.bfloat16() save_kv_cache = True use_aiter_fused = ( self.use_fused_qk_norm_mrope and forward_batch.forward_mode.is_decode() - and get_server_args().rl_on_policy_target is None + and get_exec().deterministic.rl_on_policy_target is None ) if use_aiter_fused: @@ -298,7 +303,7 @@ class Qwen3Attention(nn.Module): forward_batch=forward_batch, ) - if get_server_args().rl_on_policy_target is not None: + if get_exec().deterministic.rl_on_policy_target is not None: q = q.to(torch.bfloat16) k = k.to(torch.bfloat16) @@ -362,7 +367,7 @@ class Qwen3DecoderLayer(nn.Module): override_orig_dtype=torch.float32, fp32_residual=True, ) - if get_server_args().rl_on_policy_target is not None + if get_exec().deterministic.rl_on_policy_target is not None else {} ) self.input_layernorm = RMSNorm( diff --git a/python/sglang/srt/models/qwen3_5.py b/python/sglang/srt/models/qwen3_5.py index fab3d1bab..08c47b994 100644 --- a/python/sglang/srt/models/qwen3_5.py +++ b/python/sglang/srt/models/qwen3_5.py @@ -92,9 +92,9 @@ from sglang.srt.models.utils import ( fused_qk_gemma_rmsnorm_with_gate, ) from sglang.srt.runtime_context import ( + get_exec, get_forward, get_parallel, - get_server_args, get_stream, ) @@ -139,7 +139,7 @@ cached_get_processor = lru_cache(get_processor) def _disable_shared_experts_fusion() -> bool: # Resolved lazily: the global server args is not set at module import time # (e.g. when this module is imported by unit tests). - return get_server_args().disable_shared_experts_fusion + return get_exec().moe.disable_shared_experts_fusion if _is_cuda: @@ -181,7 +181,7 @@ def _enable_qwen35_fused_ar_quant() -> bool: return False if get_bool_env_var("SGLANG_DISABLE_FUSED_AR_QUANT", default="false"): return False - return bool(get_server_args().enable_aiter_allreduce_fusion) + return bool(get_exec().comm.enable_aiter_allreduce_fusion) def _linear_accepts_fp8_tuple(linear: nn.Module) -> bool: @@ -1313,7 +1313,7 @@ class Qwen3_5ForCausalLM(nn.Module): # so the model still gets the #25885 multi-streaming path. ROCm-only. if ( config.model_type == "qwen3_5_moe_text" - and not get_server_args().disable_shared_experts_fusion + and not get_exec().moe.disable_shared_experts_fusion and not can_fuse_shared_expert(config, quant_config) ): from sglang.srt.arg_groups.overrides import declare_load_time_override diff --git a/python/sglang/srt/models/qwen3_5_mtp.py b/python/sglang/srt/models/qwen3_5_mtp.py index ebda7517f..ded5ef3b8 100644 --- a/python/sglang/srt/models/qwen3_5_mtp.py +++ b/python/sglang/srt/models/qwen3_5_mtp.py @@ -34,7 +34,11 @@ from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.qwen3_5 import Qwen3_5ForCausalLM -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_model, + get_parallel, + get_spec, +) from sglang.srt.utils import add_prefix, is_npu logger = logging.getLogger(__name__) @@ -63,7 +67,7 @@ class Qwen3_5ForCausalLMMTP(nn.Module): "modelopt_mixed", ): quant_config = None - if is_npu() and get_server_args().speculative_draft_model_quantization is None: + if is_npu() and get_spec().speculative_draft_model_quantization is None: quant_config = None # Quark-quantized Qwen3.5 MXFP4 checkpoints ship the MTP module in @@ -153,7 +157,7 @@ class Qwen3_5ForCausalLMMTP(nn.Module): if ( is_npu() and self.quant_config is None - and get_server_args().quantization is not None + and get_model().quantization is not None ): # ascend mtp unquant exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True)) diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index 9524a7ba7..2f20631e3 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -73,6 +73,7 @@ from sglang.srt.models.utils import ( enable_fused_set_kv_buffer, ) from sglang.srt.runtime_context import ( + get_exec, get_forward, get_parallel, get_server_args, @@ -261,7 +262,7 @@ class Qwen3MoeSparseMoeBlock(nn.Module): ) self.experts = get_moe_impl_class(quant_config)( - num_experts=config.num_experts + get_server_args().ep_num_redundant_experts, + num_experts=config.num_experts + get_exec().moe.ep_num_redundant_experts, top_k=config.num_experts_per_tok, layer_id=layer_id, hidden_size=config.hidden_size, @@ -293,7 +294,7 @@ class Qwen3MoeSparseMoeBlock(nn.Module): # TODO: we will support tp < ep in the future self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.num_experts + get_server_args().ep_num_redundant_experts + config.num_experts + get_exec().moe.ep_num_redundant_experts ) self.top_k = config.num_experts_per_tok @@ -524,7 +525,7 @@ class Qwen3MoeAttention(nn.Module): ) and self.head_dim in (64, 128, 256) _yarn_factor, _, _, _ = compute_yarn_parameters(config) self.use_fused_qk_norm_rope = ( - get_server_args().enable_fused_qk_norm_rope + get_exec().kernel.enable_fused_qk_norm_rope and self.compatible_with_fused_qk_norm_rope and _is_cuda and can_use_fused_qk_norm_rope( diff --git a/python/sglang/srt/models/qwen3_next_mtp.py b/python/sglang/srt/models/qwen3_next_mtp.py index ef33eb9a1..3d86e5f94 100644 --- a/python/sglang/srt/models/qwen3_next_mtp.py +++ b/python/sglang/srt/models/qwen3_next_mtp.py @@ -32,7 +32,12 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.models.qwen3_next import Qwen3NextForCausalLM, Qwen3NextModel -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_model, + get_parallel, + get_server_args, + get_spec, +) from sglang.srt.utils import add_prefix, is_npu logger = logging.getLogger(__name__) @@ -51,7 +56,7 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM): config = copy.deepcopy(config) self.config = config self.tp_size = get_parallel().tp_size - if is_npu() and get_server_args().speculative_draft_model_quantization is None: + if is_npu() and get_spec().speculative_draft_model_quantization is None: quant_config = None self.quant_config = quant_config # if not set, model load will be broken in Qwen3NextForCausalLM load_weights() @@ -110,7 +115,7 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM): if ( is_npu() and self.quant_config is None - and get_server_args().quantization is not None + and get_model().quantization is not None ): # ascend mtp unquant exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True)) diff --git a/python/sglang/srt/models/qwen3_vl.py b/python/sglang/srt/models/qwen3_vl.py index 3e9d5d211..c66bcb547 100644 --- a/python/sglang/srt/models/qwen3_vl.py +++ b/python/sglang/srt/models/qwen3_vl.py @@ -73,7 +73,7 @@ from sglang.srt.multimodal.mm_utils import ( run_dp_sharded_mrope_vision_model, ) from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_mm, get_parallel, get_server_args from sglang.srt.utils import ( add_prefix, cpu_has_amx_support, @@ -329,7 +329,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin): self.num_position_embeddings = vision_config.num_position_embeddings self.num_grid_per_side = int(self.num_position_embeddings**0.5) self.num_grid = self.num_grid_per_side * self.num_grid_per_side - self.align_corners = get_server_args().enable_precise_embedding_interpolation + self.align_corners = get_exec().kernel.enable_precise_embedding_interpolation self.patch_size = vision_config.patch_size self.spatial_merge_size = vision_config.spatial_merge_size self.spatial_merge_unit = self.spatial_merge_size**2 @@ -372,7 +372,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin): ) workspace_buffer = None - if get_server_args().mm_attention_backend == "flashinfer_cudnn": + if get_mm().mm_attention_backend == "flashinfer_cudnn": if torch.cuda.is_available() and (not _is_npu): ws_device = torch.device("cuda", torch.cuda.current_device()) else: @@ -920,7 +920,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin): flashinfer_sequence_lengths = None flashinfer_max_seqlen = 0 - if get_server_args().mm_attention_backend == "flashinfer_cudnn": + if get_mm().mm_attention_backend == "flashinfer_cudnn": # real token lens (B,) real_seq_lens = token_cu_seqlens[1:] - token_cu_seqlens[:-1] flashinfer_max_seqlen = self.bucket_flashinfer_max_seqlen( @@ -1236,7 +1236,7 @@ class Qwen3VLForConditionalGeneration(nn.Module): self.pp_group = get_pp_group() self.quant_config = quant_config - self.use_data_parallel = get_server_args().mm_enable_dp_encoder + self.use_data_parallel = get_mm().mm_enable_dp_encoder self.visual = Qwen3VLMoeVisionModel( config.vision_config, diff --git a/python/sglang/srt/models/sarvam_moe.py b/python/sglang/srt/models/sarvam_moe.py index 11ad6468c..6906820e3 100644 --- a/python/sglang/srt/models/sarvam_moe.py +++ b/python/sglang/srt/models/sarvam_moe.py @@ -61,6 +61,7 @@ from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha imp DeepseekMHAForwardMixin, ) from sglang.srt.runtime_context import ( + get_exec, get_forward, get_model, get_parallel, @@ -274,7 +275,7 @@ class SarvamMoESparseMoeBlock(nn.Module): ) self.experts = get_moe_impl_class(quant_config)( - num_experts=config.num_experts + get_server_args().ep_num_redundant_experts, + num_experts=config.num_experts + get_exec().moe.ep_num_redundant_experts, top_k=config.num_experts_per_tok, hidden_size=config.hidden_size, intermediate_size=config.moe_intermediate_size, diff --git a/python/sglang/srt/models/sdar.py b/python/sglang/srt/models/sdar.py index 91946db54..7ae5d8b5d 100644 --- a/python/sglang/srt/models/sdar.py +++ b/python/sglang/srt/models/sdar.py @@ -42,6 +42,7 @@ from sglang.srt.models.utils import ( enable_fused_set_kv_buffer, ) from sglang.srt.runtime_context import ( + get_exec, get_forward, get_parallel, get_server_args, @@ -207,7 +208,7 @@ class SDARAttention(nn.Module): hidden_states: torch.Tensor, forward_batch: ForwardBatch, ): - if get_server_args().rl_on_policy_target is not None: + if get_exec().deterministic.rl_on_policy_target is not None: hidden_states = hidden_states.bfloat16() qkv, _ = self.qkv_proj(hidden_states) @@ -235,7 +236,7 @@ class SDARAttention(nn.Module): ), ) - if get_server_args().rl_on_policy_target is not None: + if get_exec().deterministic.rl_on_policy_target is not None: q = q.to(torch.bfloat16) k = k.to(torch.bfloat16) @@ -270,7 +271,7 @@ class SDARBlock(nn.Module): override_orig_dtype=torch.float32, fp32_residual=True, ) - if get_server_args().rl_on_policy_target is not None + if get_exec().deterministic.rl_on_policy_target is not None else {} ) self.input_layernorm = RMSNorm( @@ -395,7 +396,7 @@ class SDARModel(nn.Module): override_orig_dtype=torch.float32, fp32_residual=True, ) - if get_server_args().rl_on_policy_target is not None + if get_exec().deterministic.rl_on_policy_target is not None else {} ) self.norm = RMSNorm(self.embed_dim, eps=config.rms_norm_eps, **norm_kwargs) diff --git a/python/sglang/srt/models/sdar_moe.py b/python/sglang/srt/models/sdar_moe.py index b91db7be3..14c53099b 100644 --- a/python/sglang/srt/models/sdar_moe.py +++ b/python/sglang/srt/models/sdar_moe.py @@ -58,6 +58,7 @@ from sglang.srt.models.utils import ( enable_fused_set_kv_buffer, ) from sglang.srt.runtime_context import ( + get_exec, get_forward, get_parallel, get_server_args, @@ -101,7 +102,7 @@ class SDARMoeSparseMoeBlock(nn.Module): ) self.experts = get_moe_impl_class(quant_config)( - num_experts=config.num_experts + get_server_args().ep_num_redundant_experts, + num_experts=config.num_experts + get_exec().moe.ep_num_redundant_experts, top_k=config.num_experts_per_tok, layer_id=layer_id, hidden_size=config.hidden_size, @@ -123,7 +124,7 @@ class SDARMoeSparseMoeBlock(nn.Module): if get_moe_a2a_backend().is_deepep(): self.ep_size = get_parallel().moe_ep_size self.num_experts = ( - config.num_experts + get_server_args().ep_num_redundant_experts + config.num_experts + get_exec().moe.ep_num_redundant_experts ) self.top_k = config.num_experts_per_tok @@ -274,7 +275,7 @@ class SDARMoeAttention(nn.Module): hidden_states: torch.Tensor, forward_batch: ForwardBatch, ) -> torch.Tensor: - if get_server_args().rl_on_policy_target is not None: + if get_exec().deterministic.rl_on_policy_target is not None: hidden_states = hidden_states.bfloat16() qkv, _ = self.qkv_proj(hidden_states) @@ -302,7 +303,7 @@ class SDARMoeAttention(nn.Module): ), ) - if get_server_args().rl_on_policy_target is not None: + if get_exec().deterministic.rl_on_policy_target is not None: q = q.to(torch.bfloat16) k = k.to(torch.bfloat16) @@ -338,7 +339,7 @@ class SDARMoeBlock(nn.Module): override_orig_dtype=torch.float32, fp32_residual=True, ) - if get_server_args().rl_on_policy_target is not None + if get_exec().deterministic.rl_on_policy_target is not None else {} ) self.input_layernorm = RMSNorm( @@ -478,7 +479,7 @@ class SDARMoeModel(nn.Module): override_orig_dtype=torch.float32, fp32_residual=True, ) - if get_server_args().rl_on_policy_target is not None + if get_exec().deterministic.rl_on_policy_target is not None else {} ) self.norm = RMSNorm(self.embed_dim, eps=config.rms_norm_eps, **norm_kwargs) diff --git a/python/sglang/srt/models/step3p5.py b/python/sglang/srt/models/step3p5.py index 117b795bd..287df9f62 100644 --- a/python/sglang/srt/models/step3p5.py +++ b/python/sglang/srt/models/step3p5.py @@ -47,6 +47,7 @@ from sglang.srt.layers.vocab_parallel_embedding import ( from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.runtime_context import ( + get_exec, get_forward, get_parallel, get_server_args, @@ -153,7 +154,7 @@ class Step3p5MoEMLP(nn.Module): self.experts = get_moe_impl_class(quant_config)( num_experts=config.moe_num_experts - + get_server_args().ep_num_redundant_experts, + + get_exec().moe.ep_num_redundant_experts, top_k=config.moe_top_k, layer_id=layer_id, hidden_size=config.hidden_size, @@ -176,7 +177,7 @@ class Step3p5MoEMLP(nn.Module): # TODO: we will support tp < ep in the future self.ep_size = get_parallel().moe_ep_size self.moe_num_experts = ( - config.moe_num_experts + get_server_args().ep_num_redundant_experts + config.moe_num_experts + get_exec().moe.ep_num_redundant_experts ) self.top_k = config.moe_top_k @@ -676,7 +677,7 @@ class Step3p5Model(nn.Module): prefix=add_prefix("embed_tokens", prefix), params_dtype=( torch.float32 - if get_server_args().rl_on_policy_target is not None + if get_exec().deterministic.rl_on_policy_target is not None else None ), ) diff --git a/python/sglang/srt/models/transformers.py b/python/sglang/srt/models/transformers.py index e8e801477..a5584c85e 100644 --- a/python/sglang/srt/models/transformers.py +++ b/python/sglang/srt/models/transformers.py @@ -65,7 +65,7 @@ from sglang.srt.managers.schedule_batch import MultimodalDataItem, MultimodalInp from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.utils import AutoWeightsLoader, WeightsMapper -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.utils import get_device from sglang.srt.utils.common import direct_register_custom_op from sglang.srt.utils.hf_transformers_utils import get_hf_text_config @@ -352,7 +352,7 @@ class TransformersFusedMoE(nn.Module): expert_mapping: list, ) -> None: super().__init__() - num_redundant = get_server_args().ep_num_redundant_experts + num_redundant = get_exec().moe.ep_num_redundant_experts experts_cls = get_moe_impl_class(quant_config) self.experts = experts_cls( num_experts=num_experts + num_redundant, @@ -1231,7 +1231,7 @@ class MoEMixin: expert_mapping = self._get_expert_mapping(num_experts) # EPLB / EP tracking - num_redundant = get_server_args().ep_num_redundant_experts + num_redundant = get_exec().moe.ep_num_redundant_experts ep_size = get_parallel().moe_ep_size self.mlp_moe_layers: list[nn.Module] = [] diff --git a/python/sglang/srt/models/utils.py b/python/sglang/srt/models/utils.py index de686c226..3a11b5a6f 100644 --- a/python/sglang/srt/models/utils.py +++ b/python/sglang/srt/models/utils.py @@ -37,7 +37,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_context import get_token_to_kv_pool from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_loader.weight_utils import default_weight_loader -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec from sglang.srt.utils import get_current_device_stream_fast, is_cuda, is_hip from sglang.srt.utils.custom_op import register_custom_op @@ -447,7 +447,7 @@ def _reshape_for_qk_norm(x: torch.Tensor, head_dim: int) -> torch.Tensor: if ( _is_cuda - and get_server_args().cuda_graph_config.prefill.tc_compiler == "inductor" + and get_exec().graph.cuda_graph_config.prefill.tc_compiler == "inductor" ): return x.view(*x.shape[:-1], -1, head_dim) return x.reshape(-1, head_dim) @@ -488,7 +488,7 @@ def apply_qk_norm( and allow_inplace # TODO(dark): this can be relaxed if needed and (q_eps == k_eps) # TODO(dark): this can also be relaxed and not envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.get() - and get_server_args().cuda_graph_config.prefill.tc_compiler + and get_exec().graph.cuda_graph_config.prefill.tc_compiler != "inductor" # let inductor fuse QK norm and can_use_fused_inplace_qknorm(head_dim, q.dtype) ): diff --git a/python/sglang/srt/multimodal/internvl_vit_cuda_graph_runner.py b/python/sglang/srt/multimodal/internvl_vit_cuda_graph_runner.py index 10461fe86..23ce493eb 100644 --- a/python/sglang/srt/multimodal/internvl_vit_cuda_graph_runner.py +++ b/python/sglang/srt/multimodal/internvl_vit_cuda_graph_runner.py @@ -22,7 +22,7 @@ import torch import torch.nn as nn from sglang.srt.layers.attention.vision import VisionAttention -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_mm class InternViTCudaGraphRunner: @@ -95,7 +95,7 @@ class InternViTCudaGraphRunner: def _warmup_once(self, key: Hashable) -> None: """Run a tiny eager warmup on the preallocated buffers to trigger lazy init.""" - override_backend = get_server_args().mm_attention_backend + override_backend = get_mm().mm_attention_backend cu = self.cu[key] cu_kk = self.cu_kk[key] max_len = int(cu_kk.max().item()) if cu_kk.numel() else 0 @@ -115,7 +115,7 @@ class InternViTCudaGraphRunner: def _capture_graph(self, key: Hashable) -> None: g = torch.cuda.CUDAGraph() - override_backend = get_server_args().mm_attention_backend + override_backend = get_mm().mm_attention_backend cu = self.cu[key] cu_kk = self.cu_kk[key] diff --git a/python/sglang/srt/multimodal/vit_cuda_graph_runner.py b/python/sglang/srt/multimodal/vit_cuda_graph_runner.py index e05171e34..336a09b6a 100644 --- a/python/sglang/srt/multimodal/vit_cuda_graph_runner.py +++ b/python/sglang/srt/multimodal/vit_cuda_graph_runner.py @@ -25,7 +25,7 @@ import torch.nn as nn from sglang.srt.distributed.parallel_state import get_tp_group from sglang.srt.layers.attention.vision import VisionAttention -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_mm class ViTCudaGraphRunner: @@ -151,7 +151,7 @@ class ViTCudaGraphRunner: cu_full_kk = self.cu_full_len_kk[graph_key] max_full_len = int(cu_full_kk.max().item()) - override_backend = get_server_args().mm_attention_backend + override_backend = get_mm().mm_attention_backend if self._fullatt_block_indexes and 0 not in vit.fullatt_block_indexes: warmup_cu_ws = [cu_window, cu_window_kk, max_window_len] diff --git a/python/sglang/srt/multiplex/multiplexing_mixin.py b/python/sglang/srt/multiplex/multiplexing_mixin.py index 419dbe9b1..5a4e2ffc7 100644 --- a/python/sglang/srt/multiplex/multiplexing_mixin.py +++ b/python/sglang/srt/multiplex/multiplexing_mixin.py @@ -21,6 +21,7 @@ from sglang.srt.multiplex.pdmux_context import ( load_pdmux_config, set_current_stream_idx, ) +from sglang.srt.runtime_context import get_disagg if TYPE_CHECKING: from sglang.srt.managers.schedule_batch import ScheduleBatch @@ -36,7 +37,7 @@ class SchedulerMultiplexMixin: self.split_prefill_batch: Optional[ScheduleBatch] = None # for pd_multiplexing, Init stream_groups, exclude normal stream for prefill only and decode only - self.pdmux_config = load_pdmux_config(self.server_args.pdmux_config_path) + self.pdmux_config = load_pdmux_config(get_disagg().pdmux_config_path) initialize_stream_groups(self.gpu_id, self.pdmux_config) self.stream_groups = get_stream_groups() self.sm_counts = get_sm_counts() diff --git a/python/sglang/srt/sampling/sampling_batch_info.py b/python/sglang/srt/sampling/sampling_batch_info.py index b6005c571..a85faafd2 100644 --- a/python/sglang/srt/sampling/sampling_batch_info.py +++ b/python/sglang/srt/sampling/sampling_batch_info.py @@ -12,7 +12,7 @@ from sglang.srt.constrained.base_grammar_backend import ( GrammarMask, GrammarRow, ) -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec, get_server_args from sglang.srt.sampling.custom_logit_processor import CustomLogitProcessor from sglang.srt.sampling.penaltylib.repetition_penalty import apply_scaling_penalties from sglang.srt.sampling.sampling_params import TOP_K_ALL @@ -86,7 +86,7 @@ class SamplingBatchInfo: @classmethod def from_schedule_batch(cls, batch: ScheduleBatch, vocab_size: int): global_server_args = get_server_args() - enable_deterministic = global_server_args.enable_deterministic_inference + enable_deterministic = get_exec().deterministic.enable_deterministic_inference reqs = batch.reqs device = batch.device @@ -142,7 +142,7 @@ class SamplingBatchInfo: # Check if any request has custom logit processor has_custom_logit_processor = ( - global_server_args.enable_custom_logit_processor + get_exec().features.enable_custom_logit_processor and any(r.custom_logit_processor for r in reqs) # check the flag first. ) # then check the requests. return_sampling_masks = [r.return_sampling_mask for r in reqs] diff --git a/python/sglang/srt/speculative/base_spec_worker.py b/python/sglang/srt/speculative/base_spec_worker.py index eb2119413..b83647b3a 100644 --- a/python/sglang/srt/speculative/base_spec_worker.py +++ b/python/sglang/srt/speculative/base_spec_worker.py @@ -5,6 +5,8 @@ from typing import TYPE_CHECKING, Optional import torch +from sglang.srt.runtime_context import get_exec, get_schedule + if TYPE_CHECKING: from sglang.srt.managers.io_struct import ( UpdateWeightFromDiskReqInput, @@ -62,13 +64,13 @@ class EagleDraftWorkerBase(ABC): num_steps = self.speculative_num_steps sa = self.server_args decode_max_bs = ( - sa.cuda_graph_config.decode.max_bs - if sa.cuda_graph_config is not None + get_exec().graph.cuda_graph_config.decode.max_bs + if get_exec().graph.cuda_graph_config is not None else None ) max_bs = max( decode_max_bs or 0, - sa.max_running_requests or 0, + get_schedule().max_running_requests or 0, 1, ) # A single-step chain has no parent entries (slow path drops the last diff --git a/python/sglang/srt/speculative/dflash_info_v2.py b/python/sglang/srt/speculative/dflash_info_v2.py index 72162ba51..a59034ee8 100644 --- a/python/sglang/srt/speculative/dflash_info_v2.py +++ b/python/sglang/srt/speculative/dflash_info_v2.py @@ -9,7 +9,7 @@ import torch from sglang.srt.environ import envs from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.mem_cache.allocation import alloc_for_spec_decode -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_spec from sglang.srt.speculative.spec_info import SpecInput, SpecInputType from sglang.srt.utils.common import is_pin_memory_available @@ -134,7 +134,7 @@ class DFlashDraftInputV2(SpecInput): cur_kv_lens_cpu_t = self._prepare_cur_kv_lens_cpu_buf[:bs] # For DFLASH, each decode step needs a fixed-size verify block. - block_size = int(get_server_args().speculative_num_draft_tokens) + block_size = int(get_spec().speculative_num_draft_tokens) if block_size <= 0: raise ValueError( f"DFLASH invalid speculative_num_draft_tokens={block_size}." diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py index e4d5c2c54..5677048af 100644 --- a/python/sglang/srt/speculative/dflash_worker_v2.py +++ b/python/sglang/srt/speculative/dflash_worker_v2.py @@ -27,6 +27,7 @@ from sglang.srt.model_executor.forward_batch_info import ( ForwardMode, compute_position, ) +from sglang.srt.runtime_context import get_exec from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.base_spec_worker import BaseSpecWorker from sglang.srt.speculative.dflash_info import DFlashVerifyInput @@ -337,7 +338,7 @@ class DFlashWorkerV2(BaseSpecWorker): def init_cuda_graphs(self): capture_decode_cuda_graph = ( - self.server_args.cuda_graph_config.decode.backend != Backend.DISABLED + get_exec().graph.cuda_graph_config.decode.backend != Backend.DISABLED ) if is_cuda() and capture_decode_cuda_graph: available_mem = get_available_gpu_memory(self.device, self.gpu_id) @@ -399,7 +400,7 @@ class DFlashWorkerV2(BaseSpecWorker): block_size=self.block_size, num_org=num_org, org_vocab_start=org_vocab_start, - max_bs=max(self.server_args.cuda_graph_config.decode.bs), + max_bs=max(get_exec().graph.cuda_graph_config.decode.bs), tp_group=tp_group if tp_group.world_size > 1 else None, ) @@ -1264,7 +1265,7 @@ class DFlashWorkerV2(BaseSpecWorker): mamba_steps_to_track = None if batch.mamba_track_indices is not None: - mamba_track_interval = self.server_args.mamba_track_interval + mamba_track_interval = get_exec().mamba.mamba_track_interval to_track_mask = ( seq_lens_pre_verify // mamba_track_interval != batch.seq_lens // mamba_track_interval diff --git a/python/sglang/srt/speculative/draft_utils.py b/python/sglang/srt/speculative/draft_utils.py index c9d8cd926..c05a57709 100644 --- a/python/sglang/srt/speculative/draft_utils.py +++ b/python/sglang/srt/speculative/draft_utils.py @@ -1,3 +1,4 @@ +from sglang.srt.runtime_context import get_spec from sglang.srt.server_args import ServerArgs from sglang.srt.utils.common import ( cpu_has_amx_support, @@ -112,7 +113,7 @@ class DraftBackendFactory: } backend_name = ( "decode_attention_backend" - if self.server_args.speculative_attention_mode == "decode" + if get_spec().speculative_attention_mode == "decode" else "prefill_attention_backend" ) backend = self._create_backend( diff --git a/python/sglang/srt/speculative/dspark_components/dspark_planner.py b/python/sglang/srt/speculative/dspark_components/dspark_planner.py index 3fbd56fe2..c67da766d 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_planner.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_planner.py @@ -20,7 +20,7 @@ from sglang.srt.managers.overlap_utils import ( ResolvedConfidence, ) from sglang.srt.managers.schedule_batch import ScheduleBatch -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_disagg, get_parallel, get_schedule, get_spec from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2 from sglang.srt.speculative.dflash_utils import apply_dflash_verify_logits_adjustments @@ -158,7 +158,7 @@ class DSparkVerifyPlanner: ) relay_lag_steps = ( 0 - if self.server_args.disable_overlap_schedule + if get_schedule().disable_overlap_schedule else CONFIDENCE_RELAY_RING_LAG ) self._budget_planner = HostConfidenceBudgetPlanner( @@ -174,16 +174,15 @@ class DSparkVerifyPlanner: and get_parallel().attn_tp_size == 1 and get_parallel().attn_cp_size == 1 and require_mlp_tp_gather(self.server_args) - and not self.server_args.disable_overlap_schedule - and not self.server_args.speculative_skip_dp_mlp_sync - and self.server_args.disaggregation_mode == "null" + and not get_schedule().disable_overlap_schedule + and not get_spec().speculative_skip_dp_mlp_sync + and get_disagg().disaggregation_mode == "null" and self.server_args.pp_size == 1 and not envs.SGLANG_SCHEDULER_SKIP_ALL_GATHER.get() ) if tp_rank == 0: sps_table_source = ( - self.server_args.speculative_dspark_sps_table_path - or "uninitialized" + get_spec().speculative_dspark_sps_table_path or "uninitialized" ) logger.info( "DSpark ragged-verify scheduler enabled (mode=%s, lag=%d, " @@ -382,7 +381,7 @@ class DSparkVerifyPlanner: the draft input by prepare_verify_budget; otherwise compute it now.""" if not self.schedules_verify_budget or confidence is None: return None - if not self.server_args.disable_overlap_schedule: + if not get_schedule().disable_overlap_schedule: return draft_input.verify_token_budget return self.compute_budget_sync( confidence=confidence, 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 ae778ee1d..03c2ac155 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py @@ -14,7 +14,7 @@ from sglang.srt.model_executor.forward_batch_info import ( CaptureHiddenMode, compute_position, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.base_spec_worker import BaseSpecWorker from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2 @@ -298,7 +298,7 @@ class DSparkWorkerV2(BaseSpecWorker): self._draft_worker.init_attention_backends() def init_cuda_graphs(self): - capture_decode_cuda_graph = not self.server_args.disable_cuda_graph + capture_decode_cuda_graph = not get_exec().graph.disable_cuda_graph if is_cuda() and capture_decode_cuda_graph: available_mem = get_available_gpu_memory(self.device, self.gpu_id) if available_mem < 1.0: @@ -324,7 +324,7 @@ class DSparkWorkerV2(BaseSpecWorker): return maybe_build_draft_sampler( draft_model=self.draft_model, gamma=self.gamma, - max_bs=max(self.server_args.cuda_graph_config.decode.bs), + max_bs=max(get_exec().graph.cuda_graph_config.decode.bs), device=self.device, tp_rank=self.ps.tp_rank, confidence_fn=( 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 8de99d832..13371f432 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -34,7 +34,7 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen from sglang.srt.model_executor.runner_backend_utils import ( CUDA_GRAPH_CAPTURE_FAILED_MSG, ) -from sglang.srt.runtime_context import get_flags +from sglang.srt.runtime_context import get_flags, get_spec from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.speculative.eagle_info import EagleDraftInput from sglang.srt.speculative.eagle_utils import get_draft_recurrent_hidden_state_spec @@ -120,7 +120,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner): model_runner.server_args.enable_profile_cuda_graph ) self.speculative_num_steps = ( - model_runner.server_args.speculative_num_steps + get_spec().speculative_num_steps if speculative_num_steps is None else speculative_num_steps ) 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 667064c69..850840111 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 @@ -35,7 +35,7 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen from sglang.srt.model_executor.runner_backend_utils import ( CUDA_GRAPH_CAPTURE_FAILED_MSG, ) -from sglang.srt.runtime_context import get_flags +from sglang.srt.runtime_context import get_flags, get_spec from sglang.srt.speculative.eagle_info import EagleDraftExtendInput from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req @@ -106,7 +106,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): model_runner.server_args.enable_profile_cuda_graph ) self.speculative_num_steps = ( - model_runner.server_args.speculative_num_steps + get_spec().speculative_num_steps if speculative_num_steps is None else speculative_num_steps ) diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index 4184fe465..a927efc8b 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -6,7 +6,7 @@ import torch from sglang.kernels.ops.attention.utils import create_flashinfer_kv_indices_triton from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_spec from sglang.srt.speculative.spec_info import SpecInput, SpecInputType logger = logging.getLogger(__name__) @@ -199,7 +199,7 @@ class EagleDraftInput(SpecInput): topk_index=torch.empty((0, topk), device=device, dtype=torch.int64), draft_probs=( torch.empty((0, vocab_size), device=device, dtype=torch.float32) - if get_server_args().speculative_use_rejection_sampling + if get_spec().speculative_use_rejection_sampling else None ), capture_hidden_mode=capture_hidden_mode, diff --git a/python/sglang/srt/speculative/eagle_utils.py b/python/sglang/srt/speculative/eagle_utils.py index 078d0d1a9..e87ecd91e 100644 --- a/python/sglang/srt/speculative/eagle_utils.py +++ b/python/sglang/srt/speculative/eagle_utils.py @@ -659,7 +659,6 @@ def eagle_sample( from sglang.srt.layers.dp_attention import ( is_dp_attention_enabled, ) - from sglang.srt.runtime_context import get_server_args from sglang.srt.sampling.penaltylib.repetition_penalty import ( apply_scaling_penalties, ) @@ -746,7 +745,7 @@ def eagle_sample( chain_speculative_sampling_triton, ) - use_rejection_sampling = get_server_args().speculative_use_rejection_sampling + use_rejection_sampling = get_spec().speculative_use_rejection_sampling # Apply temperature and get target probs expanded_temperature = torch.repeat_interleave( @@ -919,9 +918,8 @@ def eagle_prepare_for_decode(batch: ScheduleBatch): # (get_alloc_reserve_per_decode) outgrows the req_to_token row: the write below # would OOB and free would leak KV. The row is widened to hold it in _init_pools # (PR #26972); fail here with a clear error, not on a later cryptic CUDA assert. - from sglang.srt.runtime_context import get_server_args - if page_size > 1 and (get_server_args().speculative_eagle_topk or 1) > 1: + if page_size > 1 and (get_spec().speculative_eagle_topk or 1) > 1: max_alloc_len = int(nxt_kv_lens_cpu.max()) row_width = batch.req_to_token_pool.req_to_token.shape[1] assert max_alloc_len <= row_width, ( diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 047260130..0139e9e94 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -43,7 +43,13 @@ from sglang.srt.model_executor.runner import ( DecodeCudaGraphRunner, get_batch_sizes_to_capture, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import ( + get_context, + get_exec, + get_model, + get_parallel, + get_spec, +) from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.adaptive_runtime_state import ( AdaptiveController, @@ -136,7 +142,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): # Args for easy access self.device = server_args.device self.topk = server_args.speculative_eagle_topk - if self.server_args.speculative_use_rejection_sampling: + if get_spec().speculative_use_rejection_sampling: assert self.topk == 1, "Chain speculative sampling supports only topk=1" self.speculative_num_steps = server_args.speculative_num_steps self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens @@ -192,7 +198,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): self.init_token_map() self.init_lm_head() - if self.server_args.speculative_use_rejection_sampling: + if get_spec().speculative_use_rejection_sampling: target_vocab_size = self.target_worker.model_config.vocab_size draft_vocab_size = ( self.hot_token_id.shape[0] @@ -251,13 +257,13 @@ class EagleDraftWorker(EagleDraftWorkerBase): def init_token_map(self): # Load hot token ids if self.speculative_algorithm.is_eagle3(): - if self.server_args.speculative_token_map is not None: + if get_spec().speculative_token_map is not None: logger.warning( "Speculative token map specified, but EAGLE3 models already have this. Ignoring the specified token map." ) self.hot_token_id = None - elif self.server_args.speculative_token_map is not None: - self.hot_token_id = load_token_map(self.server_args.speculative_token_map) + elif get_spec().speculative_token_map is not None: + self.hot_token_id = load_token_map(get_spec().speculative_token_map) self.server_args.override( "eagle_worker.hot_token_map", json_model_override_args=( @@ -342,7 +348,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): if _is_cpu or check_cuda_graph_backend(Phase.DECODE, Backend.DISABLED): return - if self.server_args.model_impl == "mindspore": + if get_model().model_impl == "mindspore": return Device2DraftCudaGraphRunner = { @@ -352,7 +358,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): "musa": EAGLEDraftCudaGraphRunner, } # Capture draft - decode_backend = self.server_args.cuda_graph_config.decode.backend + decode_backend = get_exec().graph.cuda_graph_config.decode.backend capture_bs, _ = get_batch_sizes_to_capture(self.draft_runner) if self.speculative_num_steps > 1: tic = time.perf_counter() @@ -549,7 +555,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): score_list: List[torch.Tensor] = [] token_list: List[torch.Tensor] = [] parents_list: List[torch.Tensor] = [] - if self.server_args.speculative_use_rejection_sampling: + if get_spec().speculative_use_rejection_sampling: draft_probs_list: List[torch.Tensor] = [spec_info.draft_probs] topk1_chain_fits = ( @@ -564,7 +570,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): topk1_chain_fits and _is_cuda and self.hot_token_id is None - and not self.server_args.speculative_use_rejection_sampling + and not get_spec().speculative_use_rejection_sampling ): draft_tokens_topk1 = torch.empty( (topk_index.shape[0], self.speculative_num_steps), @@ -629,7 +635,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): logits_output = self.draft_runner.forward(forward_batch).logits_output maybe_detect_nan(logits_output.next_token_logits, f"draft_forward step {i}") maybe_detect_inf(logits_output.next_token_logits, f"draft_forward step {i}") - if self.server_args.speculative_use_rejection_sampling: + if get_spec().speculative_use_rejection_sampling: probs, topk_p, topk_index = sample_draft_proposal( logits_output.next_token_logits, forward_batch.sampling_info.temperatures, @@ -655,7 +661,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): probs = renorm_draft_probs( logits_output.next_token_logits, forward_batch.sampling_info, - self.server_args.speculative_use_rejection_sampling, + get_spec().speculative_use_rejection_sampling, ) topk_p, topk_index = fast_topk(probs, self.topk, dim=-1) forward_batch.positions.add_(1) @@ -675,7 +681,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): draft_probs = ( torch.stack(draft_probs_list, dim=1) - if self.server_args.speculative_use_rejection_sampling + if get_spec().speculative_use_rejection_sampling else None ) @@ -795,7 +801,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): prefill_dsa_topk = self.dsa_extend_topk_buf[:bs].clone() # Assemble the next-iter draft spec_info from the extend output. - use_rejection_sampling = self.server_args.speculative_use_rejection_sampling + use_rejection_sampling = get_spec().speculative_use_rejection_sampling probs = renorm_draft_probs( logits_output.next_token_logits, batch.sampling_info, @@ -941,7 +947,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): ] # The draft-extend graph only anchors full logits; selected-row topk is # owned by the worker for both graph and eager paths. - if self.server_args.speculative_use_rejection_sampling: + if get_spec().speculative_use_rejection_sampling: ret_draft_probs, ret_topk_p, ret_topk_index = sample_draft_proposal( draft_logits_output.next_token_logits, batch.sampling_info.temperatures, @@ -958,7 +964,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): probs = renorm_draft_probs( draft_logits_output.next_token_logits, batch.sampling_info, - self.server_args.speculative_use_rejection_sampling, + get_spec().speculative_use_rejection_sampling, ) ret_topk_p, ret_topk_index = fast_topk(probs, self.topk, dim=-1) ret_draft_probs = None @@ -975,7 +981,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): ret_topk_index, ret_hidden_states, ) - if self.server_args.speculative_use_rejection_sampling: + if get_spec().speculative_use_rejection_sampling: next_draft_input.draft_probs = ret_draft_probs if self.seed_dsa_topk_from_draft_extend: next_draft_input.dsa_topk_indices = dsa_seed_topk_indices @@ -1078,7 +1084,7 @@ class EAGLEWorkerV2(BaseSpecWorker): cuda_graph_bs=( None if check_cuda_graph_backend(Phase.DECODE, Backend.DISABLED) - else self.server_args.cuda_graph_bs_decode + else get_exec().graph.cuda_graph_bs_decode ), ) @@ -1380,7 +1386,7 @@ class EAGLEWorkerV2(BaseSpecWorker): ) # Sync server_args - self.server_args.override( + get_context().override( "adaptive_spec.restore", speculative_num_steps=state.speculative_num_steps, speculative_num_draft_tokens=state.speculative_num_draft_tokens, @@ -1394,7 +1400,6 @@ class EAGLEWorkerV2(BaseSpecWorker): cuda_graph_bs: list[int] | None = None, ): """Temporarily override server_args and worker attributes for graph capture.""" - sa = self.server_args dw = self._draft_worker backup = ( self.speculative_num_steps, @@ -1407,17 +1412,17 @@ class EAGLEWorkerV2(BaseSpecWorker): dw.draft_runner.attn_backend, dw.cuda_graph_runner, dw.cuda_graph_runner_for_draft_extend, - sa.speculative_num_steps, - sa.speculative_num_draft_tokens, - sa.cuda_graph_bs_decode, - sa.disable_cuda_graph, + get_spec().speculative_num_steps, + get_spec().speculative_num_draft_tokens, + get_exec().graph.cuda_graph_bs_decode, + get_exec().graph.disable_cuda_graph, ) self.speculative_num_steps = speculative_num_steps self.speculative_num_draft_tokens = speculative_num_draft_tokens dw.speculative_num_steps = speculative_num_steps dw.speculative_num_draft_tokens = speculative_num_draft_tokens - sa.override( + get_context().override( "adaptive_spec.capture_override", speculative_num_steps=speculative_num_steps, speculative_num_draft_tokens=speculative_num_draft_tokens, @@ -1427,7 +1432,7 @@ class EAGLEWorkerV2(BaseSpecWorker): # for steps that no BS range uses (e.g. step=1). Disable graph # capture for those steps; restore in finally so subsequent steps # are not affected. - sa.override( + get_context().override( "adaptive_spec.capture_override", cuda_graph_bs_decode=cuda_graph_bs, **({"disable_cuda_graph": True} if not cuda_graph_bs else {}), @@ -1449,7 +1454,7 @@ class EAGLEWorkerV2(BaseSpecWorker): dw.cuda_graph_runner, dw.cuda_graph_runner_for_draft_extend, ) = backup[:10] - sa.override( + get_context().override( "adaptive_spec.capture_restore", speculative_num_steps=backup[10], speculative_num_draft_tokens=backup[11], 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 0cce8c07f..3fdf45e77 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 @@ -32,7 +32,7 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen from sglang.srt.model_executor.runner_backend_utils import ( CUDA_GRAPH_CAPTURE_FAILED_MSG, ) -from sglang.srt.runtime_context import get_flags +from sglang.srt.runtime_context import get_flags, get_spec from sglang.srt.speculative.frozen_kv_mtp_info import FrozenKVMTPDraftInput from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req from sglang.srt.utils import ( @@ -96,7 +96,7 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner): self.tp_size = self.model_runner.ps.tp_size self.attn_dp_size = self.model_runner.ps.attn_dp_size self.pp_size = model_runner.server_args.pp_size - self.speculative_num_steps = model_runner.server_args.speculative_num_steps + self.speculative_num_steps = get_spec().speculative_num_steps self.topk = model_runner.server_args.speculative_eagle_topk self.draft_attn_backend = frozen_kv_mtp_worker.draft_attn_backend self.enable_profile_cuda_graph = ( 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 54844777a..6f51e67a5 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 @@ -59,7 +59,7 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen from sglang.srt.model_executor.runner_backend_utils import ( CUDA_GRAPH_CAPTURE_FAILED_MSG, ) -from sglang.srt.runtime_context import get_flags +from sglang.srt.runtime_context import get_flags, get_spec from sglang.srt.speculative.eagle_info import EagleDraftExtendInput from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim from sglang.srt.speculative.multi_layer_eagle_utils import ( @@ -159,10 +159,8 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): 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.enable_pdmux = model_runner.server_args.enable_pdmux - self.speculative_num_steps = model_runner.server_args.speculative_num_steps - self.speculative_num_draft_tokens = ( - model_runner.server_args.speculative_num_draft_tokens - ) + self.speculative_num_steps = get_spec().speculative_num_steps + self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens self.topk = model_runner.server_args.speculative_eagle_topk self.enable_profile_cuda_graph = ( model_runner.server_args.enable_profile_cuda_graph diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index 44ccccb21..1c0f3fbb3 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -46,7 +46,7 @@ from sglang.srt.mem_cache.allocation import ( from sglang.srt.mem_cache.allocation import ( assign_req_to_token_pool_func as assign_req_to_token_pool_func, ) -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_exec, get_server_args from sglang.srt.utils import ( is_cpu, is_cuda, @@ -805,7 +805,7 @@ def _verify_commit_step_indices( return last_correct_step_indices, None seq_lens_pre_verify = batch.seq_lens seq_lens_post_verify = batch.seq_lens + accept_lens - mamba_track_interval = get_server_args().mamba_track_interval + mamba_track_interval = get_exec().mamba.mamba_track_interval to_track_mask = ( seq_lens_pre_verify // mamba_track_interval != seq_lens_post_verify // mamba_track_interval @@ -969,7 +969,7 @@ def spec_prepare_for_decode(batch: ScheduleBatch) -> None: if server_args.enable_mamba_extra_buffer_lazy(): # Scheduler phase (outside forward isolation). batch.mamba_lazy_spec_prepare( - server_args.mamba_track_interval, + get_exec().mamba.mamba_track_interval, server_args.max_speculative_num_draft_tokens, ) if batch.spec_algorithm.is_dflash_family(): diff --git a/python/sglang/srt/state_capturer/indexer_topk.py b/python/sglang/srt/state_capturer/indexer_topk.py index afa652cec..de5554648 100644 --- a/python/sglang/srt/state_capturer/indexer_topk.py +++ b/python/sglang/srt/state_capturer/indexer_topk.py @@ -6,7 +6,7 @@ import pybase64 import torch from sglang.srt.configs.model_config import ModelConfig, get_num_indexer_layers -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_exec, get_parallel, get_schedule from sglang.srt.state_capturer.base import BaseTopkCapturer logger = logging.getLogger(__name__) @@ -32,7 +32,7 @@ class IndexerTopkCapturer(BaseTopkCapturer): # DP-attention capture is per-rank-local: each rank writes [:local_batch, ...] # to its own device_cache, so the buffer only needs to fit one rank's batch. server_args = get_server_args() - max_batch_size = max(server_args.chunked_prefill_size, max_running_requests) + max_batch_size = max(get_schedule().chunked_prefill_size, max_running_requests) super().__init__( num_tokens=num_tokens, @@ -89,9 +89,8 @@ def create_indexer_capturer( max_running_requests: int, device: str, ) -> Optional[IndexerTopkCapturer]: - from sglang.srt.runtime_context import get_server_args - enable = get_server_args().enable_return_indexer_topk + enable = get_exec().features.enable_return_indexer_topk # Producer wiring is CUDA-only (Indexer.forward_cuda + MLA skip_topk # path); other backends would create a capturer but never feed it. if enable and device != "cuda": diff --git a/python/sglang/srt/state_capturer/routed_experts.py b/python/sglang/srt/state_capturer/routed_experts.py index 65f6cc052..5dccd46c4 100644 --- a/python/sglang/srt/state_capturer/routed_experts.py +++ b/python/sglang/srt/state_capturer/routed_experts.py @@ -12,7 +12,12 @@ from sglang.srt.layers.dp_attention import ( ) from sglang.srt.layers.moe import get_moe_a2a_backend from sglang.srt.model_executor.forward_batch_info import ForwardBatch -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_exec, + get_parallel, + get_schedule, + get_server_args, +) from sglang.srt.state_capturer.base import BaseTopkCapturer @@ -36,9 +41,9 @@ class RoutedExpertsCapturer(BaseTopkCapturer): device: str, ) -> Optional["RoutedExpertsCapturer"]: server_args = get_server_args() - if not server_args.enable_return_routed_experts: + if not get_exec().features.enable_return_routed_experts: return None - if not server_args.disable_shared_experts_fusion and hasattr( + if not get_exec().moe.disable_shared_experts_fusion and hasattr( model, "num_fused_shared_experts" ): num_fused_shared_experts = model.num_fused_shared_experts @@ -71,7 +76,7 @@ class RoutedExpertsCapturer(BaseTopkCapturer): # chunked_prefill_size. # FIXME: spec decoding's num_verify_tokens is still not accounted for. max_batch_size = max( - server_args.chunked_prefill_size * server_args.dp_size, + get_schedule().chunked_prefill_size * server_args.dp_size, max_running_requests * server_args.dp_size, ) diff --git a/python/sglang/srt/utils/profile_utils.py b/python/sglang/srt/utils/profile_utils.py index b3946eeb7..562ffd28d 100644 --- a/python/sglang/srt/utils/profile_utils.py +++ b/python/sglang/srt/utils/profile_utils.py @@ -13,7 +13,7 @@ from sglang.srt.environ import envs from sglang.srt.managers.io_struct import ProfileReqOutput from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.platforms import current_platform -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_device from sglang.srt.utils import is_npu from sglang.srt.utils.torch_npu_patch_utils import apply_torch_npu_patches @@ -62,7 +62,7 @@ class ProfileManager: ) self.ps = ps self.cpu_group = cpu_group - self.first_rank_in_node = ps.gpu_id == get_server_args().base_gpu_id + self.first_rank_in_node = ps.gpu_id == get_device().base_gpu_id self.profiler_kwargs = None self.profiler = None diff --git a/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_runner.py b/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_runner.py index e90576627..eb150e7c9 100644 --- a/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_runner.py +++ b/python/sglang/test/kits/attention_unittest/runner_modes/speculative_draft_runner.py @@ -327,6 +327,10 @@ def _configure_runner_for_eagle_draft( "use_mla_backend": runner.use_mla_backend, } server_args.override(source="attention-unittest-eagle-draft", **updates) + # Re-publish so the bags pick up the overrides. + from sglang.srt.runtime_context import get_context + + get_context().set_server_args(server_args) runner.spec_algorithm = SpeculativeAlgorithm.EAGLE runner.is_draft_worker = True @@ -390,6 +394,9 @@ def _build_frozen_kv_mtp_fixture( fixture.runner.server_args.override( "attention_unittest.frozen_kv_draft", speculative_algorithm="FROZEN_KV_MTP" ) + from sglang.srt.runtime_context import get_context + + get_context().set_server_args(fixture.runner.server_args) fixture.runner.spec_algorithm = SpeculativeAlgorithm.FROZEN_KV_MTP fixture.runner.draft_attn_backend = fixture.backend fixture.runner.attn_backend = fixture.backend diff --git a/test/registered/dcp/test_dcp_layout_unit.py b/test/registered/dcp/test_dcp_layout_unit.py index 2ee7e44e4..c310eb6b1 100644 --- a/test/registered/dcp/test_dcp_layout_unit.py +++ b/test/registered/dcp/test_dcp_layout_unit.py @@ -19,6 +19,7 @@ from unittest.mock import MagicMock, patch import torch +from sglang.srt import runtime_context as rc from sglang.srt.layers.dcp.layout import get_dcp_lens from sglang.srt.mem_cache.allocator.paged import PagedTokenToKVPoolAllocator from sglang.srt.mem_cache.kv_cache_configurator import KVCacheConfigurator @@ -136,6 +137,17 @@ class TestGetDcpLens(CustomTestCase): ) allocators = {} + # The configurator's bag reads (disaggregation_mode / page_size / + # enable_hisparse) come from the published context; the per-iteration + # dcp_size stays on the injected instance stand-in. + self._sa_override = rc.get_context().override_server_args( + disaggregation_mode="null", + page_size=physical_page_size, + enable_hisparse=False, + ) + self._sa_override.install() + self.addCleanup(self._sa_override.restore) + for dcp_size in (1, 4): configurator = SimpleNamespace( server_args=SimpleNamespace( diff --git a/test/registered/unit/mem_cache/test_mamba_donated_alloc_ratio.py b/test/registered/unit/mem_cache/test_mamba_donated_alloc_ratio.py index cba3c2543..4f3937562 100644 --- a/test/registered/unit/mem_cache/test_mamba_donated_alloc_ratio.py +++ b/test/registered/unit/mem_cache/test_mamba_donated_alloc_ratio.py @@ -117,8 +117,17 @@ class TestMambaRatioEnvGate(unittest.TestCase): enable_mamba_extra_buffer_lazy=lambda: lazy, ) fake = SimpleNamespace(server_args=server_args) + # The bag reads (disable_radix_cache / disable_overlap_schedule) come + # from the published context; the derived-method calls stay on the + # injected stand-in. + from sglang.srt import runtime_context as rc + with envs.SGLANG_OPT_MAMBA_SKIP_DECODE_LOCK.override(skip): - return KVCacheConfigurator._calculate_mamba_ratio(fake) + with rc.get_context().override_server_args( + disable_radix_cache=False, + disable_overlap_schedule=disable_overlap, + ): + return KVCacheConfigurator._calculate_mamba_ratio(fake) def test_flag_off_restores_original_ratios(self): r = lambda **kw: self._ratio(skip=False, **kw) diff --git a/test/registered/unit/model_loader/test_prefetch_checkpoints.py b/test/registered/unit/model_loader/test_prefetch_checkpoints.py index 367c44312..399260809 100644 --- a/test/registered/unit/model_loader/test_prefetch_checkpoints.py +++ b/test/registered/unit/model_loader/test_prefetch_checkpoints.py @@ -340,6 +340,10 @@ class TestPrefetchDispatch(CustomTestCase): drop_cache, ), ), + patch( + "sglang.srt.model_loader.loader.get_model", + return_value=self._server_args(prefetch, disable_mmap, drop_cache), + ), patch( "sglang.srt.model_loader.loader." "buffered_multi_thread_safetensors_weights_iterator", @@ -356,12 +360,13 @@ class TestPrefetchDispatch(CustomTestCase): """Prefetch on + no explicit multithread config -> single-threaded, and the opt-out warning fires once.""" loader = self._make_loader({}) - p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch( + p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch( prefetch=True ) with ( p_prep, p_args, + p_model, p_buffered as mock_buffered, p_single as mock_single, p_warn as mock_warning, @@ -375,12 +380,13 @@ class TestPrefetchDispatch(CustomTestCase): """Explicit enable_multithread_load=true is the escape hatch; the override and its warning must not fire.""" loader = self._make_loader({"enable_multithread_load": True}) - p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch( + p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch( prefetch=True ) with ( p_prep, p_args, + p_model, p_buffered as mock_buffered, p_single as mock_single, p_warn as mock_warning, @@ -395,12 +401,13 @@ class TestPrefetchDispatch(CustomTestCase): default) also signals multi-thread intent, so the override must not fire and num_threads stays live.""" loader = self._make_loader({"num_threads": 64}) - p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch( + p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch( prefetch=True ) with ( p_prep, p_args, + p_model, p_buffered as mock_buffered, p_single as mock_single, p_warn as mock_warning, @@ -416,12 +423,13 @@ class TestPrefetchDispatch(CustomTestCase): """Prefetch off -> multi-threaded iterator is used (default), no override warning.""" loader = self._make_loader({}) - p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch( + p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch( prefetch=False ) with ( p_prep, p_args, + p_model, p_buffered as mock_buffered, p_single as mock_single, p_warn as mock_warning, @@ -435,12 +443,13 @@ class TestPrefetchDispatch(CustomTestCase): """Prefetch is a no-op without mmap, so the override and its warning must not fire.""" loader = self._make_loader({}) - p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch( + p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch( prefetch=True, disable_mmap=True ) with ( p_prep, p_args, + p_model, p_buffered as mock_buffered, p_single as mock_single, p_warn as mock_warning, @@ -454,7 +463,7 @@ class TestPrefetchDispatch(CustomTestCase): """FASTSAFETENSORS ignores both flags; override + warning must not fire.""" loader = self._make_loader({}, load_format=LoadFormat.FASTSAFETENSORS) - p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch( + p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch( prefetch=True ) with ( @@ -464,6 +473,7 @@ class TestPrefetchDispatch(CustomTestCase): ) as mock_fast, p_prep, p_args, + p_model, p_buffered as mock_buffered, p_single as mock_single, p_warn as mock_warning, @@ -482,7 +492,7 @@ class TestPrefetchDispatch(CustomTestCase): loader = self._make_loader( {"enable_gds": False}, load_format=LoadFormat.FASTSAFETENSORS ) - p_prep, p_args, p_buffered, p_single, p_warn = self._patch_dispatch( + p_prep, p_args, p_model, p_buffered, p_single, p_warn = self._patch_dispatch( prefetch=False, drop_cache=True, ) @@ -493,6 +503,7 @@ class TestPrefetchDispatch(CustomTestCase): ) as mock_fast, p_prep, p_args, + p_model, p_buffered, p_single, p_warn, diff --git a/test/registered/unit/model_loader/test_presharded_loader.py b/test/registered/unit/model_loader/test_presharded_loader.py index b2a8dda86..8d0c241d8 100644 --- a/test/registered/unit/model_loader/test_presharded_loader.py +++ b/test/registered/unit/model_loader/test_presharded_loader.py @@ -819,6 +819,16 @@ class TestShardConfig(unittest.TestCase): ), mock.patch( "sglang.srt.model_loader.loader.get_parallel", return_value=parallel, + ), mock.patch( + "sglang.srt.model_loader.loader.get_exec", + return_value=SimpleNamespace( + features=SimpleNamespace(enable_fp32_lm_head=True), + moe=SimpleNamespace( + ep_num_redundant_experts=4, + enable_eplb=True, + init_expert_location="trivial", + ), + ), ), mock.patch.object( loader, "_compute_structural_signature", return_value="sig16" ): diff --git a/test/registered/unit/sampling/test_sampling_batch_info.py b/test/registered/unit/sampling/test_sampling_batch_info.py index 66159b2fb..7d1e5ab17 100644 --- a/test/registered/unit/sampling/test_sampling_batch_info.py +++ b/test/registered/unit/sampling/test_sampling_batch_info.py @@ -6,6 +6,7 @@ register_cpu_ci(est_time=9, suite="base-a-test-cpu") register_cpu_ci(est_time=8, suite="base-c-test-cpu") import unittest +from types import SimpleNamespace from unittest.mock import MagicMock, patch import torch @@ -455,6 +456,22 @@ class TestCopyForForward(CustomTestCase): # from_schedule_batch class TestFromScheduleBatch(CustomTestCase): + def setUp(self): + super().setUp() + # from_schedule_batch reads these two flags from the exec bag; give + # each test a mutable stand-in so it does not depend on a published + # (or leaked) process context. + self._exec_ns = SimpleNamespace( + deterministic=SimpleNamespace(enable_deterministic_inference=False), + features=SimpleNamespace(enable_custom_logit_processor=False), + ) + exec_patch = patch( + "sglang.srt.sampling.sampling_batch_info.get_exec", + return_value=self._exec_ns, + ) + exec_patch.start() + self.addCleanup(exec_patch.stop) + def _make_req( self, temp=1.0, @@ -537,6 +554,7 @@ class TestFromScheduleBatch(CustomTestCase): """Test that explicit seed=123 is kept and missing seed defaults to 42.""" mock_server_args.return_value.enable_deterministic_inference = True mock_server_args.return_value.enable_custom_logit_processor = False + self._exec_ns.deterministic.enable_deterministic_inference = True reqs = [self._make_req(seed=123), self._make_req(seed=None)] batch = MagicMock() @@ -585,6 +603,7 @@ class TestFromScheduleBatch(CustomTestCase): mock_server_args.return_value.enable_deterministic_inference = False mock_server_args.return_value.enable_custom_logit_processor = True + self._exec_ns.features.enable_custom_logit_processor = True proc_str = DisallowedTokensLogitsProcessor.to_str() req1 = self._make_req() diff --git a/test/registered/unit/spec/test_ngram_mamba_verify_update.py b/test/registered/unit/spec/test_ngram_mamba_verify_update.py index 7d147b1ee..c9f70fa91 100644 --- a/test/registered/unit/spec/test_ngram_mamba_verify_update.py +++ b/test/registered/unit/spec/test_ngram_mamba_verify_update.py @@ -200,6 +200,9 @@ class TestNgramMambaVerifyUpdate(CustomTestCase): ), patch( "sglang.srt.speculative.spec_utils.get_server_args", return_value=MagicMock(mamba_track_interval=256), + ), patch( + "sglang.srt.speculative.spec_utils.get_exec", + return_value=MagicMock(mamba=MagicMock(mamba_track_interval=256)), ): commit_mamba_states_after_verify( target_worker, diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index a1b715b2e..669422c30 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -10,7 +10,7 @@ import unittest from unittest.mock import patch import sglang.srt.server_args as server_args_module -from sglang.srt.arg_groups.arg_utils import A, Arg +from sglang.srt.arg_groups.arg_utils import NS, A, Arg from sglang.srt.runtime_context import ( Flags, ParallelContext, @@ -19,6 +19,7 @@ from sglang.srt.runtime_context import ( get_context, get_flags, get_parallel, + get_schedule, get_server_args, reset_context, ) @@ -375,8 +376,10 @@ class TestFlagsTier(_IsolatedServerArgs): class _FakeResolvedArgs: """Publishable fixture with a resolvable whitelist (real flat leaves).""" - page_size: A[int | None, Arg(help="p", resolvable=True)] = None - sampling_backend: A[str | None, Arg(help="s", resolvable=True)] = None + page_size: A[int | None, Arg(help="p", resolvable=True), NS("schedule")] = None + sampling_backend: A[ + str | None, Arg(help="s", resolvable=True), NS("exec.kernel") + ] = None _resolved_overrides: list = dataclasses.field(default_factory=list) @@ -935,12 +938,15 @@ class TestPublishLifecycle(_IsolatedServerArgs): get_context().set_server_args(object()) self.assertFalse(get_flags().capture.enable_torch_compile) - def test_declare_load_time_override_writes_through(self): + def test_declare_load_time_override_writes_the_bag(self): from sglang.srt.arg_groups.overrides import declare_load_time_override args = self._publish(page_size=1) declare_load_time_override("model.load_time", {"page_size": 64}) - self.assertEqual(args.page_size, 64) + # The declaration lands on the config bag; the pristine startup record + # (server_args) is untouched. + self.assertEqual(get_schedule().page_size, 64) + self.assertEqual(args.page_size, 1) def test_declare_load_time_override_validates_whitelist(self): from sglang.srt.arg_groups.overrides import declare_load_time_override @@ -952,16 +958,14 @@ class TestPublishLifecycle(_IsolatedServerArgs): def test_declare_load_time_override_records_provenance(self): from sglang.srt.arg_groups.overrides import declare_load_time_override - from sglang.srt.server_args import ServerArgs - class _Args(_FakeResolvedArgs): - override = ServerArgs.override - - args = _Args(page_size=1) - get_context().set_server_args(args) + self._publish(page_size=1) declare_load_time_override("model.load_time", {"page_size": 64}) - self.assertEqual(args.page_size, 64) - self.assertIn(("model.load_time", {"page_size": 64}), args._resolved_overrides) + self.assertEqual(get_schedule().page_size, 64) + self.assertIn( + ("model.load_time", {"page_size": 64}), + get_context().overrides_log(), + ) if __name__ == "__main__": diff --git a/test/registered/unit/test_server_args_writer_ratchet.py b/test/registered/unit/test_server_args_writer_ratchet.py index 3082b5654..a151bfd2e 100644 --- a/test/registered/unit/test_server_args_writer_ratchet.py +++ b/test/registered/unit/test_server_args_writer_ratchet.py @@ -49,7 +49,7 @@ _EXCLUDED = ( "multimodal_gen", ) -_BASELINE = 49 +_BASELINE = 39 class TestServerArgsWriterRatchet(CustomTestCase):