From 11a4c2d05771303f02dc28c28af5e710ea3f5e1f Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Wed, 22 Jul 2026 01:18:05 -0700 Subject: [PATCH] config: read resolved config via namespace accessors (#31814) --- .../srt/batch_overlap/two_batch_overlap.py | 15 +- python/sglang/srt/configs/inkling.py | 4 +- .../sglang/srt/constrained/grammar_manager.py | 3 +- .../sglang/srt/disaggregation/common/conn.py | 9 +- python/sglang/srt/disaggregation/decode.py | 10 +- .../srt/disaggregation/encode_grpc_server.py | 7 +- .../srt/disaggregation/encode_server.py | 39 ++- .../srt/disaggregation/mooncake/conn.py | 11 +- python/sglang/srt/disaggregation/nixl/conn.py | 16 +- python/sglang/srt/disaggregation/prefill.py | 3 +- .../device_communicators/pymscclpp.py | 4 +- .../device_communicators/pynccl_allocator.py | 4 +- .../device_communicators/torch_symm_mem.py | 4 +- python/sglang/srt/dllm/mixin/scheduler.py | 5 +- python/sglang/srt/elastic_ep/elastic_ep.py | 2 +- .../srt/elastic_ep/expert_backup_client.py | 8 +- python/sglang/srt/entrypoints/engine.py | 17 +- python/sglang/srt/entrypoints/grpc_bridge.py | 19 +- python/sglang/srt/entrypoints/http_server.py | 17 +- .../srt/entrypoints/http_server_engine.py | 2 + .../entrypoints/openai/realtime/session.py | 5 +- python/sglang/srt/eplb/eplb_manager.py | 6 +- .../srt/eplb/expert_location_dispatch.py | 4 +- .../srt/eplb/expert_location_updater.py | 4 +- .../hardware_backend/mlx/model_runner_stub.py | 17 +- .../srt/hardware_backend/mlx/tp_worker.py | 21 +- .../musa/attention/flashattention_backend.py | 8 +- .../npu/attention/ascend_dsv4_backend.py | 8 +- .../npu/graph_runner/vit_npu_graph_runner.py | 4 +- .../srt/hardware_backend/npu/moe/fuseep.py | 8 +- python/sglang/srt/layers/activation.py | 8 +- .../srt/layers/attention/dsa/dsa_indexer.py | 20 +- .../srt/layers/attention/dsv4/indexer.py | 13 +- .../attention/flashattention_backend.py | 15 +- .../attention/flashinfer_mla_backend.py | 12 +- .../attention/hybrid_linear_attn_backend.py | 6 +- .../layers/attention/trtllm_mla_backend.py | 10 +- python/sglang/srt/layers/attention/vision.py | 15 +- .../srt/layers/attention/xpu_backend.py | 11 +- python/sglang/srt/layers/communicator.py | 22 +- python/sglang/srt/layers/cp/zigzag.py | 10 +- python/sglang/srt/layers/dcp/planner.py | 14 +- python/sglang/srt/layers/layernorm.py | 23 +- python/sglang/srt/layers/linear.py | 16 +- python/sglang/srt/layers/logits_processor.py | 8 +- python/sglang/srt/layers/moe/hash_topk.py | 8 +- .../moe/moe_runner/triton_utils/fused_moe.py | 4 +- .../triton_utils/fused_moe_triton_config.py | 6 +- .../layers/moe/token_dispatcher/flashinfer.py | 12 +- python/sglang/srt/layers/moe/topk.py | 16 +- .../srt/layers/quantization/fp8_utils.py | 9 +- .../sglang/srt/layers/quantization/mxfp4.py | 8 +- .../mxfp4_flashinfer_trtllm_moe.py | 8 +- .../srt/layers/rotary_embedding/base.py | 12 +- .../srt/layers/rotary_embedding/mrope.py | 7 +- python/sglang/srt/layers/sampler.py | 33 +-- .../srt/managers/data_parallel_controller.py | 5 +- python/sglang/srt/managers/mm_utils.py | 16 +- .../srt/managers/multi_tokenizer_mixin.py | 12 +- python/sglang/srt/managers/schedule_batch.py | 16 +- python/sglang/srt/managers/schedule_policy.py | 4 +- python/sglang/srt/managers/scheduler.py | 140 ++++----- .../batch_result_processor.py | 41 ++- .../managers/scheduler_components/dp_attn.py | 7 +- .../scheduler_components/load_inquirer.py | 3 +- .../logprob_result_processor.py | 13 +- .../scheduler_components/output_streamer.py | 17 +- .../scheduler_components/profiler_manager.py | 14 +- .../scheduler_components/request_receiver.py | 24 +- .../sglang/srt/managers/scheduler_pp_mixin.py | 5 +- .../srt/managers/tokenizer_control_mixin.py | 23 +- .../sglang/srt/managers/tokenizer_manager.py | 74 ++--- .../managers/tokenizer_manager_score_mixin.py | 5 +- python/sglang/srt/managers/tp_worker.py | 21 +- python/sglang/srt/mem_cache/allocation.py | 6 +- python/sglang/srt/mem_cache/common.py | 4 +- .../srt/mem_cache/deepseek_v4_memory_pool.py | 4 +- .../srt/mem_cache/kv_cache_configurator.py | 267 +++++++++--------- .../storage/lmcache/lmc_radix_cache.py | 4 +- .../srt/model_executor/forward_batch_info.py | 13 +- .../sglang/srt/model_executor/model_runner.py | 88 +++--- .../model_runner_components/misc_utils.py | 6 + .../remote_instance_weight_transporter.py | 5 +- python/sglang/srt/model_loader/loader.py | 9 +- python/sglang/srt/models/bailing_moe.py | 11 +- .../sglang/srt/models/bailing_moe_linear.py | 14 +- python/sglang/srt/models/bert.py | 6 +- .../attention_backend_handler.py | 6 +- .../attention_forward_methods/forward_mha.py | 14 +- .../attention_forward_methods/forward_mla.py | 12 +- python/sglang/srt/models/deepseek_nextn.py | 14 +- python/sglang/srt/models/deepseek_v2.py | 35 +-- python/sglang/srt/models/deepseek_v4.py | 31 +- python/sglang/srt/models/exaone_moe.py | 20 +- python/sglang/srt/models/gemma4_causal.py | 18 +- python/sglang/srt/models/gemma4_vision.py | 5 +- python/sglang/srt/models/glm4_moe.py | 7 +- python/sglang/srt/models/glm4_moe_lite.py | 9 +- .../sglang/srt/models/glm4_moe_lite_nextn.py | 6 +- python/sglang/srt/models/glm4_moe_nextn.py | 6 +- python/sglang/srt/models/glm4v.py | 4 +- python/sglang/srt/models/glm4v_moe.py | 6 +- python/sglang/srt/models/glm_image_vl.py | 9 +- python/sglang/srt/models/glm_ocr.py | 4 +- python/sglang/srt/models/glm_ocr_nextn.py | 4 +- python/sglang/srt/models/gpt_oss.py | 9 +- python/sglang/srt/models/inkling.py | 17 +- .../sglang/srt/models/inkling_common/attn.py | 4 +- .../srt/models/inkling_common/dense_mlp.py | 6 +- .../srt/models/inkling_common/kernels/comm.py | 10 +- .../sglang/srt/models/inkling_common/moe.py | 9 +- .../sglang/srt/models/inkling_common/sconv.py | 4 +- .../sglang/srt/models/inkling_common/util.py | 4 +- python/sglang/srt/models/internvl.py | 4 +- python/sglang/srt/models/kimi_k25.py | 4 +- python/sglang/srt/models/kimi_vl.py | 4 +- python/sglang/srt/models/laguna.py | 23 +- python/sglang/srt/models/llada2.py | 11 +- python/sglang/srt/models/llama_eagle3.py | 4 +- python/sglang/srt/models/mellum.py | 4 +- python/sglang/srt/models/mimo_audio.py | 4 +- python/sglang/srt/models/mimo_v2.py | 14 +- python/sglang/srt/models/mimo_vl.py | 4 +- python/sglang/srt/models/minimax_m2.py | 19 +- python/sglang/srt/models/minimax_m3.py | 13 +- python/sglang/srt/models/minimax_m3_vl.py | 13 +- python/sglang/srt/models/minimax_vl_common.py | 15 +- python/sglang/srt/models/mllama4.py | 6 +- python/sglang/srt/models/moss_vl.py | 9 +- python/sglang/srt/models/nemotron_h.py | 14 +- python/sglang/srt/models/qwen2.py | 13 +- python/sglang/srt/models/qwen2_5_vl.py | 4 +- python/sglang/srt/models/qwen2_moe.py | 19 +- python/sglang/srt/models/qwen3.py | 21 +- python/sglang/srt/models/qwen3_5.py | 10 +- python/sglang/srt/models/qwen3_5_mtp.py | 10 +- python/sglang/srt/models/qwen3_moe.py | 7 +- python/sglang/srt/models/qwen3_next_mtp.py | 11 +- python/sglang/srt/models/qwen3_vl.py | 22 +- python/sglang/srt/models/sarvam_moe.py | 12 +- python/sglang/srt/models/sdar.py | 13 +- python/sglang/srt/models/sdar_moe.py | 22 +- python/sglang/srt/models/step3p5.py | 16 +- python/sglang/srt/models/transformers.py | 10 +- python/sglang/srt/models/utils.py | 6 +- .../internvl_vit_cuda_graph_runner.py | 6 +- .../multimodal/processors/base_processor.py | 14 +- .../srt/multimodal/processors/kimi_k25.py | 7 +- .../srt/multimodal/vit_cuda_graph_runner.py | 4 +- .../srt/multiplex/multiplexing_mixin.py | 3 +- .../sglang/srt/speculative/dflash_info_v2.py | 4 +- .../srt/speculative/dflash_worker_v2.py | 7 +- python/sglang/srt/speculative/draft_utils.py | 5 +- .../dspark_components/dspark_planner.py | 15 +- .../dspark_components/dspark_worker_v2.py | 6 +- python/sglang/srt/speculative/eagle_info.py | 4 +- python/sglang/srt/speculative/eagle_utils.py | 23 +- .../sglang/srt/speculative/eagle_worker_v2.py | 73 +++-- .../speculative/frozen_kv_mtp_worker_v2.py | 7 +- python/sglang/srt/speculative/spec_utils.py | 4 +- .../sglang/srt/state_capturer/indexer_topk.py | 5 +- python/sglang/srt/utils/profile_utils.py | 4 +- 162 files changed, 1103 insertions(+), 1209 deletions(-) diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index ca44b0021..e2c410996 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -39,7 +39,12 @@ 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_exec, + 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 @@ -183,7 +188,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: @@ -335,7 +340,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): @@ -633,7 +638,7 @@ class TboForwardBatchPreparer: sum_field=None, ) _, child_b.extend_start_loc = compute_position( - get_server_args().attention_backend, + get_exec().kernel.attention_backend, child_b.extend_prefix_lens, child_b.extend_seq_lens, child_b.extend_num_tokens, @@ -832,7 +837,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/constrained/grammar_manager.py b/python/sglang/srt/constrained/grammar_manager.py index b039020fd..010185c25 100644 --- a/python/sglang/srt/constrained/grammar_manager.py +++ b/python/sglang/srt/constrained/grammar_manager.py @@ -14,6 +14,7 @@ from sglang.srt.constrained.base_grammar_backend import ( from sglang.srt.constrained.reasoner_grammar_backend import ReasonerGrammarObject from sglang.srt.distributed.communication_tags import P2PTag from sglang.srt.environ import envs +from sglang.srt.runtime_context import get_serving if TYPE_CHECKING: from sglang.srt.managers.io_struct import AbortReq @@ -28,7 +29,7 @@ class GrammarManager: self.scheduler = scheduler self.server_args = scheduler.server_args self.grammar_queue: List[Req] = [] - if not self.server_args.skip_tokenizer_init: + if not get_serving().skip_tokenizer_init: self.grammar_backend = create_grammar_backend( self.server_args, scheduler.tokenizer, diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 064165f46..f4a373ce5 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -32,11 +32,8 @@ from sglang.srt.disaggregation.utils import ( ) from sglang.srt.distributed import get_pp_group, get_world_group from sglang.srt.environ import envs -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.layers.dp_attention import get_attention_dp_rank, get_attention_dp_size +from sglang.srt.runtime_context import get_model, get_parallel, get_serving from sglang.srt.server_args import ServerArgs from sglang.srt.utils.network import ( NetworkAddress, @@ -634,7 +631,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 5224920c4..25c5d0db8 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -86,7 +86,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 from sglang.srt.utils.network import NetworkAddress from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method @@ -2151,7 +2151,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 @@ -2191,7 +2191,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 @@ -2203,9 +2203,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 875a80f2d..1b0a8af8c 100644 --- a/python/sglang/srt/disaggregation/encode_server.py +++ b/python/sglang/srt/disaggregation/encode_server.py @@ -59,14 +59,9 @@ from sglang.srt.model_loader import get_model from sglang.srt.multimodal.processors.qwen_vl import preprocess_video from sglang.srt.observability.metrics_collector import EncoderMetricsCollector from sglang.srt.observability.req_time_stats import EncoderReqTimeStats -from sglang.srt.observability.trace import ( - process_tracing_init, - trace_set_thread_info, -) -from sglang.srt.server_args import ( - PortArgs, - ServerArgs, -) +from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info +from sglang.srt.runtime_context import get_disagg, get_exec, get_mm +from sglang.srt.server_args import PortArgs, ServerArgs from sglang.srt.utils import ( add_prometheus_middleware, configure_logger, @@ -350,7 +345,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, ) @@ -368,15 +363,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() @@ -389,8 +384,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 ), ) @@ -399,7 +394,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 @@ -413,12 +408,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 @@ -1688,7 +1683,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]) @@ -1784,7 +1779,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: @@ -1855,7 +1850,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) @@ -1887,11 +1882,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 7d306a4d0..2ba2323c2 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -55,6 +55,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 @@ -314,9 +315,7 @@ class MooncakeKVManager(CommonKVManager): self.kv_buffer_tensors = None def _handle_staging_req(self, msg): - from sglang.srt.disaggregation.common.staging_handler import ( - handle_staging_req, - ) + from sglang.srt.disaggregation.common.staging_handler import handle_staging_req room = int(msg[1].decode("ascii")) session_id = msg[4].decode("ascii") @@ -350,9 +349,7 @@ class MooncakeKVManager(CommonKVManager): def _is_watermark_ready( self, session_id: str, alloc_round: int, alloc_end: int ) -> bool: - from sglang.srt.disaggregation.common.staging_handler import ( - is_watermark_ready, - ) + from sglang.srt.disaggregation.common.staging_handler import is_watermark_ready return is_watermark_ready(self._staging_ctx, session_id, alloc_round, alloc_end) @@ -469,7 +466,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 f60a8766a..c1f8b31ce 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -13,6 +13,8 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple import numpy as np import numpy.typing as npt +from sglang.srt.runtime_context import get_schedule + if TYPE_CHECKING: from sglang.srt.disaggregation.common.staging_handler import StagingTransferInfo @@ -535,9 +537,7 @@ class NixlKVManager(CommonKVManager): def _is_watermark_ready( self, agent_name: str, alloc_round: int, alloc_end: int ) -> bool: - from sglang.srt.disaggregation.common.staging_handler import ( - is_watermark_ready, - ) + from sglang.srt.disaggregation.common.staging_handler import is_watermark_ready return is_watermark_ready(self._staging_ctx, agent_name, alloc_round, alloc_end) @@ -558,9 +558,7 @@ class NixlKVManager(CommonKVManager): threading.Thread(target=decode_staging_thread, daemon=True).start() def _handle_staging_req(self, msg): - from sglang.srt.disaggregation.common.staging_handler import ( - handle_staging_req, - ) + from sglang.srt.disaggregation.common.staging_handler import handle_staging_req room = int(msg[1].decode("ascii")) session_id = msg[4].decode("ascii") @@ -625,7 +623,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, ) @@ -1739,9 +1737,7 @@ class NixlKVManager(CommonKVManager): req, page_start, num_pages, session_id=req.agent_name ) if not ready: - from sglang.srt.disaggregation.common.staging_buffer import ( - StagingAllocator, - ) + from sglang.srt.disaggregation.common.staging_buffer import StagingAllocator if c_offset == StagingAllocator.ALLOC_OVERSIZED: raise RuntimeError( diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 1f4c947fa..841e23a7f 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -64,6 +64,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.nvtx_utils import scheduler_nvtx_method if TYPE_CHECKING: @@ -1181,7 +1182,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/elastic_ep.py b/python/sglang/srt/elastic_ep/elastic_ep.py index 8ceb53a9e..11886fdc6 100644 --- a/python/sglang/srt/elastic_ep/elastic_ep.py +++ b/python/sglang/srt/elastic_ep/elastic_ep.py @@ -442,7 +442,7 @@ def get_healthy_expert_location_src_rank( *, invoked_in_elastic_ep_rejoin_path: bool ) -> int: world_group = get_world_group() - # NOTE: do not key off `self.server_args.elastic_ep_rejoin` here. + # NOTE: do not key off `get_exec().moe.elastic_ep_rejoin` here. # A rank that was started as a rejoin rank may later act as a healthy # rank in a subsequent recovery cycle. local_rejoin_flag = bool(invoked_in_elastic_ep_rejoin_path) diff --git a/python/sglang/srt/elastic_ep/expert_backup_client.py b/python/sglang/srt/elastic_ep/expert_backup_client.py index 8b77f7f07..7a20e3057 100644 --- a/python/sglang/srt/elastic_ep/expert_backup_client.py +++ b/python/sglang/srt/elastic_ep/expert_backup_client.py @@ -7,13 +7,11 @@ from typing import Any, Callable import torch import zmq -from sglang.srt.distributed.parallel_state import ( - get_world_group, - get_world_size, -) +from sglang.srt.distributed.parallel_state import get_world_group, get_world_size 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 +109,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/engine.py b/python/sglang/srt/entrypoints/engine.py index 5c158535e..d9975d981 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -877,7 +877,14 @@ class Engine(EngineScoreMixin, EngineBase): server_args, port_args ) else: - # Launch multi-tokenizer router + # Launch multi-tokenizer router. Unlike TokenizerManager, the router + # does not publish; but it runs in this parent process and reads + # resolved config through the namespace accessors (e.g. get_parallel() + # for routed_dp_rank), so publish here. The child TokenizerWorkers + # publish independently in their own processes. + from sglang.srt.runtime_context import publish + + publish(server_args, role="tokenizer") tokenizer_manager = MultiTokenizerRouter(server_args, port_args) template_manager = None @@ -997,12 +1004,18 @@ class Engine(EngineScoreMixin, EngineBase): ) def get_server_info(self): + from sglang.srt.runtime_context import get_context + internal_states = self.loop.run_until_complete( self.tokenizer_manager.get_internal_state() ) return msgspec_to_builtins( { - **dataclasses.asdict(self.tokenizer_manager.server_args), + # Overlay post-publish overrides so the report reflects current + # config (weight version, model path, runtime tunables). + **get_context().resolved_server_args_dict( + base=dataclasses.asdict(self.tokenizer_manager.server_args) + ), **self._scheduler_init_result.scheduler_infos[0], "internal_states": internal_states, "version": __version__, diff --git a/python/sglang/srt/entrypoints/grpc_bridge.py b/python/sglang/srt/entrypoints/grpc_bridge.py index fa0ab61b8..eca22c3b8 100644 --- a/python/sglang/srt/entrypoints/grpc_bridge.py +++ b/python/sglang/srt/entrypoints/grpc_bridge.py @@ -15,6 +15,7 @@ from typing import Any, Awaitable, Callable, Dict, List, Optional from pydantic import ValidationError +from sglang.srt.runtime_context import get_context, get_lora, get_serving from sglang.srt.utils.msgspec_utils import msgspec_to_builtins logger = logging.getLogger(__name__) @@ -229,9 +230,7 @@ class RuntimeHandle: return self._openai_serving_classes from sglang.srt.entrypoints.openai.serving_chat import OpenAIServingChat - from sglang.srt.entrypoints.openai.serving_classify import ( - OpenAIServingClassify, - ) + from sglang.srt.entrypoints.openai.serving_classify import OpenAIServingClassify from sglang.srt.entrypoints.openai.serving_completions import ( OpenAIServingCompletion, ) @@ -376,16 +375,20 @@ 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": get_serving().tokenizer_path, "is_generation": self.tokenizer_manager.is_generation, - "weight_version": self.server_args.weight_version, + "weight_version": get_serving().weight_version, "model_type": getattr(model_config.hf_config, "model_type", None), "architectures": getattr(model_config.hf_config, "architectures", None), } return json.dumps(result, default=str) def get_server_info(self) -> str: - result: Dict[str, Any] = dataclasses.asdict(self.server_args) + # Overlay post-publish overrides (weight version, model path, runtime + # tunables) so the report reflects current config, not the startup record. + result: Dict[str, Any] = get_context().resolved_server_args_dict( + base=dataclasses.asdict(self.server_args) + ) result.update(self.scheduler_info) return json.dumps(msgspec_to_builtins(result), default=str) @@ -424,9 +427,7 @@ class RuntimeHandle: "max_model_len": self.tokenizer_manager.model_config.context_len, } ] - if self.server_args.enable_lora and hasattr( - self.tokenizer_manager, "lora_registry" - ): + if get_lora().enable_lora and hasattr(self.tokenizer_manager, "lora_registry"): lora_registry = self.tokenizer_manager.lora_registry for _, lora_ref in lora_registry.get_all_adapters().items(): models.append( diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index 9a1aa72c5..ea0216fcf 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -693,18 +693,19 @@ async def get_model_info(): @app.get("/model_info") async def model_info(): """Get the model information.""" + from sglang.srt.runtime_context import get_serving + model_config = _global_state.tokenizer_manager.model_config result = { "model_path": _global_state.tokenizer_manager.model_path, "tokenizer_path": _global_state.tokenizer_manager.server_args.tokenizer_path, "is_generation": _global_state.tokenizer_manager.is_generation, "preferred_sampling_params": _global_state.tokenizer_manager.server_args.preferred_sampling_params, - "weight_version": _global_state.tokenizer_manager.server_args.weight_version, + "weight_version": get_serving().weight_version, "has_image_understanding": model_config.is_image_understandable_model, "has_audio_understanding": model_config.is_audio_understandable_model, "model_type": getattr(model_config.hf_config, "model_type", None), "architectures": getattr(model_config.hf_config, "architectures", None), - "weight_version": _global_state.tokenizer_manager.server_args.weight_version, # "hf_config": model_config.hf_config.to_dict(), } return result @@ -738,12 +739,18 @@ async def server_info(): await _global_state.tokenizer_manager.get_internal_state() ) + from sglang.srt.runtime_context import get_context + server_args = _global_state.tokenizer_manager.server_args # server_args.model_config is not serializable but should be excluded by asdict. + # Overlay post-publish overrides so runtime updates (weight version, model + # path/load format) are reported, not the startup record. return msgspec_to_builtins( { - **dataclasses.asdict(server_args), + **get_context().resolved_server_args_dict( + base=dataclasses.asdict(server_args) + ), **_global_state.scheduler_info, "internal_states": internal_states, "version": __version__, @@ -1368,7 +1375,9 @@ async def update_weight_version( # since weight_version update is a simple operation that doesn't affect model weights try: # Update the weight version in server args (the single source of truth) - _global_state.tokenizer_manager.server_args.override( + from sglang.srt.runtime_context import get_context + + get_context().override( "http.update_weight_version", weight_version=obj.new_version ) diff --git a/python/sglang/srt/entrypoints/http_server_engine.py b/python/sglang/srt/entrypoints/http_server_engine.py index 4a4996743..02705a847 100644 --- a/python/sglang/srt/entrypoints/http_server_engine.py +++ b/python/sglang/srt/entrypoints/http_server_engine.py @@ -55,6 +55,8 @@ class HttpServerEngineAdapter(EngineBase): def __init__(self, **kwargs): self.server_args = ServerArgs(**kwargs) + # Read host/port from the adapter's own args: no config is published yet + # in this process (publish happens in the child from launch_server_process). print( f"Launch HttpServerEngineAdapter at: {self.server_args.host}:{self.server_args.port}" ) diff --git a/python/sglang/srt/entrypoints/openai/realtime/session.py b/python/sglang/srt/entrypoints/openai/realtime/session.py index c5951993e..eb122f21a 100644 --- a/python/sglang/srt/entrypoints/openai/realtime/session.py +++ b/python/sglang/srt/entrypoints/openai/realtime/session.py @@ -72,6 +72,7 @@ from sglang.srt.entrypoints.openai.transcription_adapters.base import ( TranscriptionAdapter, ) from sglang.srt.managers.tokenizer_manager import TokenizerManager +from sglang.srt.runtime_context import get_serving from sglang.srt.server_args import ServerArgs from sglang.srt.utils import random_uuid @@ -338,12 +339,12 @@ class RealtimeConnection: if ( transcription is not None and transcription.model - and transcription.model != self.server_args.served_model_name + and transcription.model != get_serving().served_model_name ): await self._send_error( "not_supported", f"Model {transcription.model!r} is not served by this endpoint " - f"(serving {self.server_args.served_model_name!r}); set " + f"(serving {get_serving().served_model_name!r}); set " f"transcription.model to null or to the server's model name.", param="session.audio.input.transcription.model", ) diff --git a/python/sglang/srt/eplb/eplb_manager.py b/python/sglang/srt/eplb/eplb_manager.py index 360ac93b7..b2b1db792 100644 --- a/python/sglang/srt/eplb/eplb_manager.py +++ b/python/sglang/srt/eplb/eplb_manager.py @@ -16,7 +16,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 @@ -274,8 +274,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 bf1890a5a..ba95644c8 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 @@ -34,7 +34,7 @@ class ExpertLocationDispatchInfo: @classmethod def init_new(cls, layer_id: int): - ep_dispatch_algorithm = get_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 diff --git a/python/sglang/srt/eplb/expert_location_updater.py b/python/sglang/srt/eplb/expert_location_updater.py index 7873223f0..a5ba50923 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__) @@ -107,7 +107,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 84f58eaeb..455a8715d 100644 --- a/python/sglang/srt/hardware_backend/mlx/model_runner_stub.py +++ b/python/sglang/srt/hardware_backend/mlx/model_runner_stub.py @@ -16,9 +16,8 @@ from sglang.srt.hardware_backend.mlx.kv_cache.auxiliary_state import ( from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator from sglang.srt.mem_cache.memory_pool import KVCache, ReqToTokenPool from sglang.srt.model_executor.model_runner import ModelRunner -from sglang.srt.model_executor.model_runner_components.layer_setup import ( - ModelLayerInfo, -) +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,7 +143,7 @@ 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 @@ -165,7 +164,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) @@ -173,7 +172,7 @@ class MlxModelRunnerStub(ModelRunner): requested_per_worker = requested // self.dp_size resolved = min(requested_per_worker, capacity_cap) - aux_state_size = self.server_args.max_mamba_cache_size + aux_state_size = get_schedule().max_mamba_cache_size if ( mambaish_config(self.model_config) is not None and aux_state_size is not None @@ -209,7 +208,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) @@ -241,7 +240,7 @@ class MlxModelRunnerStub(ModelRunner): # Create minimal pools if mambaish_config(self.model_config) is not None: - auxiliary_state_size = self.server_args.max_mamba_cache_size + auxiliary_state_size = get_schedule().max_mamba_cache_size if auxiliary_state_size is None: auxiliary_state_size = ( self.max_running_requests * self._aux_state_slots_per_request() @@ -255,7 +254,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..cc65a61b2 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__) @@ -47,25 +48,23 @@ class MlxTpModelWorker(TpModelWorker): def _init_model_runner(self): """Create MLX runner first (auto-sizes pool), then stub with matching size.""" from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner - from sglang.srt.hardware_backend.mlx.model_runner_stub import ( - MlxModelRunnerStub, - ) + from sglang.srt.hardware_backend.mlx.model_runner_stub import MlxModelRunnerStub 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..c48aff4c9 100644 --- a/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py +++ b/python/sglang/srt/hardware_backend/musa/attention/flashattention_backend.py @@ -19,11 +19,9 @@ from sglang.srt.layers.attention.flashattention_backend import ( merge_state_v2_wrapper, ) from sglang.srt.layers.radix_attention import AttentionType -from sglang.srt.layers.utils.cp_utils import ( - cp_allgather_and_save_kv_cache, -) +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 +513,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_dsv4_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_dsv4_backend.py index d45ec4fda..971315be2 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 @@ -12,7 +12,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 @@ -1362,9 +1362,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 ) @@ -1409,9 +1408,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 941c84685..ad4fc3e87 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 a7db9d27a..36fa364c3 100644 --- a/python/sglang/srt/layers/activation.py +++ b/python/sglang/srt/layers/activation.py @@ -22,9 +22,7 @@ import torch.nn as nn import torch.nn.functional as F from transformers import PretrainedConfig -from sglang.srt.distributed import ( - divide, -) +from sglang.srt.distributed import divide from sglang.srt.environ import envs from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.utils import MultiPlatformOp @@ -33,7 +31,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, @@ -89,7 +87,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/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index 196639f18..eb34874a9 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -37,10 +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.state_capturer.indexer_topk import ( - maybe_capture_indexer_topk, +from sglang.srt.runtime_context import ( + get_device, + get_exec, + get_parallel, + get_schedule, + get_server_args, ) +from sglang.srt.state_capturer.indexer_topk import maybe_capture_indexer_topk from sglang.srt.utils import ( add_prefix, ceil_align, @@ -105,9 +109,7 @@ if is_npu(): import torch_npu from sglang.srt.hardware_backend.npu.utils import get_indexer_weight_stream -from sglang.srt.distributed import ( - get_attn_tp_group, -) +from sglang.srt.distributed import get_attn_tp_group from sglang.srt.distributed.parallel_state import get_pp_group from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers.communicator import ScatterMode @@ -458,7 +460,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 @@ -469,7 +471,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 @@ -1055,7 +1057,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/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index 0ecd103ae..30c88f03a 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -28,16 +28,14 @@ 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 if TYPE_CHECKING: from sglang.srt.layers.attention.base_attn_backend import AttentionBackend - from sglang.srt.layers.attention.dsv4.compressor import ( - CompressorBackendMixin, - ) + from sglang.srt.layers.attention.dsv4.compressor import CompressorBackendMixin from sglang.srt.layers.quantization import QuantizationConfig from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch @@ -129,9 +127,7 @@ def _aiter_fp8_paged_mqa_logits( clean_logits: bool = False, ) -> torch.Tensor: """Wrapper adapting aiter's deepgemm_fp8_paged_mqa_logits to SGLang's interface.""" - from aiter.ops.triton.attention.pa_mqa_logits import ( - deepgemm_fp8_paged_mqa_logits, - ) + from aiter.ops.triton.attention.pa_mqa_logits import deepgemm_fp8_paged_mqa_logits batch_size = q_fp8.shape[0] next_n = q_fp8.shape[1] @@ -838,9 +834,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 3036fa619..292ca9bf0 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -13,9 +13,7 @@ from sglang.kernels.ops.attention.metadata import ( ) from sglang.kernels.ops.attention.pa_page_table import _build_pa_page_table from sglang.kernels.ops.attention.utils import assert_buffer_fits -from sglang.kernels.ops.kvcache.trtllm_mha_page_table import ( - build_trtllm_mha_page_table, -) +from sglang.kernels.ops.kvcache.trtllm_mha_page_table import build_trtllm_mha_page_table from sglang.srt.configs.model_config import AttentionArch from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.cp.base import CPAttentionBackendKind, get_cp_strategy @@ -28,7 +26,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 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 @@ -166,9 +164,12 @@ class FlashAttentionBackend(AttentionBackend): self.token_to_kv_pool = model_runner.token_to_kv_pool self.req_to_token = model_runner.req_to_token_pool.req_to_token self.kv_cache_dtype = model_runner.kv_cache_dtype - from sglang.srt.runtime_context import get_model - self.kv_cache_dtype_str = get_model().kv_cache_dtype + self.kv_cache_dtype_str = getattr( + model_runner, + "kv_cache_dtype_str", + model_runner.server_args.kv_cache_dtype, + ) self.kv_cache_is_mxfp8 = self.kv_cache_dtype_str == "mxfp8" self.page_size = model_runner.page_size # Static page-table width (upper bound). The device-side page-table build @@ -1479,7 +1480,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 0d68b303c..27599bfca 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. @@ -32,7 +32,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, @@ -223,9 +223,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 @@ -401,7 +401,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/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index 6e9e57adf..8fbeb9ff6 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -19,7 +19,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 @@ -350,7 +350,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) @@ -764,7 +764,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/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index 14fba1e25..82650ea99 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -38,7 +38,11 @@ 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, +) from sglang.srt.utils import is_flashinfer_available, is_float4_e2m1fn_x2 if is_flashinfer_available(): @@ -197,9 +201,7 @@ 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.cuda_graph_custom_mask = None diff --git a/python/sglang/srt/layers/attention/vision.py b/python/sglang/srt/layers/attention/vision.py index 1b94d92a0..a6b653633 100644 --- a/python/sglang/srt/layers/attention/vision.py +++ b/python/sglang/srt/layers/attention/vision.py @@ -15,7 +15,7 @@ from einops import rearrange from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm as can_use_jit_qk_norm 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, @@ -69,9 +69,7 @@ if _is_npu: if _is_xpu: from sgl_kernel.flash_attn import flash_attn_varlen_func -from sglang.kernels.ops.attention.prefill_attention import ( - context_attention_fwd, -) +from sglang.kernels.ops.attention.prefill_attention import context_attention_fwd from sglang.srt.distributed import ( split_tensor_along_last_dim, tensor_model_parallel_all_gather, @@ -86,7 +84,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 @@ -1045,7 +1042,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.") @@ -1124,7 +1121,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( @@ -1152,7 +1149,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: @@ -1257,7 +1254,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/xpu_backend.py b/python/sglang/srt/layers/attention/xpu_backend.py index 27ad57ce5..60cc92d8c 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 if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention @@ -69,9 +69,12 @@ class XPUAttentionBackend(AttentionBackend): self.token_to_kv_pool = model_runner.token_to_kv_pool self.req_to_token = model_runner.req_to_token_pool.req_to_token self.kv_cache_dtype = model_runner.kv_cache_dtype - from sglang.srt.runtime_context import get_model - self.kv_cache_dtype_str = get_model().kv_cache_dtype + self.kv_cache_dtype_str = getattr( + model_runner, + "kv_cache_dtype_str", + model_runner.server_args.kv_cache_dtype, + ) self.page_size = model_runner.page_size self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA self.skip_prefill = skip_prefill @@ -640,7 +643,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 7a5157dba..774ccdcf2 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -72,7 +72,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, @@ -170,7 +176,7 @@ def apply_flashinfer_allreduce_fusion(batch_size: int): and batch_size > 0 and batch_size <= FUSE_ALLREDUCE_MAX_BATCH_SIZE 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() ) @@ -186,7 +192,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 ) @@ -274,7 +280,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: @@ -407,7 +413,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 @@ -467,7 +473,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): @@ -815,7 +821,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) @@ -1120,7 +1126,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 b7fe868d6..06f3bda8b 100644 --- a/python/sglang/srt/layers/cp/zigzag.py +++ b/python/sglang/srt/layers/cp/zigzag.py @@ -48,12 +48,10 @@ from sglang.srt.layers.cp.base import ( CPAttentionBackendKind, ) from sglang.srt.layers.cp.padding import pad_local_rows -from sglang.srt.layers.dp_attention import ( - is_allocation_symmetric, -) +from sglang.srt.layers.dp_attention import is_allocation_symmetric 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 +206,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 8b0ca5141..aa151b720 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -31,7 +31,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,9 +130,7 @@ if _is_cuda: # BEFORE the weight multiply, so the multiply is done in the narrow dtype. _jit_rmsnorm_hf_available = False try: - from sglang.jit_kernel.rmsnorm_hf import ( - is_supported_rmsnorm_hf_hidden_size, - ) + from sglang.jit_kernel.rmsnorm_hf import is_supported_rmsnorm_hf_hidden_size from sglang.jit_kernel.rmsnorm_hf import rmsnorm_hf as _jit_rmsnorm_hf _jit_rmsnorm_hf_available = True @@ -144,9 +142,7 @@ if _is_cuda: _jit_rmsnorm_hf = None from sglang.jit_kernel.norm import fused_add_rmsnorm as _jit_fused_add_rmsnorm - from sglang.jit_kernel.norm import ( - is_supported_jit_fused_add_rmsnorm_hidden_size, - ) + from sglang.jit_kernel.norm import is_supported_jit_fused_add_rmsnorm_hidden_size logger = logging.getLogger(__name__) @@ -206,7 +202,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) @@ -284,7 +280,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( @@ -391,7 +387,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) @@ -452,7 +448,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( @@ -579,7 +575,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 57aae27ba..674ea2c65 100644 --- a/python/sglang/srt/layers/linear.py +++ b/python/sglang/srt/layers/linear.py @@ -25,9 +25,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) from sglang.srt.environ import envs -from sglang.srt.layers.dp_attention import ( - is_allocation_symmetric, -) +from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.moe.utils import should_skip_mlp_all_reduce from sglang.srt.layers.parameter import ( BasevLLMParameter, @@ -39,7 +37,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: @@ -759,9 +757,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear): shard_offsets.append((i, current_shard_offset, output_size)) current_shard_offset += output_size if _is_cpu: - from sglang.srt.model_loader.weight_utils import ( - pad_loaded_weight, - ) + from sglang.srt.model_loader.weight_utils import pad_loaded_weight loaded_weight = pad_loaded_weight( loaded_weight, param.output_dim, output_sizes @@ -805,9 +801,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear): current_block_offset += shard_block_size if _is_cpu: - from sglang.srt.model_loader.weight_utils import ( - pad_loaded_weight, - ) + from sglang.srt.model_loader.weight_utils import pad_loaded_weight loaded_weight = pad_loaded_weight( loaded_weight, param.output_dim, shard_block_sizes @@ -1596,7 +1590,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 cf55b97cf..a377a6054 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -47,7 +47,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, @@ -346,7 +346,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 = ( @@ -370,8 +370,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/hash_topk.py b/python/sglang/srt/layers/moe/hash_topk.py index 42a2caa90..5b371ac97 100644 --- a/python/sglang/srt/layers/moe/hash_topk.py +++ b/python/sglang/srt/layers/moe/hash_topk.py @@ -7,9 +7,7 @@ import torch from torch import nn from sglang.srt.environ import envs -from sglang.srt.eplb.expert_distribution import ( - get_global_expert_distribution_recorder, -) +from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location_dispatch import ( ExpertLocationDispatchInfo, topk_ids_logical_to_physical, @@ -22,6 +20,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 +43,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 3160467e4..87f0e1040 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, @@ -506,7 +506,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..8e08b5b81 100644 --- a/python/sglang/srt/layers/moe/token_dispatcher/flashinfer.py +++ b/python/sglang/srt/layers/moe/token_dispatcher/flashinfer.py @@ -21,13 +21,9 @@ from sglang.srt.layers.moe.token_dispatcher import ( from sglang.srt.layers.moe.token_dispatcher.flashinfer_utils import ( TorchDistributedCommBackend, ) -from sglang.srt.layers.moe.topk import ( - StandardTopKOutput, - TopKOutput, - TopKOutputChecker, -) +from sglang.srt.layers.moe.topk import StandardTopKOutput, TopKOutput, 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 +119,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 +128,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 82a0ffe15..bfcf92d60 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 @@ -83,9 +83,7 @@ except ImportError: pass from sglang.jit_kernel.dsv4 import mask_topk_ids -from sglang.srt.distributed import ( - get_tp_group, -) +from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) @@ -98,9 +96,7 @@ from sglang.srt.eplb.expert_location_dispatch import ( ) from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.moe import get_moe_runner_backend -from sglang.srt.layers.moe.utils import ( - has_per_rank_fused_shared_slots, -) +from sglang.srt.layers.moe.utils import has_per_rank_fused_shared_slots from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer from sglang.srt.utils import ( @@ -419,10 +415,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 @@ -496,9 +491,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 = ( diff --git a/python/sglang/srt/layers/quantization/fp8_utils.py b/python/sglang/srt/layers/quantization/fp8_utils.py index de940c480..f8686a382 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, @@ -1470,9 +1469,7 @@ def requant_block_scale_ue8m0_for_deepgemm( scales are not already UE8M0, and DeepGEMM can run the layer (bf16 output, aligned shape). Returns True when it requantizes. """ - from sglang.srt.model_loader.utils import ( - should_deepgemm_weight_requant_ue8m0, - ) + from sglang.srt.model_loader.utils import should_deepgemm_weight_requant_ue8m0 if ( not use_deepgemm_runner @@ -1794,7 +1791,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 822998fc2..458730649 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, @@ -77,9 +77,7 @@ if is_flashinfer_available(): nvfp4_block_scale_interleave, trtllm_fp4_block_scale_moe, ) - from flashinfer.fused_moe.core import ( - get_w2_permute_indices_with_cache, - ) + from flashinfer.fused_moe.core import get_w2_permute_indices_with_cache # SM90 mixed-input helpers landed in FlashInfer #3084 (post-0.6.10). Older # versions don't ship them; gate at import so unrelated code paths still load. @@ -334,7 +332,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 two 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 36f79d628..89cf50f74 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): @@ -376,9 +376,7 @@ def maybe_fuse_routed_scale_and_shared_add( from sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe import ( Mxfp4FlashinferCutlassMoEMethod, ) - from sglang.srt.layers.quantization.mxfp4_marlin_moe import ( - Mxfp4MarlinMoEMethod, - ) + from sglang.srt.layers.quantization.mxfp4_marlin_moe import Mxfp4MarlinMoEMethod fused = isinstance( experts.quant_method, diff --git a/python/sglang/srt/layers/rotary_embedding/base.py b/python/sglang/srt/layers/rotary_embedding/base.py index 2f68e4965..67abc2ca7 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, @@ -65,9 +65,7 @@ if _is_npu: ) if _is_hip: - from sglang.kernels.ops.attention.utils import ( - fused_qk_rope_reshape_and_cache, - ) + from sglang.kernels.ops.attention.utils import fused_qk_rope_reshape_and_cache if _is_xpu: from sgl_kernel import fused_qk_rope_with_cos_sin_cache_inplace @@ -127,7 +125,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, @@ -151,7 +149,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 @@ -162,7 +160,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 979b9741b..7ed800dfa 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 from sglang.srt.utils import ( cpu_has_amx_support, is_cuda, @@ -42,7 +42,6 @@ if _is_xpu: from sgl_kernel import multimodal_rotary_embedding from sglang.kernels.ops.attention.mrope import apply_interleaved_rope_triton -from sglang.srt.runtime_context import get_server_args def apply_interleaved_rope(x: torch.Tensor, mrope_section: list) -> torch.Tensor: @@ -132,7 +131,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): @@ -144,7 +143,7 @@ class MRotaryEmbedding(RotaryEmbedding): last_dim = cos_sin.size()[-1] cos, sin = cos_sin.chunk(2, dim=-1) if self.mrope_interleaved: - if support_triton(get_server_args().attention_backend): + if support_triton(get_exec().kernel.attention_backend): cos = apply_interleaved_rope_triton(cos, self.mrope_section) sin = apply_interleaved_rope_triton(sin, self.mrope_section) else: diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index de587ec5f..97868104b 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -8,34 +8,21 @@ from torch import nn from sglang.kernels.ops.sampling.murmur_hash import murmur_hash32 from sglang.srt.distributed import get_tp_group -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled 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.layers.logprob_processor import OutputLogprobProcessor +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 -from sglang.srt.utils.common import ( - get_bool_env_var, - is_cuda, - is_hip, - is_musa, - is_npu, -) +from sglang.srt.utils.common import get_bool_env_var, is_cuda, is_hip, is_musa, is_npu if is_cuda(): from flashinfer.sampling import ( min_p_sampling_from_probs, top_k_top_p_sampling_from_probs, ) - from sgl_kernel import ( - top_k_renorm_prob, - top_p_renorm_prob, - ) + from sgl_kernel import top_k_renorm_prob, top_p_renorm_prob if is_musa(): from sgl_kernel import ( @@ -74,12 +61,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() @@ -245,7 +234,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 diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index 35efe50d8..3457dd509 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -48,6 +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 get_exec from sglang.srt.server_args import ( DP_ATTENTION_HANDSHAKE_PORT_DELTA, PortArgs, @@ -231,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; " @@ -484,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 891d8df6c..cdf99935f 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -33,7 +33,13 @@ 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, + get_serving, +) 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 @@ -878,7 +884,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" @@ -1287,7 +1293,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 ) @@ -1967,7 +1973,7 @@ def wrap_shm_features(obj): """ Scan the object for multimodal tensors and wrap them in SHM pointers. """ - if _get_is_default_transport() or get_server_args().skip_tokenizer_init: + if _get_is_default_transport() or get_serving().skip_tokenizer_init: return obj if obj.mm_inputs: @@ -2028,7 +2034,7 @@ def unwrap_shm_features(obj): Restore ShmPointerMMData wrappers back into standard torch.Tensors. Handles both single requests and batch requests. """ - if _get_is_default_transport() or get_server_args().skip_tokenizer_init: + if _get_is_default_transport() or get_serving().skip_tokenizer_init: return obj # Handle batch requests if isinstance(obj, BaseBatchReq): diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py index 8a1ba1839..3381ad2ba 100644 --- a/python/sglang/srt/managers/multi_tokenizer_mixin.py +++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py @@ -1,5 +1,7 @@ from __future__ import annotations +from sglang.srt.runtime_context import get_disagg + # Copyright 2023-2024 SGLang Team # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. @@ -645,15 +647,15 @@ class TokenizerWorker(TokenizerManager): self.tokenizer_ipc_name = port_args.tokenizer_ipc_name # For PD disaggregtion - self.server_args.override( + from sglang.srt.runtime_context import get_context + + get_context().override( "tokenizer_worker.restore_disaggregation_mode", disaggregation_mode=disaggregation_mode, ) - self.disaggregation_mode = DisaggregationMode( - self.server_args.disaggregation_mode - ) + self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode) self.disaggregation_transfer_backend = TransferBackend( - self.server_args.disaggregation_transfer_backend + get_disagg().disaggregation_transfer_backend ) # Register this worker with the router for pause/continue broadcasting diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 97a3dceb6..1f11d03ed 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -77,10 +77,7 @@ from sglang.srt.managers.embed_types import PositionalEmbeds from sglang.srt.managers.scheduler_components.new_token_ratio_tracker import ( NewTokenRatioTracker, ) -from sglang.srt.mem_cache.allocation import ( - alloc_for_decode, - alloc_for_extend, -) +from sglang.srt.mem_cache.allocation import alloc_for_decode, alloc_for_extend from sglang.srt.mem_cache.allocation_sizing import get_alloc_reserve_per_decode from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.base_prefix_cache import ( @@ -105,7 +102,12 @@ from sglang.srt.observability.req_time_stats import ( DPControllerReqTimeStats, SchedulerReqTimeStats, ) -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_parallel, + get_server_args, + get_serving, + get_spec, +) from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.server_args import ServerArgs @@ -1094,7 +1096,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 @@ -1115,7 +1117,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 diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index fa929fd38..4cb347de4 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -56,7 +56,7 @@ 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.runtime_context import get_disagg from sglang.srt.server_args import ServerArgs if TYPE_CHECKING: @@ -193,7 +193,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 0b96a4de5..0dfdc8207 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -210,9 +210,7 @@ from sglang.srt.managers.scheduler_components.pool_stats_observer import ( from sglang.srt.managers.scheduler_components.profiler_manager import ( SchedulerProfilerManager, ) -from sglang.srt.managers.scheduler_components.recv_skipper import ( - SchedulerRecvSkipper, -) +from sglang.srt.managers.scheduler_components.recv_skipper import SchedulerRecvSkipper from sglang.srt.managers.scheduler_components.request_receiver import ( SchedulerRequestReceiver, ) @@ -241,7 +239,20 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa from sglang.srt.parser.reasoning_parser import ReasoningParser from sglang.srt.platforms import current_platform from sglang.srt.plugins import load_plugins -from sglang.srt.runtime_context import get_context, get_parallel +from sglang.srt.runtime_context import ( + get_context, + get_device, + get_disagg, + get_exec, + get_lora, + get_memory, + get_mm, + get_observability, + get_parallel, + get_schedule, + get_serving, + get_spec, +) from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_params import TOP_K_ALL from sglang.srt.server_args import PortArgs, ServerArgs @@ -443,9 +454,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 @@ -471,8 +482,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, @@ -583,7 +594,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 ) @@ -611,11 +622,11 @@ class Scheduler( self.ipc_channels = SchedulerIpcChannels.create( 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 + skip_tokenizer_init=get_serving().skip_tokenizer_init, + 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(), ) @@ -631,7 +642,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) @@ -641,7 +652,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=[ @@ -712,9 +723,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, ) @@ -785,7 +796,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 @@ -793,10 +804,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) @@ -887,7 +898,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(), ) @@ -933,14 +944,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. @@ -987,7 +998,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 ) @@ -1007,13 +1018,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: @@ -1049,8 +1059,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)." @@ -1067,15 +1077,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( @@ -1091,12 +1101,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 @@ -1118,11 +1128,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? @@ -1192,10 +1200,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, ) @@ -1221,7 +1229,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, @@ -1235,11 +1243,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, @@ -1320,7 +1327,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 @@ -1329,7 +1336,7 @@ class Scheduler( "triton": ("SGLANG_TRITON_PREFILL_TRUNCATION_ALIGN_SIZE", 4096), } env_var, default_size = backend_sizes.get( - self.server_args.attention_backend, (None, None) + get_exec().kernel.attention_backend, (None, None) ) self.truncation_align_size = ( get_int_env_var(env_var, default_size) if env_var else None @@ -1725,10 +1732,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 @@ -1830,7 +1837,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, @@ -2006,7 +2013,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) @@ -2053,7 +2060,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() ): @@ -2099,8 +2106,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: @@ -2112,7 +2118,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, @@ -2265,7 +2271,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 = ( @@ -2314,7 +2320,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) @@ -2592,7 +2598,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) @@ -2804,7 +2810,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: @@ -2878,7 +2884,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: @@ -2945,7 +2951,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), @@ -3516,7 +3522,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 @@ -3792,7 +3798,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, ) @@ -3912,7 +3918,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() @@ -4445,10 +4451,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..d46197727 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -2,14 +2,7 @@ from __future__ import annotations import logging from dataclasses import dataclass -from typing import ( - TYPE_CHECKING, - Callable, - List, - Optional, - Tuple, - Union, -) +from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Union import torch @@ -23,11 +16,14 @@ from sglang.srt.managers.schedule_batch import ( ScheduleBatch, mamba_lazy_spec_in_window, ) -from sglang.srt.mem_cache.common import ( - maybe_cache_unfinished_req, - release_kv_cache, +from sglang.srt.mem_cache.common import maybe_cache_unfinished_req, release_kv_cache +from sglang.srt.runtime_context import ( + get_disagg, + get_exec, + get_memory, + get_observability, + get_server_args, ) -from sglang.srt.runtime_context import 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 @@ -48,10 +44,7 @@ if TYPE_CHECKING: SchedulerOutputStreamer, ) from sglang.srt.managers.tp_worker import BaseTpWorker - from sglang.srt.managers.utils import ( - EmbeddingBatchResult, - GenerationBatchResult, - ) + from sglang.srt.managers.utils import EmbeddingBatchResult, GenerationBatchResult from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.mem_cache.memory_pool import ReqToTokenPool @@ -84,7 +77,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 +85,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 +236,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 +749,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 +932,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 +952,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 @@ -1102,7 +1095,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..8503c7b09 100644 --- a/python/sglang/srt/managers/scheduler_components/dp_attn.py +++ b/python/sglang/srt/managers/scheduler_components/dp_attn.py @@ -12,9 +12,7 @@ from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs from sglang.srt.layers.dp_attention import world_dp_gather_enabled from sglang.srt.managers.schedule_batch import ScheduleBatch -from sglang.srt.managers.scheduler_components.recv_skipper import ( - SchedulerRecvSkipper, -) +from sglang.srt.managers.scheduler_components.recv_skipper import SchedulerRecvSkipper from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.mem_cache.memory_pool import ReqToTokenPool @@ -26,6 +24,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 +384,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 e8619a2f7..5652b740e 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 @@ -144,7 +145,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..85fb1ef59 100644 --- a/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py @@ -1,20 +1,15 @@ from __future__ import annotations from dataclasses import dataclass -from typing import ( - List, - Tuple, -) +from typing import List, Tuple 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.server_args import ( - MIS_DELIMITER_TOKEN_ID, - ServerArgs, -) +from sglang.srt.runtime_context import get_exec +from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID, ServerArgs @dataclass(kw_only=True, slots=True, frozen=True) @@ -164,7 +159,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/output_streamer.py b/python/sglang/srt/managers/scheduler_components/output_streamer.py index 278ccf428..d0b82044c 100644 --- a/python/sglang/srt/managers/scheduler_components/output_streamer.py +++ b/python/sglang/srt/managers/scheduler_components/output_streamer.py @@ -2,12 +2,7 @@ from __future__ import annotations import logging from dataclasses import dataclass, field -from typing import ( - Any, - Callable, - List, - Optional, -) +from typing import Any, Callable, List, Optional import torch import zmq @@ -21,11 +16,9 @@ from sglang.srt.managers.io_struct import ( CachedTokensDetails, wrap_as_pickle, ) -from sglang.srt.managers.schedule_batch import ( - BaseFinishReason, - Req, -) +from sglang.srt.managers.schedule_batch import BaseFinishReason, 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 @@ -144,7 +137,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, ) @@ -171,7 +164,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 f1c65a3b1..190583aaa 100644 --- a/python/sglang/srt/managers/scheduler_components/profiler_manager.py +++ b/python/sglang/srt/managers/scheduler_components/profiler_manager.py @@ -5,13 +5,7 @@ import os import time from dataclasses import dataclass from pathlib import Path -from typing import ( - TYPE_CHECKING, - Any, - Callable, - List, - Optional, -) +from typing import TYPE_CHECKING, Any, Callable, List, Optional import torch @@ -19,7 +13,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 @@ -255,7 +249,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 @@ -365,7 +359,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 fccd74c1d..47c145adc 100644 --- a/python/sglang/srt/managers/scheduler_components/request_receiver.py +++ b/python/sglang/srt/managers/scheduler_components/request_receiver.py @@ -2,14 +2,7 @@ from __future__ import annotations from dataclasses import dataclass from http import HTTPStatus -from typing import ( - TYPE_CHECKING, - Any, - Callable, - List, - Optional, - Union, -) +from typing import TYPE_CHECKING, Any, Callable, List, Optional, Union import zmq from torch.distributed import barrier @@ -22,14 +15,9 @@ from sglang.srt.managers.io_struct import ( TokenizedGenerateReqInput, sock_recv, ) -from sglang.srt.managers.mm_utils import ( - has_shm_features, - unwrap_shm_features, -) -from sglang.srt.utils import ( - broadcast_pyobj, - point_to_point_pyobj, -) +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 from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method if TYPE_CHECKING: @@ -220,8 +208,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 0a55601d4..a9b9e1875 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/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index 86b7b378f..cdb9cfea5 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -74,6 +74,7 @@ from sglang.srt.managers.io_struct import ( UpdateWeightsFromTensorReqOutput, ) from sglang.srt.managers.load_snapshot import LoadSnapshot +from sglang.srt.runtime_context import get_lora from sglang.srt.server_args import LoRARef, ServerArgs from sglang.srt.utils import ( get_bool_env_var, @@ -569,7 +570,7 @@ class TokenizerControlMixin: self.auto_create_handle_loop() try: - if not self.server_args.enable_lora: + if not get_lora().enable_lora: raise ValueError( "LoRA is not enabled. Please set `--enable-lora` to enable LoRA." ) @@ -602,10 +603,10 @@ class TokenizerControlMixin: await self.lora_registry.register(new_adapter) self.lora_ref_cache[obj.lora_name] = new_adapter - if self.server_args.max_loaded_loras is not None: + if get_lora().max_loaded_loras is not None: while ( self.lora_registry.num_registered_loras - > self.server_args.max_loaded_loras + > get_lora().max_loaded_loras ): lru_lora_name = await self.lora_registry.lru_lora_name( exclude_pinned=True @@ -619,7 +620,7 @@ class TokenizerControlMixin: logger.info( f"Unloading least recently used LoRA adapter '{lru_lora_name}' " f"(current number of adapters: {self.lora_registry.num_registered_loras}, " - f"max allowed: {self.server_args.max_loaded_loras})" + f"max allowed: {get_lora().max_loaded_loras})" ) unload_result = await self._unload_lora_adapter_locked( @@ -647,7 +648,7 @@ class TokenizerControlMixin: self.auto_create_handle_loop() try: - if not self.server_args.enable_lora: + if not get_lora().enable_lora: raise ValueError( "LoRA is not enabled. Please set `--enable-lora` to enable LoRA." ) @@ -672,10 +673,10 @@ class TokenizerControlMixin: if result.success: await self.lora_registry.register(new_adapter) self.lora_ref_cache[obj.lora_name] = new_adapter - if self.server_args.max_loaded_loras is not None: + if get_lora().max_loaded_loras is not None: while ( self.lora_registry.num_registered_loras - > self.server_args.max_loaded_loras + > get_lora().max_loaded_loras ): lru_lora_name = await self.lora_registry.lru_lora_name( exclude_pinned=True @@ -689,7 +690,7 @@ class TokenizerControlMixin: logger.info( f"Unloading least recently used LoRA adapter '{lru_lora_name}' " f"(current number of adapters: {self.lora_registry.num_registered_loras}, " - f"max allowed: {self.server_args.max_loaded_loras})" + f"max allowed: {get_lora().max_loaded_loras})" ) unload_result = await self._unload_lora_adapter_locked( @@ -717,7 +718,7 @@ class TokenizerControlMixin: self.auto_create_handle_loop() try: - if not self.server_args.enable_lora: + if not get_lora().enable_lora: raise ValueError( "LoRA is not enabled. Please set `--enable-lora` to enable LoRA." ) @@ -893,6 +894,8 @@ class TokenizerControlMixin: ) -> None: """Update weight version if provided.""" if weight_version is not None: - self.server_args.override( + from sglang.srt.runtime_context import get_context + + get_context().override( "tokenizer.weight_version", weight_version=weight_version ) diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 60b35ea02..17036d6ed 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -110,6 +110,14 @@ from sglang.srt.observability.request_metrics_exporter import ( RequestMetricsExporterManager, ) from sglang.srt.observability.trace import SpanAttributes, extract_trace_headers +from sglang.srt.runtime_context import ( + get_device, + get_disagg, + get_lora, + get_model, + get_observability, + get_serving, +) from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.server_args import ( PortArgs, @@ -463,10 +471,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # TODO: Refactor and organize the log export code. # Request logging self.request_logger = RequestLogger( - log_requests=self.server_args.log_requests, - log_requests_level=self.server_args.log_requests_level, - log_requests_format=self.server_args.log_requests_format, - log_requests_target=self.server_args.log_requests_target, + log_requests=get_observability().log_requests, + log_requests_level=get_observability().log_requests_level, + log_requests_format=get_observability().log_requests_format, + log_requests_target=get_observability().log_requests_target, ) # Dumping @@ -489,7 +497,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): def init_weight_update(self): # Initial weights status self.initial_weights_loaded = True - if self.server_args.checkpoint_engine_wait_weights_before_ready: + if get_model().checkpoint_engine_wait_weights_before_ready: self.initial_weights_loaded = False # Weight updates @@ -509,7 +517,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # The registry dynamically updates as adapters are loaded / unloaded during runtime. It # serves as the source of truth for available adapters and maps user-friendly LoRA names # to internally used unique LoRA IDs. - self.lora_registry = LoRARegistry(self.server_args.lora_paths) + self.lora_registry = LoRARegistry(get_lora().lora_paths) # Lock to serialize LoRA update operations. # Please note that, unlike `model_update_lock`, this does not block inference, allowing # LoRA updates and inference to overlap. @@ -518,15 +526,13 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # point to their latest LoRARef objects, so that they can be # dynamically loaded if needed for inference self.lora_ref_cache: Dict[str, LoRARef] = {} - if self.server_args.lora_paths is not None: - for lora_ref in self.server_args.lora_paths: + if get_lora().lora_paths is not None: + for lora_ref in get_lora().lora_paths: self.lora_ref_cache[lora_ref.lora_name] = lora_ref def init_disaggregation(self): # PD Disaggregation - self.disaggregation_mode = DisaggregationMode( - self.server_args.disaggregation_mode - ) + self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode) # Keep a reference so the bootstrap server is not garbage-collected. self.bootstrap_server = start_disagg_service(self.server_args) # Single-source counter for auto-assigning fake bootstrap_room. @@ -535,18 +541,16 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # Encoder Disaggregation self.encoder_bootstrap_server = None if self.server_args.language_only: - from sglang.srt.disaggregation.encode_receiver import ( - EncoderBootstrapServer, - ) + from sglang.srt.disaggregation.encode_receiver import EncoderBootstrapServer # Shared mutable URL list: the bootstrap server appends / removes # entries as encoders register, the receiver reads from the same # list. Pre-populated with static --encoder-urls so the legacy # CLI flag still works (alongside dynamic registrations). - self.encoder_urls: List[str] = list(self.server_args.encoder_urls) + self.encoder_urls: List[str] = list(get_disagg().encoder_urls) self.encoder_bootstrap_server = EncoderBootstrapServer( - host=self.server_args.host, - port=self.server_args.encoder_bootstrap_port, + host=get_serving().host, + port=get_disagg().encoder_bootstrap_port, urls=self.encoder_urls, ) self.mm_receiver = create_mm_receiver( @@ -560,20 +564,22 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # Metrics if self.enable_metrics: engine_type = DisaggregationMode.to_engine_type( - self.server_args.disaggregation_mode + get_disagg().disaggregation_mode ) labels = { - "model_name": self.server_args.served_model_name, + "model_name": get_serving().served_model_name, "engine_type": engine_type, } if self.enable_priority_scheduling: labels["priority"] = "" - if self.server_args.tokenizer_metrics_allowed_custom_labels: - for label in self.server_args.tokenizer_metrics_allowed_custom_labels: + if get_observability().tokenizer_metrics_allowed_custom_labels: + for ( + label + ) in get_observability().tokenizer_metrics_allowed_custom_labels: labels[label] = "" - if self.server_args.extra_metric_labels: - labels.update(self.server_args.extra_metric_labels) + if get_observability().extra_metric_labels: + labels.update(get_observability().extra_metric_labels) tokenizer_collector_cls = resolve_collector_class( self.server_args, STAT_LOGGER_ROLE_TOKENIZER, @@ -582,18 +588,18 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): self.metrics_collector = tokenizer_collector_cls( server_args=self.server_args, labels=labels, - bucket_time_to_first_token=self.server_args.bucket_time_to_first_token, - bucket_e2e_request_latency=self.server_args.bucket_e2e_request_latency, - bucket_inter_token_latency=self.server_args.bucket_inter_token_latency, + bucket_time_to_first_token=get_observability().bucket_time_to_first_token, + bucket_e2e_request_latency=get_observability().bucket_e2e_request_latency, + bucket_inter_token_latency=get_observability().bucket_inter_token_latency, ) start_cpu_monitor_thread("tokenizer") - if self.server_args.gc_warning_threshold_secs > 0.0: - configure_gc_warning(self.server_args.gc_warning_threshold_secs) + if get_observability().gc_warning_threshold_secs > 0.0: + configure_gc_warning(get_observability().gc_warning_threshold_secs) self.soft_watchdog = Watchdog.create( debug_name="TokenizerManager", - watchdog_timeout=self.server_args.soft_watchdog_timeout, + watchdog_timeout=get_device().soft_watchdog_timeout, soft=True, test_stuck_time=envs.SGLANG_TEST_STUCK_TOKENIZER.get(), ) @@ -1757,7 +1763,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # default the load format to the server_args if obj.load_format is None: - obj.load_format = self.server_args.load_format + obj.load_format = get_model().load_format logger.info("Start update_weights. Load format=%s", obj.load_format) if obj.abort_all_requests: @@ -1783,7 +1789,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): def _update_model_path_info(self, model_path: str, load_format: str): self.served_model_name = model_path - self.server_args.override( + from sglang.srt.runtime_context import get_context + + get_context().override( "tokenizer.update_weights", model_path=model_path, load_format=load_format ) self.model_path = model_path @@ -1927,7 +1935,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): "id": rid, "finish_reason": recv_obj.finished_reasons[i], "prompt_tokens": recv_obj.prompt_tokens[i], - "weight_version": self.server_args.weight_version, + "weight_version": get_serving().weight_version, "num_retractions": recv_obj.retraction_counts[i], } @@ -2801,7 +2809,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): meta_info = { "id": recv_obj.rid, "finish_reason": finish_reason, - "weight_version": self.server_args.weight_version, + "weight_version": get_serving().weight_version, "e2e_latency": state.time_stats.get_e2e_latency(), } is_stream = getattr(state.obj, "stream", False) diff --git a/python/sglang/srt/managers/tokenizer_manager_score_mixin.py b/python/sglang/srt/managers/tokenizer_manager_score_mixin.py index 6d5475517..4a1b11110 100644 --- a/python/sglang/srt/managers/tokenizer_manager_score_mixin.py +++ b/python/sglang/srt/managers/tokenizer_manager_score_mixin.py @@ -597,7 +597,10 @@ class TokenizerManagerScoreMixin: f"Token ID {token_id} is out of vocabulary (vocab size: {vocab_size})" ) - # Check if multi-item scoring is enabled + # Check if multi-item scoring is enabled. enable_mis is a static startup + # feature flag (never overridden post-publish), and score_request is also + # exercised on a bare mixin without a published context, so read it off + # server_args rather than the resolved-config bag. use_multi_item_scoring = self.server_args.enable_mis input_ids = None diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 6d09e8133..4234ab41c 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 ( @@ -405,14 +406,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, @@ -423,7 +424,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, @@ -439,11 +440,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, @@ -459,7 +460,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 @@ -485,9 +486,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/allocation.py b/python/sglang/srt/mem_cache/allocation.py index 3d6c8cf86..99a8e1f02 100644 --- a/python/sglang/srt/mem_cache/allocation.py +++ b/python/sglang/srt/mem_cache/allocation.py @@ -26,7 +26,7 @@ from sglang.srt.mem_cache.common import ( evict_from_tree_cache, ) 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_exec, get_server_args from sglang.srt.utils import ( is_cpu, is_cuda, @@ -65,7 +65,7 @@ def write_cache_indices( prefix_tensors: list[torch.Tensor], req_to_token_pool: ReqToTokenPool, ): - if support_triton(get_server_args().attention_backend): + if support_triton(get_exec().kernel.attention_backend): prefix_pointers = torch.tensor( [t.data_ptr() for t in prefix_tensors], dtype=torch.uint64, @@ -106,7 +106,7 @@ def get_last_loc( req_pool_indices_tensor: torch.Tensor, prefix_lens_tensor: torch.Tensor, ) -> torch.Tensor: - attn_backend = get_server_args().attention_backend + attn_backend = get_exec().kernel.attention_backend uses_triton_dispatch = attn_backend not in ("ascend", "torch_native") if _is_hip and uses_triton_dispatch: diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 7a3e1c39d..91cb8615e 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -16,7 +16,7 @@ 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_server_args, get_serving from sglang.srt.utils.common import ceil_align if TYPE_CHECKING: @@ -183,7 +183,7 @@ def _release_overallocated_kv_indices( # 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 c36ee39ea..f4166cf43 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 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() diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index cd6b55ec4..542e2c628 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_disagg, + get_exec, + get_memory, + get_model, + 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 ( @@ -115,9 +123,7 @@ if TYPE_CHECKING: from sglang.srt.model_executor.model_runner_components.spec_aux_hidden_state import ( SpecAuxHiddenStateConfig, ) - from sglang.srt.model_executor.pool_configurator import ( - MemoryPoolConfig, - ) + from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig class KVCacheConfigResult(msgspec.Struct, frozen=True, kw_only=True): @@ -308,8 +314,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: @@ -358,13 +364,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 @@ -394,7 +400,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) ): @@ -432,8 +438,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 @@ -462,14 +468,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. @@ -502,13 +508,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) @@ -558,8 +564,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, @@ -579,7 +585,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 @@ -588,7 +594,7 @@ class KVCacheConfigurator: elif current_platform.is_out_of_tree() and not self.mambaish_config: unsupported_pool_family = "out-of-tree platform KV pool" elif ( - self.server_args.attention_backend == "ascend" and not self.mambaish_config + get_exec().kernel.attention_backend == "ascend" and not self.mambaish_config ): unsupported_pool_family = "NPU/Ascend KV pool" elif self.use_mla_backend and is_dsa_model: @@ -614,9 +620,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, @@ -648,15 +654,13 @@ class KVCacheConfigurator: extra_max_context_len: int, pre_alloc_size: int, ) -> ReqToTokenPool: - from sglang.srt.disaggregation.decode import ( - HybridMambaDecodeReqToTokenPool, - ) + from sglang.srt.disaggregation.decode import HybridMambaDecodeReqToTokenPool req_to_token_pool = HybridMambaDecodeReqToTokenPool( 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=( [ @@ -666,11 +670,11 @@ 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, ) return req_to_token_pool @@ -688,7 +692,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 @@ -701,11 +705,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=( [ @@ -717,18 +721,18 @@ 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, # ReplaySSM spec-verify is GDN-only: activate the pool machinery # (rings + cursors + the intermediate_ssm gate) only for GDN-hybrid # models, so any other mamba-ish model (Mamba2/Nemotron, lightning, # ...) run with the flag set stays byte-identical to flag-off. 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 ), ) @@ -754,7 +758,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 @@ -770,7 +774,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 ) @@ -802,7 +806,7 @@ class KVCacheConfigurator: max_total_num_tokens=sizes.max_total_num_tokens, ) elif ( - self.server_args.attention_backend == "ascend" and not self.mambaish_config + get_exec().kernel.attention_backend == "ascend" and not self.mambaish_config ): if self.is_hybrid_swa: token_to_kv_pool = self._build_ascend_swa_kv_pool( @@ -878,14 +882,12 @@ 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." if self.is_draft_worker: - from sglang.srt.models.deepseek_v4_nextn import ( - COMPRESS_RATIO_NEXTN_LAYER, - ) + from sglang.srt.models.deepseek_v4_nextn import COMPRESS_RATIO_NEXTN_LAYER compression_ratios = [ COMPRESS_RATIO_NEXTN_LAYER @@ -912,12 +914,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: @@ -935,7 +937,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, @@ -946,11 +948,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 ), @@ -961,7 +963,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, @@ -972,7 +974,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), @@ -985,14 +987,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, ) @@ -1002,13 +1004,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, ) @@ -1020,9 +1022,7 @@ class KVCacheConfigurator: full_max_total_num_tokens: Optional[int], swa_max_total_num_tokens: Optional[int], ) -> KVCache: - from sglang.srt.hardware_backend.npu.memory_pool_npu import ( - NPUMHATokenToKVPool, - ) + from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMHATokenToKVPool kwargs = {} if self.is_hybrid_swa_compress: @@ -1039,7 +1039,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), @@ -1055,39 +1055,35 @@ class KVCacheConfigurator: def _build_ascend_mla_kv_pool( self, *, max_total_num_tokens: int, is_dsa_model: bool ) -> KVCache: - from sglang.srt.hardware_backend.npu.memory_pool_npu import ( - NPUMLATokenToKVPool, - ) + from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMLATokenToKVPool 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, ) return token_to_kv_pool def _build_ascend_mha_kv_pool(self, *, max_total_num_tokens: int) -> KVCache: - from sglang.srt.hardware_backend.npu.memory_pool_npu import ( - NPUMHATokenToKVPool, - ) + from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMHATokenToKVPool 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, ) @@ -1101,7 +1097,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 @@ -1121,7 +1117,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, @@ -1132,7 +1128,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), @@ -1143,13 +1139,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, ) @@ -1158,13 +1154,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, ) @@ -1221,7 +1217,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), @@ -1229,7 +1225,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, ) @@ -1244,7 +1240,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), @@ -1254,7 +1250,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, ) @@ -1293,7 +1289,7 @@ class KVCacheConfigurator: 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), @@ -1302,8 +1298,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, @@ -1316,18 +1312,18 @@ 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 @@ -1339,7 +1335,7 @@ class KVCacheConfigurator: 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 = {} @@ -1349,18 +1345,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 @@ -1375,20 +1371,20 @@ 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, need_sort=need_sort, ) elif _is_npu and ( - self.server_args.attention_backend == "ascend" + get_exec().kernel.attention_backend == "ascend" or is_dsv4_model or self.hybrid_gdn_config is not None ): @@ -1406,7 +1402,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, @@ -1419,7 +1415,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, @@ -1429,7 +1425,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, @@ -1439,22 +1435,20 @@ 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: - from sglang.srt.mem_cache.sparsity import ( - parse_hisparse_config, - ) + if get_memory().enable_hisparse: + from sglang.srt.mem_cache.sparsity import parse_hisparse_config 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, @@ -1462,8 +1456,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, @@ -1475,7 +1468,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, @@ -1483,7 +1476,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 @@ -1535,7 +1528,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( @@ -1559,7 +1552,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 = " @@ -1570,14 +1563,14 @@ 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 additional_ratio = 0 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: @@ -1596,7 +1589,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: @@ -1626,7 +1619,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) @@ -1637,13 +1630,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 " @@ -1673,7 +1666,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( @@ -1689,18 +1682,20 @@ 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 def _handle_max_mamba_cache(self, total_rest_memory): + from sglang.srt.runtime_context import get_context + config = self.mambaish_config server_args = self.server_args assert config is not None @@ -1710,11 +1705,11 @@ class KVCacheConfigurator: assert server_args.speculative_num_draft_tokens is not None assert server_args.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 @@ -1722,7 +1717,7 @@ class KVCacheConfigurator: 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_mamba_cache_size // ratio, ) intermediate_size = ( config.mamba2_cache_params.mamba_cache_per_req @@ -1735,7 +1730,7 @@ class KVCacheConfigurator: and server_args.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 // self.ps.attn_dp_size, @@ -1744,7 +1739,7 @@ class KVCacheConfigurator: if has_spec_dec: intermediate_size = ( config.mamba2_cache_params.mamba_cache_per_req - * server_args.max_mamba_cache_size + * get_schedule().max_mamba_cache_size * server_args.speculative_num_draft_tokens ) total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) @@ -1769,7 +1764,7 @@ class KVCacheConfigurator: ratio = self._calculate_mamba_ratio() D = server_args.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 / ratio)) @@ -1779,12 +1774,12 @@ class KVCacheConfigurator: # 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_mamba_cache_size // ratio, ) intermediate_size = per_req * capped_reqs * D total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30)) else: - server_args.override( + get_context().override( "mamba_pool.memory_budget", max_mamba_cache_size=int(mamba_budget_bytes // per_req), ) @@ -1793,10 +1788,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, " @@ -1806,7 +1801,7 @@ class KVCacheConfigurator: ) mamba_state_memory = ( - server_args.max_mamba_cache_size + get_schedule().max_mamba_cache_size * config.mamba2_cache_params.mamba_cache_per_req / (1 << 30) ) 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 897a5e61b..8460a6d62 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 try: from lmcache.integration.sglang.multi_process_adapter import LMCacheMPConnector @@ -108,7 +108,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( diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 58ccdebf2..9473ea985 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -51,13 +51,8 @@ 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.utils import ( - is_cuda, - is_hip, - is_npu, - support_triton, -) +from sglang.srt.runtime_context import get_exec, get_parallel +from sglang.srt.utils import is_cuda, is_hip, is_npu, support_triton from sglang.srt.utils.common import ceil_align, is_pin_memory_available if TYPE_CHECKING: @@ -941,7 +936,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( @@ -1110,7 +1105,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 fdf2e69bd..414286793 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -26,11 +26,7 @@ import torch import torch.distributed as dist from sglang.srt.configs.load_config import LoadConfig -from sglang.srt.configs.model_config import ( - AttentionArch, - ModelConfig, - ModelImpl, -) +from sglang.srt.configs.model_config import AttentionArch, ModelConfig, ModelImpl from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp from sglang.srt.debug_utils.dumper import dumper from sglang.srt.distributed import bootstrap @@ -74,9 +70,7 @@ from sglang.srt.kv_canary.runner.canary_manager import context_tuple from sglang.srt.kv_canary.token_oracle.install import install_token_oracle_from_env from sglang.srt.layers import deep_gemm_wrapper, model_parallel from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp -from sglang.srt.layers.cp.utils import ( - get_cp_strategy, -) +from sglang.srt.layers.cp.utils import get_cp_strategy from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.sampler import create_sampler from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model @@ -86,17 +80,10 @@ from sglang.srt.lora.lora_registry import LoRARef from sglang.srt.managers.schedule_batch import sanity_check_mm_pad_shift_value from sglang.srt.mem_cache import kv_cache_dtype from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator -from sglang.srt.mem_cache.kv_cache_configurator import ( - KVCacheConfigurator, -) +from sglang.srt.mem_cache.kv_cache_configurator import KVCacheConfigurator from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool -from sglang.srt.model_executor.cuda_graph_config import ( - cuda_graph_fully_disabled, -) -from sglang.srt.model_executor.forward_batch_info import ( - ForwardBatch, - PPProxyTensors, -) +from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled +from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_executor.forward_context import ( ForwardContext, forward_context, @@ -155,14 +142,15 @@ from sglang.srt.model_executor.model_runner_components.weight_updater import ( WeightUpdater, ) from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig -from sglang.srt.model_executor.runner import ( - EagerRunner, - get_batch_sizes_to_capture, -) +from sglang.srt.model_executor.runner import EagerRunner, get_batch_sizes_to_capture from sglang.srt.platforms import current_platform from sglang.srt.runtime_context import ( + get_device, + get_exec, get_global_dwdp_manager, - get_server_args, + get_lora, + get_model, + get_schedule, set_global_dwdp_manager, ) from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo @@ -319,7 +307,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) @@ -396,12 +384,12 @@ 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_joiner ): return - is_scale_join = self.server_args.ep_join_mode == "scale" + is_scale_join = get_exec().moe.ep_join_mode == "scale" if is_scale_join: join_effective_ep_size = ( self.server_args.ep_join_rank_offset + self.ps.tp_size @@ -484,7 +472,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, @@ -561,7 +549,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( @@ -618,7 +606,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): @@ -654,7 +642,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): @@ -668,12 +656,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): @@ -692,8 +680,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 ) @@ -702,17 +690,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() @@ -973,7 +961,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") @@ -1030,7 +1018,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_scale_joiner=self.server_args.is_ep_scale_joiner, ) @@ -1050,16 +1038,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( @@ -1331,7 +1319,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 @@ -1361,7 +1349,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 @@ -1765,7 +1753,7 @@ class ModelRunner: recovered = maybe_recover_ep_ranks( tp_group=self.tp_group, eplb_manager=self.eplb_manager, - random_seed=self.server_args.random_seed, + random_seed=get_device().random_seed, ) if recovered: self.forward_pass_id = 0 @@ -1774,7 +1762,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) @@ -1842,7 +1830,9 @@ class ModelRunner: load_config: LoadConfig, ) -> None: self.model = new_model - self.server_args.override( + from sglang.srt.runtime_context import get_context + + 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/misc_utils.py b/python/sglang/srt/model_executor/model_runner_components/misc_utils.py index f0261d3ae..cf82daa3d 100644 --- a/python/sglang/srt/model_executor/model_runner_components/misc_utils.py +++ b/python/sglang/srt/model_executor/model_runner_components/misc_utils.py @@ -24,6 +24,12 @@ def maybe_disable_chunked_prefix_cache( # model's (often non-MLA) config must not flip the shared setting. if is_draft_worker: return + + # This is a load-time gate that runs in ModelRunner.__init__ BEFORE the + # runner publishes its config (and direct/benchmark construction never + # publishes earlier), so read/write the supplied server_args. The runner's + # subsequent publish snapshots this into the schedule bag for get_schedule() + # readers. if ( not use_mla_backend or server_args.attention_backend 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_loader/loader.py b/python/sglang/srt/model_loader/loader.py index 13c0c07d4..cce163b4e 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -44,7 +44,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_server_args from sglang.srt.utils import get_available_gpu_memory # Try to import accelerate (optional dependency) @@ -71,9 +71,7 @@ from sglang.srt.connector import ( get_connector_type, ) from sglang.srt.connector.utils import parse_model_name -from sglang.srt.distributed import ( - model_parallel_is_initialized, -) +from sglang.srt.distributed import model_parallel_is_initialized from sglang.srt.layers.modelopt_utils import QUANT_CFG_CHOICES from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.model_loader.remote_instance_weight_loader_utils import ( @@ -865,9 +863,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) diff --git a/python/sglang/srt/models/bailing_moe.py b/python/sglang/srt/models/bailing_moe.py index 9abccdbe0..cc58be6c0 100644 --- a/python/sglang/srt/models/bailing_moe.py +++ b/python/sglang/srt/models/bailing_moe.py @@ -41,9 +41,7 @@ from sglang.srt.layers.communicator import ( LayerScatterModes, enable_moe_dense_fully_dp, ) -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( MergedColumnParallelLinear, @@ -78,6 +76,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 +208,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 +222,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..7044f0a9a 100644 --- a/python/sglang/srt/models/bailing_moe_linear.py +++ b/python/sglang/srt/models/bailing_moe_linear.py @@ -12,17 +12,12 @@ from transformers import PretrainedConfig from sglang.kernels.ops.attention.fla.layernorm_gated import RMSNorm as RMSNormGated from sglang.kernels.ops.attention.fla.layernorm_gated import layernorm_fn from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz -from sglang.srt.distributed import ( - get_pp_group, - tensor_model_parallel_all_reduce, -) +from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.layers import deep_gemm_wrapper from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( ColumnParallelLinear, @@ -59,6 +54,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 +525,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 +686,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 6e24068a1..5b3f8a95b 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_server_args +from sglang.srt.runtime_context import get_exec from sglang.srt.utils import 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") @@ -187,7 +187,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 31f7e73d0..5985ffe66 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 bc26c7522..2117e1593 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 @@ -65,10 +65,8 @@ 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.state_capturer.indexer_topk import ( - maybe_capture_indexer_topk, -) +from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args +from sglang.srt.state_capturer.indexer_topk import maybe_capture_indexer_topk from sglang.srt.utils import BumpAllocator from sglang.srt.utils.custom_op import register_custom_op @@ -153,7 +151,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( @@ -990,8 +988,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 ( diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py index 222effb03..0fa9172ff 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 f2001231b..10488e61f 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -29,10 +29,7 @@ import torch.nn.functional as F from torch import nn from transformers import PretrainedConfig -from sglang.jit_kernel.dsv4 import ( - silu_and_mul_clamp, - silu_and_mul_contig_post_quant, -) +from sglang.jit_kernel.dsv4 import silu_and_mul_clamp, silu_and_mul_contig_post_quant from sglang.kernels.ops.quantization.fp8_kernel import ( create_per_token_group_quant_fp8_output_scale, ) @@ -78,9 +75,7 @@ from sglang.srt.layers.communicator_dsa_cp import ( maybe_prefetch_next_full_attention_kv, ) from sglang.srt.layers.cp.utils import is_cp_v2_active -from sglang.srt.layers.dcp.planner import ( - prepare_decode_context_parallel_metadata, -) +from sglang.srt.layers.dcp.planner import prepare_decode_context_parallel_metadata from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( ColumnParallelLinear, @@ -115,9 +110,7 @@ from sglang.srt.layers.moe.utils import ( ) from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.fp8 import Fp8Config -from sglang.srt.layers.quantization.fp8_utils import ( - materialize_bpreshuffle_fp8_scale, -) +from sglang.srt.layers.quantization.fp8_utils import materialize_bpreshuffle_fp8_scale from sglang.srt.layers.quantization.mxfp4_flashinfer_trtllm_moe import ( maybe_fuse_routed_scale_and_shared_add, ) @@ -183,11 +176,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 ( @@ -381,9 +377,7 @@ class DeepseekV2MLP(nn.Module): return down_output if self.use_fused_clamp_act_mul and self.swiglu_limit is not None: - from aiter.ops.triton.fusions.fused_clamp_act_mul import ( - fused_clamp_act_mul, - ) + from aiter.ops.triton.fusions.fused_clamp_act_mul import fused_clamp_act_mul if not self._fused_clamp_fp8_checked: from sglang.srt.layers.quantization.fp8 import Fp8LinearMethod @@ -494,7 +488,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 ( @@ -560,7 +554,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. @@ -630,8 +624,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, @@ -804,7 +797,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 @@ -1718,7 +1711,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): @@ -2069,7 +2062,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 @@ -2765,7 +2758,7 @@ 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 diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 264806123..733744257 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -29,18 +29,11 @@ from sglang.jit_kernel.dsv4 import ( fused_rope_inplace, sglang_per_token_group_quant_fp8_dsv4_wo_a, ) -from sglang.kernels.ops.attention.deepseek_v4_rope import ( - v4_rope_inplace_npu, -) -from sglang.kernels.ops.quantization.fp8_kernel import ( - sglang_per_token_group_quant_fp8, -) +from sglang.kernels.ops.attention.deepseek_v4_rope import v4_rope_inplace_npu +from sglang.kernels.ops.quantization.fp8_kernel import sglang_per_token_group_quant_fp8 from sglang.srt.compilation.compilation_config import register_split_op from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config -from sglang.srt.distributed import ( - get_pp_group, - get_tp_group, -) +from sglang.srt.distributed import get_pp_group, get_tp_group from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) @@ -136,7 +129,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 ( @@ -311,9 +310,7 @@ def _freqs_cis_to_cos_sin( if TYPE_CHECKING: - from sglang.srt.layers.attention.deepseek_v4_backend import ( - DeepseekV4AttnBackend, - ) + from sglang.srt.layers.attention.deepseek_v4_backend import DeepseekV4AttnBackend from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import ( DeepseekV4HipRadixBackend, ) @@ -575,7 +572,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_hip: @@ -2458,11 +2455,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..888edb932 100755 --- a/python/sglang/srt/models/exaone_moe.py +++ b/python/sglang/srt/models/exaone_moe.py @@ -24,17 +24,12 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import ( - get_pp_group, - tensor_model_parallel_all_reduce, -) +from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo from sglang.srt.layers.activation import SiluAndMul -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( MergedColumnParallelLinear, @@ -62,7 +57,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 +165,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 +206,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..9d4e1f5e1 100644 --- a/python/sglang/srt/models/gemma4_causal.py +++ b/python/sglang/srt/models/gemma4_causal.py @@ -18,11 +18,7 @@ from typing import Iterable, List, Optional, Set, Tuple, Union import torch from torch import nn -from transformers import ( - Gemma4TextConfig, - PretrainedConfig, - PreTrainedModel, -) +from transformers import Gemma4TextConfig, PretrainedConfig, PreTrainedModel from sglang.kernels.ops.layernorm.gemma4_fused_ops import ( gemma4_fused_routing, @@ -31,9 +27,7 @@ from sglang.kernels.ops.layernorm.gemma4_fused_ops import ( gemma_rmsnorm_residual_scalar, gemma_routing_post_topk, ) -from sglang.srt.distributed import ( - get_pp_group, -) +from sglang.srt.distributed import get_pp_group from sglang.srt.layers.layernorm import Gemma4RMSNorm, RMSNorm from sglang.srt.layers.linear import ( QKVParallelLinear, @@ -55,10 +49,8 @@ from sglang.srt.model_loader.weight_utils import ( maybe_remap_kv_scale_name, ) from sglang.srt.models.gemma3_causal import Gemma3MLP, Gemma3TextScaledWordEmbedding -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.models.utils import create_fused_set_kv_buffer_arg +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 +246,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, 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 c650585a5..2b987ab42 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, @@ -189,7 +190,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 @@ -216,7 +217,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, @@ -284,7 +285,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 @@ -928,7 +929,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..4c3c96024 100644 --- a/python/sglang/srt/models/glm_image_vl.py +++ b/python/sglang/srt/models/glm_image_vl.py @@ -34,10 +34,7 @@ from sglang.srt.layers.attention.vision import ( ) from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm -from sglang.srt.layers.linear import ( - QKVParallelLinear, - RowParallelLinear, -) +from sglang.srt.layers.linear import QKVParallelLinear, RowParallelLinear from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.radix_attention import RadixAttention @@ -57,7 +54,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 +1015,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 cd5731671..bd671c571 100644 --- a/python/sglang/srt/models/gpt_oss.py +++ b/python/sglang/srt/models/gpt_oss.py @@ -34,9 +34,7 @@ from sglang.srt.distributed import ( from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( QKVParallelLinear, @@ -69,6 +67,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 +229,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, @@ -420,7 +419,7 @@ class GptOssAttention(nn.Module): # Choose dtype of sinks based on attention backend: trtllm_mha requires float32, # others can use bfloat16 - attn_backend = get_server_args().attention_backend + attn_backend = get_exec().kernel.attention_backend sinks_dtype = torch.float32 if attn_backend == "trtllm_mha" else torch.bfloat16 self.sinks = nn.Parameter( torch.empty(self.num_heads, dtype=sinks_dtype), requires_grad=False diff --git a/python/sglang/srt/models/inkling.py b/python/sglang/srt/models/inkling.py index 2086723de..e8ebcb182 100644 --- a/python/sglang/srt/models/inkling.py +++ b/python/sglang/srt/models/inkling.py @@ -14,9 +14,7 @@ from sglang.srt.configs.inkling import ( InklingModelConfig, InklingVisionConfig, ) -from sglang.srt.distributed import ( - get_tensor_model_parallel_group, -) +from sglang.srt.distributed import get_tensor_model_parallel_group from sglang.srt.environ import envs from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.logits_processor import LogitsProcessor @@ -72,7 +70,12 @@ 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_exec, + get_model, + get_parallel, + get_server_args, +) from sglang.srt.utils import add_prefix, is_cuda, make_layers logger = logging.getLogger(__name__) @@ -218,7 +221,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" @@ -1276,9 +1279,7 @@ class InklingForConditionalGeneration(nn.Module): and not self.text_config.inference_moe_w13_interleaved and weight_loader is not default_weight_loader ): - from sglang.srt.layers.quantization.modelopt_quant import ( - deinterleave_w13, - ) + from sglang.srt.layers.quantization.modelopt_quant import deinterleave_w13 loaded_weight = deinterleave_w13(loaded_weight) if ( diff --git a/python/sglang/srt/models/inkling_common/attn.py b/python/sglang/srt/models/inkling_common/attn.py index df2963795..5d47318a6 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 0509cb29b..d4b6d357a 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_model 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, @@ -484,7 +484,7 @@ class InklingBatchDenseMLP(nn.Module, FusedMoELoadingMixin): # All shared experts must share one global weight scale (reshard with # single_global_scale=True). ModelOpt's input_scale = amax / (6 * 448). flat2 = scale2.reshape(-1).float() - if get_server_args().load_format == "dummy" and not bool( + if get_model().load_format == "dummy" and not bool( torch.all(flat2 == flat2[0]) ): # Dummy loading uses per-element noise; replace it with a valid scale. diff --git a/python/sglang/srt/models/inkling_common/kernels/comm.py b/python/sglang/srt/models/inkling_common/kernels/comm.py index dba1feade..41dc048b6 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: @@ -250,7 +250,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 @@ -395,7 +395,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 ( @@ -694,7 +694,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() ): @@ -1031,7 +1031,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 0c5e8c357..6673e6577 100644 --- a/python/sglang/srt/models/inkling_common/moe.py +++ b/python/sglang/srt/models/inkling_common/moe.py @@ -16,9 +16,7 @@ from sglang.jit_kernel.inkling_gate_topk_renorm import ( ) from sglang.kernels.jit.utils import is_arch_support_pdl from sglang.srt.configs.inkling import InklingModelConfig -from sglang.srt.distributed import ( - get_tensor_model_parallel_group, -) +from sglang.srt.distributed import get_tensor_model_parallel_group from sglang.srt.environ import GateGemvMode, envs from sglang.srt.layers.moe import get_moe_runner_backend from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE @@ -59,7 +57,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 +889,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 8bb501934..8aa9435c7 100644 --- a/python/sglang/srt/models/inkling_common/sconv.py +++ b/python/sglang/srt/models/inkling_common/sconv.py @@ -26,7 +26,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.speculative.eagle_info import EagleDraftExtendInput from sglang.srt.utils import is_cuda, set_weight_attrs @@ -482,7 +482,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 27d1e1845..529f7a469 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..3669a60a9 100644 --- a/python/sglang/srt/models/laguna.py +++ b/python/sglang/srt/models/laguna.py @@ -17,19 +17,11 @@ import torch.nn.functional as F from torch import nn from sglang.srt.configs.laguna import LagunaConfig, normalize_gating -from sglang.srt.distributed import ( - get_pp_group, - tensor_model_parallel_all_reduce, -) +from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce from sglang.srt.environ import envs from sglang.srt.layers.activation import SiluAndMul -from sglang.srt.layers.communicator import ( - LayerCommunicator, - LayerScatterModes, -) -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( ColumnParallelLinear, @@ -53,7 +45,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 +152,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..52dde19a5 100644 --- a/python/sglang/srt/models/llada2.py +++ b/python/sglang/srt/models/llada2.py @@ -41,9 +41,7 @@ from sglang.srt.layers.communicator import ( LayerScatterModes, enable_moe_dense_fully_dp, ) -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( MergedColumnParallelLinear, @@ -77,6 +75,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 +230,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 +244,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 294710d11..995cc2161 100644 --- a/python/sglang/srt/models/llama_eagle3.py +++ b/python/sglang/srt/models/llama_eagle3.py @@ -38,7 +38,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.models.llama import LlamaDecoderLayer, LlamaForCausalLM, LlamaMLP -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_spec class LlamaDecoderLayer(LlamaDecoderLayer): @@ -258,7 +258,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 e71e42791..5962f1733 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..3a292b1ce 100644 --- a/python/sglang/srt/models/mimo_v2.py +++ b/python/sglang/srt/models/mimo_v2.py @@ -23,10 +23,7 @@ from torch import nn from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo from sglang.srt.configs.model_config import get_mimo_v2_fused_qkv_expected_tp_size -from sglang.srt.distributed import ( - get_pp_group, - tensor_model_parallel_all_reduce, -) +from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo @@ -37,9 +34,7 @@ from sglang.srt.layers.communicator import ( ScatterMode, enable_moe_dense_fully_dp, ) -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( MergedColumnParallelLinear, @@ -79,6 +74,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 +409,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 +444,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 bb7503b1d..7fc27d1cd 100644 --- a/python/sglang/srt/models/minimax_m2.py +++ b/python/sglang/srt/models/minimax_m2.py @@ -32,10 +32,7 @@ from sglang.jit_kernel.all_reduce import ( ) from sglang.kernel_api_logging import debug_kernel_api from sglang.srt.batch_overlap.two_batch_overlap import model_forward_maybe_tbo -from sglang.srt.distributed import ( - get_pp_group, - tensor_model_parallel_all_reduce, -) +from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo from sglang.srt.layers.communicator import ( @@ -43,10 +40,7 @@ from sglang.srt.layers.communicator import ( LayerScatterModes, ScatterMode, ) -from sglang.srt.layers.dp_attention import ( - attn_tp_all_reduce, - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import attn_tp_all_reduce, is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( QKVParallelLinear, @@ -80,7 +74,12 @@ 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_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 @@ -513,7 +512,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 49241cfa0..9735e2fd7 100644 --- a/python/sglang/srt/models/minimax_m3.py +++ b/python/sglang/srt/models/minimax_m3.py @@ -28,10 +28,7 @@ from sglang.srt.configs.model_config import ( get_minimax_sparse_disable_value_layer_ids, get_minimax_sparse_layer_ids, ) -from sglang.srt.distributed import ( - get_pp_group, - tensor_model_parallel_all_reduce, -) +from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce from sglang.srt.environ import envs from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo @@ -79,7 +76,7 @@ from sglang.srt.model_loader.weight_utils import ( maybe_remap_kv_scale_name, ) from sglang.srt.models.minimax_m2 import MiniMaxM2RMSNormTP -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, @@ -287,7 +284,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 ) @@ -311,7 +308,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, @@ -1454,7 +1451,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 c3d563962..811b9b4ec 100644 --- a/python/sglang/srt/models/minimax_m3_vl.py +++ b/python/sglang/srt/models/minimax_m3_vl.py @@ -6,9 +6,7 @@ from typing import Iterable, List, Optional, Tuple import torch import torch.nn as nn -from sglang.srt.distributed import ( - get_pp_group, -) +from sglang.srt.distributed import get_pp_group from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.moe.utils import get_moe_a2a_backend from sglang.srt.layers.quantization.base_config import QuantizationConfig @@ -19,10 +17,7 @@ from sglang.srt.managers.mm_utils import ( MultiModalityDataPaddingPatternMultimodalTokens, general_mm_embed_routine, ) -from sglang.srt.managers.schedule_batch import ( - MultimodalDataItem, - MultimodalInputs, -) +from sglang.srt.managers.schedule_batch import MultimodalDataItem, MultimodalInputs from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors from sglang.srt.model_loader.weight_utils import ( default_weight_loader, @@ -42,7 +37,7 @@ from sglang.srt.models.minimax_vl_common import ( load_vision_weight, merge_vit_qkv_weights, ) -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, get_device_sm, is_cuda, log_info_on_rank0 from sglang.srt.utils.hf_transformers_utils import get_rope_config @@ -65,7 +60,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() diff --git a/python/sglang/srt/models/minimax_vl_common.py b/python/sglang/srt/models/minimax_vl_common.py index 0987165b7..68f811578 100644 --- a/python/sglang/srt/models/minimax_vl_common.py +++ b/python/sglang/srt/models/minimax_vl_common.py @@ -18,16 +18,13 @@ from sglang.srt.layers.attention.vision import ( prepare_vision_attention_metadata, ) from sglang.srt.layers.dp_attention import is_dp_attention_enabled -from sglang.srt.layers.linear import ( - ColumnParallelLinear, - RowParallelLinear, -) +from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear from sglang.srt.layers.quantization.base_config import QuantizationConfig 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 +410,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 +676,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 +688,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 +720,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..68976d6b4 100644 --- a/python/sglang/srt/models/moss_vl.py +++ b/python/sglang/srt/models/moss_vl.py @@ -34,10 +34,7 @@ from sglang.srt.layers.linear import ( from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.radix_attention import RadixAttention -from sglang.srt.layers.rotary_embedding import ( - MRotaryEmbedding, - get_rope, -) +from sglang.srt.layers.rotary_embedding import MRotaryEmbedding, get_rope from sglang.srt.layers.rotary_embedding.mrope import apply_interleaved_rope from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb from sglang.srt.layers.vocab_parallel_embedding import ( @@ -48,7 +45,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 +999,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 58253830b..b76425ad6 100644 --- a/python/sglang/srt/models/nemotron_h.py +++ b/python/sglang/srt/models/nemotron_h.py @@ -36,10 +36,7 @@ from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( Mamba2AttnBackend, ) from sglang.srt.layers.attention.mamba.mamba import MambaMixer2 -from sglang.srt.layers.dp_attention import ( - attn_tp_all_reduce, - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import attn_tp_all_reduce, is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( ColumnParallelLinear, @@ -89,7 +86,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 +202,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..7e61333f2 100644 --- a/python/sglang/srt/models/qwen2.py +++ b/python/sglang/srt/models/qwen2.py @@ -22,10 +22,7 @@ from typing import Any, Dict, Iterable, List, Optional, Tuple, Union import torch from torch import nn -from sglang.srt.distributed import ( - get_pp_group, - get_pp_indices, -) +from sglang.srt.distributed import get_pp_group, get_pp_indices from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm @@ -50,7 +47,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 +93,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 +327,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 +363,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 72e8ddee4..786624eda 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -46,9 +46,7 @@ from sglang.srt.layers.communicator import ( ScatterMode, ) from sglang.srt.layers.cp.utils import is_cp_v2_active -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( MergedColumnParallelLinear, @@ -91,7 +89,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, @@ -146,7 +149,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() @@ -271,10 +274,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, @@ -333,7 +336,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..ff1209b23 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -5,9 +5,7 @@ from typing import Any, Dict, Iterable, List, Optional, Tuple import torch from torch import nn -from sglang.srt.distributed import ( - get_pp_group, -) +from sglang.srt.distributed import get_pp_group from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import QKVParallelLinear, RowParallelLinear @@ -33,7 +31,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 +114,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 +274,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 +301,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 +365,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 fe26e906e..d402d95c8 100644 --- a/python/sglang/srt/models/qwen3_5.py +++ b/python/sglang/srt/models/qwen3_5.py @@ -43,9 +43,7 @@ from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_r from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.layers.attention.mamba.mamba import mamba_v2_sharded_weight_loader from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled # Layers - Others from sglang.srt.layers.layernorm import GemmaRMSNorm @@ -92,9 +90,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 +137,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: @@ -1197,7 +1195,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 9cee1d21d..0c509fa1d 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, @@ -283,7 +284,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 @@ -514,7 +515,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 7036cd717..9dca3f92e 100644 --- a/python/sglang/srt/models/qwen3_vl.py +++ b/python/sglang/srt/models/qwen3_vl.py @@ -38,9 +38,7 @@ from sglang.srt.layers.attention.vision import ( prepare_vision_attention_metadata, ) from sglang.srt.layers.conv import Conv3dLayer -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.pooler import Pooler, PoolingType @@ -70,14 +68,8 @@ from sglang.srt.models.utils import ( ) 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.utils import ( - add_prefix, - cpu_has_amx_support, - is_cpu, - is_npu, - round_up, -) +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, is_cpu, is_npu, round_up from sglang.srt.utils.hf_transformers_utils import get_processor _is_npu = is_npu() @@ -326,7 +318,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 @@ -369,7 +361,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: @@ -917,7 +909,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( @@ -1233,7 +1225,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 1a977b7aa..cac464596 100644 --- a/python/sglang/srt/models/sarvam_moe.py +++ b/python/sglang/srt/models/sarvam_moe.py @@ -13,10 +13,7 @@ from torch import nn from transformers import PretrainedConfig from sglang.kernels.ops.attention.utils import concat_and_cast_mha_k_triton -from sglang.srt.distributed import ( - get_pp_group, - tensor_model_parallel_all_reduce, -) +from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.layers.activation import SiluAndMul @@ -25,9 +22,7 @@ from sglang.srt.layers.communicator import ( LayerScatterModes, enable_moe_dense_fully_dp, ) -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( ColumnParallelLinear, @@ -61,6 +56,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 +270,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..ed9dcb852 100644 --- a/python/sglang/srt/models/sdar.py +++ b/python/sglang/srt/models/sdar.py @@ -13,9 +13,7 @@ from transformers import PretrainedConfig from sglang.srt.distributed import get_pp_group from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( MergedColumnParallelLinear, @@ -42,6 +40,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 +206,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 +234,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 +269,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 +394,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..7e715b789 100644 --- a/python/sglang/srt/models/sdar_moe.py +++ b/python/sglang/srt/models/sdar_moe.py @@ -10,17 +10,12 @@ import torch from torch import nn from transformers import PretrainedConfig -from sglang.srt.distributed import ( - get_pp_group, - tensor_model_parallel_all_reduce, -) +from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ( QKVParallelLinear, @@ -58,6 +53,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 +97,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 +119,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 +270,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 +298,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 +334,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 +474,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..e3c0e720e 100644 --- a/python/sglang/srt/models/step3p5.py +++ b/python/sglang/srt/models/step3p5.py @@ -4,18 +4,13 @@ import torch import torch.nn.functional as F from torch import nn -from sglang.srt.distributed import ( - get_pp_group, - tensor_model_parallel_all_reduce, -) +from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo from sglang.srt.layers.activation import SiluAndMul from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes -from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, -) +from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.layers.layernorm import GemmaRMSNorm from sglang.srt.layers.linear import ( ColumnParallelLinear, @@ -47,6 +42,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 +149,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 +172,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 +672,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..3ee32ba40 100644 --- a/python/sglang/srt/models/transformers.py +++ b/python/sglang/srt/models/transformers.py @@ -58,14 +58,12 @@ from sglang.srt.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) -from sglang.srt.managers.mm_utils import ( - MultiModalityDataPaddingPatternMultimodalTokens, -) +from sglang.srt.managers.mm_utils import MultiModalityDataPaddingPatternMultimodalTokens from sglang.srt.managers.schedule_batch import MultimodalDataItem, MultimodalInputs 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 +350,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 +1229,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 38b16e514..acfd8650c 100644 --- a/python/sglang/srt/models/utils.py +++ b/python/sglang/srt/models/utils.py @@ -34,7 +34,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 @@ -444,7 +444,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) @@ -485,7 +485,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/processors/base_processor.py b/python/sglang/srt/multimodal/processors/base_processor.py index 04c181aaa..b264dc400 100644 --- a/python/sglang/srt/multimodal/processors/base_processor.py +++ b/python/sglang/srt/multimodal/processors/base_processor.py @@ -20,7 +20,7 @@ from sglang.srt.managers.schedule_batch import ( MultimodalProcessorOutput, ) from sglang.srt.multimodal.processors.executor import MultimodalProcessorExecutor -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import get_device, get_exec, get_mm, get_serving from sglang.srt.utils import ( envs, is_cpu, @@ -205,7 +205,7 @@ class BaseMultimodalProcessor(ABC): self.disable_fast_image_processor = server_args.disable_fast_image_processor self.skip_tokenizer_init = server_args.skip_tokenizer_init - mm_process_config = self.server_args.mm_process_config + mm_process_config = get_mm().mm_process_config self.image_config = mm_process_config.get("image", {}) self.video_config = mm_process_config.get("video", {}) self.audio_config = mm_process_config.get("audio", {}) @@ -337,7 +337,7 @@ class BaseMultimodalProcessor(ABC): # SGLANG_MM_FEATURE_CACHE_MB is the total pool budget across all # tokenizer workers. Each worker gets an equal share so that adding # workers doesn't multiply the GPU-side footprint. - worker_num = self.server_args.tokenizer_worker_num + worker_num = get_serving().tokenizer_worker_num per_worker_pool_size = get_mm_feature_pool_size_per_worker( MM_FEATURE_CACHE_SIZE, worker_num ) @@ -347,7 +347,7 @@ class BaseMultimodalProcessor(ABC): "GPU %d (%.0f MiB per tokenizer worker × %d; configured " "budget %.0f MiB).", total_pool_size / (1024 * 1024), - self.server_args.base_gpu_id, + get_device().base_gpu_id, per_worker_pool_size / (1024 * 1024), worker_num, MM_FEATURE_CACHE_SIZE / (1024 * 1024), @@ -355,7 +355,7 @@ class BaseMultimodalProcessor(ABC): self.cudaipc_mmfeature_pool = MmItemMemoryPool( per_worker_pool_size, MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL, - self.server_args.base_gpu_id, + get_device().base_gpu_id, ) def compute_mrope_positions(self, input_ids, mm_items): @@ -528,12 +528,12 @@ class BaseMultimodalProcessor(ABC): and isinstance(processor.image_processor, BaseImageProcessor) and not self.disable_fast_image_processor ): - if _is_cpu or get_server_args().rl_on_policy_target is not None: + if _is_cpu or get_exec().deterministic.rl_on_policy_target is not None: kwargs["device"] = "cpu" elif _is_xpu: kwargs["device"] = "xpu" elif not _is_npu: - base_gpu_id = get_server_args().base_gpu_id + base_gpu_id = get_device().base_gpu_id kwargs["device"] = f"cuda:{base_gpu_id}" elif processor.__class__.__name__ not in { "Glm4vProcessor", diff --git a/python/sglang/srt/multimodal/processors/kimi_k25.py b/python/sglang/srt/multimodal/processors/kimi_k25.py index e15935620..8167604d6 100644 --- a/python/sglang/srt/multimodal/processors/kimi_k25.py +++ b/python/sglang/srt/multimodal/processors/kimi_k25.py @@ -8,9 +8,7 @@ import torch import torch.nn.functional as F from PIL import Image -from sglang.srt.managers.schedule_batch import ( - MultimodalProcessorOutput, -) +from sglang.srt.managers.schedule_batch import MultimodalProcessorOutput from sglang.srt.models.kimi_k25 import KimiK25ForConditionalGeneration from sglang.srt.multimodal.processors.base_processor import ( BaseMultimodalProcessor as SGLangBaseProcessor, @@ -19,6 +17,7 @@ from sglang.srt.multimodal.processors.base_processor import ( MultimodalSpecialTokens, ) from sglang.srt.multimodal.processors.kimi_common import KimiGridMMDataMixin +from sglang.srt.runtime_context import get_mm from sglang.srt.utils.cuda_ipc_transport_utils import ( DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY, ) @@ -464,7 +463,7 @@ class KimiK2_5VLImageProcessor(KimiGridMMDataMixin, SGLangBaseProcessor): # its IPC proxy lazy until that assignment is known, avoiding a full # image copy to every rank. The scheduler only honors this marker once # the processor has already set the item's hash and pad value. - if self.use_cuda_ipc and self.server_args.mm_enable_dp_encoder: + if self.use_cuda_ipc and get_mm().mm_enable_dp_encoder: for item in mm_items: item.model_specific_data[DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY] = ( True 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/speculative/dflash_info_v2.py b/python/sglang/srt/speculative/dflash_info_v2.py index 8893a6541..4646cb7cb 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 a1095828f..a5fe3e1fd 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 @@ -329,7 +330,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) @@ -391,7 +392,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, ) @@ -1256,7 +1257,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 bb7106eda..3f7840bb3 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_exec, get_spec from sglang.srt.server_args import ServerArgs from sglang.srt.utils.common import ( cpu_has_amx_support, @@ -34,7 +35,7 @@ class DraftBackendFactory: else getattr(self.server_args, backend_name) ) if backend_type is None: - backend_type = self.server_args.attention_backend + backend_type = get_exec().kernel.attention_backend if backend_type not in backend_map: raise ValueError(error_template.format(backend_type=backend_type)) @@ -93,7 +94,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" ) return 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 e76206d5c..54f7df8a0 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_planner.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_planner.py @@ -16,7 +16,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 @@ -147,7 +147,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( @@ -163,16 +163,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, " @@ -371,7 +370,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 dfbff5730..b4eca59e5 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 @@ -294,7 +294,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: @@ -320,7 +320,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_info.py b/python/sglang/srt/speculative/eagle_info.py index 60c702d19..a997e76f5 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -8,7 +8,7 @@ from sglang.kernels.ops.attention.utils import create_flashinfer_kv_indices_trit from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject from sglang.srt.environ import envs 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__) @@ -202,7 +202,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 c4725ebcc..8f4c78442 100644 --- a/python/sglang/srt/speculative/eagle_utils.py +++ b/python/sglang/srt/speculative/eagle_utils.py @@ -17,14 +17,7 @@ from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import ( from sglang.srt.mem_cache.allocation import alloc_for_spec_decode from sglang.srt.mem_cache.allocation_sizing import get_alloc_reserve_per_decode from sglang.srt.runtime_context import get_parallel, get_spec -from sglang.srt.utils import ( - is_cpu, - is_cuda, - is_hip, - is_musa, - is_npu, - is_xpu, -) +from sglang.srt.utils import is_cpu, is_cuda, is_hip, is_musa, is_npu, is_xpu from sglang.srt.utils.async_probe import maybe_detect_oob if TYPE_CHECKING: @@ -488,9 +481,7 @@ def eagle_prepare_for_verify( batch: ScheduleBatch, target_worker: TpModelWorker, ): - from sglang.kernels.ops.speculative.cache_locs import ( - assign_extend_cache_locs_func, - ) + from sglang.kernels.ops.speculative.cache_locs import assign_extend_cache_locs_func from sglang.srt.model_executor.forward_batch_info import ( CaptureHiddenMode, ForwardBatch, @@ -579,10 +570,7 @@ def eagle_sample( import torch.nn.functional as F from sglang.srt.distributed import get_tp_group - from sglang.srt.layers.dp_attention import ( - is_dp_attention_enabled, - ) - from sglang.srt.runtime_context import get_server_args + from sglang.srt.layers.dp_attention import is_dp_attention_enabled from sglang.srt.sampling.penaltylib.repetition_penalty import ( apply_scaling_penalties, ) @@ -672,7 +660,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( @@ -842,9 +830,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 3c743f7fb..c178734ec 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -21,9 +21,7 @@ from sglang.srt.layers.attention.flashinfer_backend import FlashInferAttnBackend from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend from sglang.srt.layers.attention.triton_backend import TritonAttnBackend from sglang.srt.layers.attention.trtllm_mha_backend import TRTLLMHAAttnBackend -from sglang.srt.layers.attention.trtllm_mla_backend import ( - TRTLLMMLABackend, -) +from sglang.srt.layers.attention.trtllm_mla_backend import TRTLLMMLABackend from sglang.srt.layers.moe.utils import ( speculative_moe_a2a_backend_context, speculative_moe_backend_context, @@ -43,7 +41,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 +140,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 @@ -195,7 +199,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] @@ -288,13 +292,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=( @@ -379,7 +383,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 = { @@ -389,7 +393,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() @@ -586,7 +590,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 = ( @@ -601,7 +605,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), @@ -666,7 +670,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, @@ -692,7 +696,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) @@ -712,7 +716,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 ) @@ -832,7 +836,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, @@ -978,7 +982,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, @@ -995,7 +999,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 @@ -1012,7 +1016,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 @@ -1115,7 +1119,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 ), ) @@ -1418,13 +1422,30 @@ class EAGLEWorkerV2(BaseSpecWorker): state.target_graph_runner ) - # Sync server_args - self.server_args.override( + # Sync the step/draft-token counts on both config stores. + self._apply_adaptive_config( "adaptive_spec.restore", speculative_num_steps=state.speculative_num_steps, speculative_num_draft_tokens=state.speculative_num_draft_tokens, ) + def _apply_adaptive_config(self, source: str, **fields) -> None: + """Rebind adaptive-spec config on both stores that back these fields. + + Adaptive speculative decoding temporarily changes the step / draft-token + counts (and the decode graph-capture knobs) while it prebuilds per-step + runtime states. Those fields are read from two places: the resolved + config bag (runtime readers such as ``tp_worker`` / ``tokenizer_manager`` + via ``get_spec()``) *and* ``server_args`` — attention backends and graph + runners read the count as ``model_runner.server_args.`` and are + not migrated to the bag. Write both so an in-flight CUDA-graph capture + and later reads agree; updating only the bag leaves a rebuilt attention + backend sizing its metadata from the stale ``server_args`` count while + the graph is captured for the overridden count, which corrupts the + capture (illegal memory access).""" + get_context().override(source, **fields) + self.server_args.override(source, **fields) + @contextlib.contextmanager def _override_worker_state( self, @@ -1456,7 +1477,7 @@ class EAGLEWorkerV2(BaseSpecWorker): 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( + self._apply_adaptive_config( "adaptive_spec.capture_override", speculative_num_steps=speculative_num_steps, speculative_num_draft_tokens=speculative_num_draft_tokens, @@ -1466,7 +1487,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( + self._apply_adaptive_config( "adaptive_spec.capture_override", cuda_graph_bs_decode=cuda_graph_bs, **({"disable_cuda_graph": True} if not cuda_graph_bs else {}), @@ -1488,7 +1509,7 @@ class EAGLEWorkerV2(BaseSpecWorker): dw.cuda_graph_runner, dw.cuda_graph_runner_for_draft_extend, ) = backup[:10] - sa.override( + self._apply_adaptive_config( "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_worker_v2.py b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py index 98c210b4e..0dd5f5483 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py @@ -42,6 +42,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.pool_configurator import MemoryPoolConfig +from sglang.srt.runtime_context import get_exec, get_spec from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.base_spec_worker import EagleDraftWorkerBase from sglang.srt.speculative.eagle_utils import ( @@ -223,9 +224,9 @@ class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker): def _resolve_draft_backend_type(self) -> str: return ( - self.server_args.speculative_draft_attention_backend - or self.server_args.decode_attention_backend - or self.server_args.attention_backend + get_spec().speculative_draft_attention_backend + or get_exec().kernel.decode_attention_backend + or get_exec().kernel.attention_backend ) def _init_draft_attn_backend(self): diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index 4671818d7..29419d12a 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -44,7 +44,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, @@ -801,7 +801,7 @@ def commit_mamba_states_after_verify( # we need to update the mamba state for the request at the crossing point. 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 diff --git a/python/sglang/srt/state_capturer/indexer_topk.py b/python/sglang/srt/state_capturer/indexer_topk.py index afa652cec..2dd33b3f4 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 from sglang.srt.state_capturer.base import BaseTopkCapturer logger = logging.getLogger(__name__) @@ -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/utils/profile_utils.py b/python/sglang/srt/utils/profile_utils.py index 693a7418d..4e3e24266 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