From 63833f8034fbfff45d1357af38aa4da568eb9a0a Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Sun, 9 Aug 2026 14:44:08 -0700 Subject: [PATCH] config: business code no longer reads the published ServerArgs (#34081) --- .../srt/hardware_backend/mlx/model_runner.py | 6 +- .../srt/layers/attention/dsa/dsa_indexer.py | 4 +- .../sglang/srt/layers/attention/dsa/utils.py | 7 +- .../attention/hybrid_linear_attn_backend.py | 22 +- .../attention/linear/inkling_sconv_backend.py | 8 +- .../linear/kernels/gdn_flashinfer.py | 6 +- .../attention/linear/kernels/kda_ptx.py | 6 +- .../attention/linear/lightning_backend.py | 11 +- python/sglang/srt/layers/cp/base.py | 13 +- python/sglang/srt/layers/dp_attention.py | 9 +- python/sglang/srt/layers/k3_ar_fusion.py | 15 +- .../moe/moe_runner/flashinfer_cutedsl.py | 12 +- python/sglang/srt/layers/utils/cp_utils.py | 8 +- python/sglang/srt/managers/schedule_batch.py | 26 +- .../batch_result_processor.py | 5 +- .../sglang/srt/mem_cache/mamba_radix_cache.py | 6 +- .../components/mamba_component.py | 7 +- python/sglang/srt/model_loader/loader.py | 4 +- python/sglang/srt/models/deepseek_v2.py | 5 +- python/sglang/srt/models/kimi_k25.py | 14 +- python/sglang/srt/models/kimi_k3.py | 12 +- python/sglang/srt/models/minimax_m2.py | 5 +- python/sglang/srt/models/sarvam_moe.py | 22 +- python/sglang/srt/runtime_context.py | 147 ++++++ python/sglang/srt/server_args.py | 29 +- python/sglang/srt/speculative/spec_utils.py | 5 +- .../srt/utils/cuda_ipc_transport_utils.py | 6 +- .../test_gdn_prefill_backend_policy.py | 13 +- .../model_loader/test_presharded_loader.py | 11 +- test/registered/unit/models/test_kimi_k25.py | 21 +- .../unit/models/test_kimi_k3_vision.py | 17 +- .../spec/test_ngram_mamba_verify_update.py | 3 - .../unit/test_context_accessor_shadowing.py | 191 ++++++++ .../unit/test_global_config_read_ratchet.py | 450 ++++++++++++++++-- test/registered/unit/test_runtime_context.py | 188 ++++++++ 35 files changed, 1101 insertions(+), 213 deletions(-) create mode 100644 test/registered/unit/test_context_accessor_shadowing.py diff --git a/python/sglang/srt/hardware_backend/mlx/model_runner.py b/python/sglang/srt/hardware_backend/mlx/model_runner.py index 86674adee..57e396afa 100644 --- a/python/sglang/srt/hardware_backend/mlx/model_runner.py +++ b/python/sglang/srt/hardware_backend/mlx/model_runner.py @@ -49,7 +49,9 @@ from sglang.srt.hardware_backend.mlx.kv_cache import ( uses_sliding_window_attention, ) from sglang.srt.mem_cache.memory_pool import ReqToTokenPool -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import ( + mamba_cache_chunk_size, +) logger = logging.getLogger(__name__) @@ -281,7 +283,7 @@ class MlxModelRunner: ): return None - chunk_size = get_server_args().mamba_cache_chunk_size + chunk_size = mamba_cache_chunk_size() track_len = prefix_len + (new_token_count // chunk_size) * chunk_size branching_len = getattr(req, "mamba_branching_seqlen", None) if ( diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index 1cf7a4c5f..05eb1819e 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -40,11 +40,11 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo is_in_tc_piecewise_cuda_graph, ) from sglang.srt.runtime_context import ( + configured_pp_size, get_device, get_exec, get_parallel, get_schedule, - get_server_args, ) from sglang.srt.state_capturer.indexer_topk import ( maybe_capture_indexer_topk, @@ -249,7 +249,7 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp): if _is_cuda: self.sm_count = deep_gemm.get_num_sms() self.half_device_sm_count = ceil_align(self.sm_count // 2, 8) - pp_size = get_server_args().pp_size + pp_size = configured_pp_size() self.logits_with_pp_recv = pp_size > 1 and not get_pp_group().is_last_rank else: self.logits_with_pp_recv = False diff --git a/python/sglang/srt/layers/attention/dsa/utils.py b/python/sglang/srt/layers/attention/dsa/utils.py index dcd023e80..abdacfd56 100644 --- a/python/sglang/srt/layers/attention/dsa/utils.py +++ b/python/sglang/srt/layers/attention/dsa/utils.py @@ -12,7 +12,10 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph 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_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_parallel, + process_model_config, +) from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip from sglang.srt.utils.common import ceil_align, ceil_div @@ -113,7 +116,7 @@ def is_dsa_enable_prefill_cp(): return False from sglang.srt.configs.model_config import is_deepseek_dsa, is_deepseek_v4 - hf_config = get_server_args().get_model_config().hf_config + hf_config = process_model_config().hf_config return is_deepseek_dsa(hf_config) or is_deepseek_v4(hf_config) 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 a3d8f1fae..54a06c73f 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -25,7 +25,11 @@ 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_exec, get_memory, get_server_args +from sglang.srt.runtime_context import ( + get_exec, + get_memory, + mamba_cache_chunk_size, +) from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput from sglang.srt.speculative.spec_info import SpecInput @@ -297,8 +301,8 @@ class MambaAttnBackendBase(AttentionBackend): lens_to_track = ( forward_batch.mamba_track_seqlens - forward_batch.extend_prefix_lens ) - mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size - aligned_len = (lens_to_track // mamba_cache_chunk_size) * mamba_cache_chunk_size + chunk_size = mamba_cache_chunk_size() + aligned_len = (lens_to_track // chunk_size) * chunk_size start_indices = query_start_loc[:-1] + aligned_len - conv_state_len start_indices = start_indices[forward_batch.mamba_track_mask] @@ -316,7 +320,7 @@ class MambaAttnBackendBase(AttentionBackend): """src/dst indices to track SSM states for prefix caching: aligned seqs cache last_recurrent_state, unaligned cache intermediate `h` at the last chunk boundary.""" - mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size + chunk_size = mamba_cache_chunk_size() # CPU to avoid kernel launches for the masking ops mamba_track_mask = forward_batch.mamba_track_mask.cpu() extend_seq_lens = forward_batch.extend_seq_lens.cpu() @@ -326,9 +330,9 @@ class MambaAttnBackendBase(AttentionBackend): prefix_lens = forward_batch.extend_prefix_lens.cpu() if isinstance(self, Mamba2AttnBackend): - num_h_states = extend_seq_lens // mamba_cache_chunk_size + num_h_states = extend_seq_lens // chunk_size else: - num_h_states = (extend_seq_lens - 1) // mamba_cache_chunk_size + 1 + num_h_states = (extend_seq_lens - 1) // chunk_size + 1 track_ssm_src_offset = torch.zeros_like(num_h_states) track_ssm_src_offset[1:] = torch.cumsum(num_h_states[:-1], dim=0) @@ -338,17 +342,17 @@ class MambaAttnBackendBase(AttentionBackend): offset_masked = track_ssm_src_offset[mamba_track_mask] dst_masked = mamba_track_indices[mamba_track_mask] - is_aligned = (lens_masked % mamba_cache_chunk_size) == 0 + is_aligned = (lens_masked % chunk_size) == 0 # Aligned: last_recurrent_state from ssm_states. track_ssm_final_src = mamba_cache_indices[mamba_track_mask][is_aligned] track_ssm_final_dst = dst_masked[is_aligned] # Unaligned: intermediate state from h. - # TODO: handle mamba_cache_chunk_size % page size != 0 + # TODO: handle chunk_size % page size != 0 not_aligned = ~is_aligned track_ssm_h_src = offset_masked[not_aligned] + ( - lens_masked[not_aligned] // mamba_cache_chunk_size + lens_masked[not_aligned] // chunk_size ) track_ssm_h_dst = dst_masked[not_aligned] diff --git a/python/sglang/srt/layers/attention/linear/inkling_sconv_backend.py b/python/sglang/srt/layers/attention/linear/inkling_sconv_backend.py index a4d0677d3..2ef578591 100644 --- a/python/sglang/srt/layers/attention/linear/inkling_sconv_backend.py +++ b/python/sglang/srt/layers/attention/linear/inkling_sconv_backend.py @@ -62,7 +62,11 @@ from sglang.srt.models.inkling_common.kernels.sconv import ( fused_extend_sconv_metadata, precompute_helion_extend_metadata, ) -from sglang.srt.runtime_context import get_exec, get_server_args, get_spec +from sglang.srt.runtime_context import ( + get_exec, + get_spec, + mamba_cache_chunk_size, +) from sglang.srt.speculative.eagle_info import EagleDraftExtendInput if TYPE_CHECKING: @@ -99,7 +103,7 @@ class InklingShortConvAttnBackend(ShortConvAttnBackend): # [n_layers, n_slots, conv_kernel - 1, conv_dim]. self._mamba_cache = self.req_to_token_pool.mamba_pool.mamba_cache self.conv_state_len: int = self.conv_states_shape[2] - self.mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size + self.mamba_cache_chunk_size = mamba_cache_chunk_size() # A plain table lookup is recordable; the unified pool's translate is an # allocator lookup and must stay in the out-of-graph replay prep. self._slot_gather_recordable = ( diff --git a/python/sglang/srt/layers/attention/linear/kernels/gdn_flashinfer.py b/python/sglang/srt/layers/attention/linear/kernels/gdn_flashinfer.py index f7f4d974d..7a9e75da5 100644 --- a/python/sglang/srt/layers/attention/linear/kernels/gdn_flashinfer.py +++ b/python/sglang/srt/layers/attention/linear/kernels/gdn_flashinfer.py @@ -19,7 +19,9 @@ import torch from sglang.srt.layers.attention.linear.kernels.kernel_backend import ( LinearAttnKernelBase, ) -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import ( + mamba_cache_chunk_size, +) from sglang.srt.utils import is_cuda if TYPE_CHECKING: @@ -50,7 +52,7 @@ def maybe_build_flashinfer_checkpoint_plan( ): return - checkpoint_every_n_tokens = get_server_args().mamba_cache_chunk_size + checkpoint_every_n_tokens = mamba_cache_chunk_size() extend_seq_lens = forward_batch.extend_seq_lens.to(device="cpu", dtype=torch.int64) track_mask = forward_batch.mamba_track_mask.to(device="cpu", dtype=torch.bool) relative_track_lens = forward_batch.mamba_track_seqlens.to( diff --git a/python/sglang/srt/layers/attention/linear/kernels/kda_ptx.py b/python/sglang/srt/layers/attention/linear/kernels/kda_ptx.py index 2ea19c465..57cd2fdb3 100644 --- a/python/sglang/srt/layers/attention/linear/kernels/kda_ptx.py +++ b/python/sglang/srt/layers/attention/linear/kernels/kda_ptx.py @@ -37,6 +37,7 @@ from sglang.srt.layers.attention.linear.kernels.kda_triton import TritonKDAKerne from sglang.srt.layers.attention.linear.kernels.kernel_backend import ( LinearAttnKernelBase, ) +from sglang.srt.runtime_context import mamba_cache_chunk_size logger = logging.getLogger(__name__) @@ -194,11 +195,8 @@ class PtxKDAKernel(LinearAttnKernelBase): # which may be larger (for example 256 for Nemotron-H). Returning # the raw 64-token rows in that case would silently select the # wrong boundary and compute wrong offsets for later sequences. - from sglang.srt.runtime_context import get_server_args - intermediate_stride_supported = ( - get_server_args().mamba_cache_chunk_size == _CHUNK - ) + intermediate_stride_supported = mamba_cache_chunk_size() == _CHUNK seq_lens = ( [int(length) for length in seq_lens_cpu] if seq_lens_cpu is not None else [] ) diff --git a/python/sglang/srt/layers/attention/linear/lightning_backend.py b/python/sglang/srt/layers/attention/linear/lightning_backend.py index 6cd3c629e..e5f867caa 100644 --- a/python/sglang/srt/layers/attention/linear/lightning_backend.py +++ b/python/sglang/srt/layers/attention/linear/lightning_backend.py @@ -15,7 +15,10 @@ from sglang.srt.layers.attention.linear.linear_metadata import ( from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.model_runner import ModelRunner -from sglang.srt.runtime_context import get_parallel, get_server_args +from sglang.srt.runtime_context import ( + get_parallel, + mamba_cache_chunk_size, +) logger = logging.getLogger(__name__) @@ -299,7 +302,7 @@ class LightningAttentionBackend(MambaAttnBackendBase): if h_dst is None or h_dst.numel() == 0: return None, None - mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size + chunk_size = mamba_cache_chunk_size() num_prefills = metadata.num_prefills track_mask = forward_batch.mamba_track_mask[:num_prefills] extend_lens = forward_batch.extend_seq_lens[:num_prefills] @@ -307,9 +310,7 @@ class LightningAttentionBackend(MambaAttnBackendBase): track_seqlens = forward_batch.mamba_track_seqlens[:num_prefills] lens_to_track = track_seqlens - prefix_lens - boundary_lens = ( - lens_to_track // mamba_cache_chunk_size - ) * mamba_cache_chunk_size + boundary_lens = (lens_to_track // chunk_size) * chunk_size track_rows = (track_mask & (boundary_lens < extend_lens)).nonzero( as_tuple=True )[0] diff --git a/python/sglang/srt/layers/cp/base.py b/python/sglang/srt/layers/cp/base.py index d0070788d..5134a4bc7 100644 --- a/python/sglang/srt/layers/cp/base.py +++ b/python/sglang/srt/layers/cp/base.py @@ -230,12 +230,15 @@ class ContextParallelStrategy(ABC): def _is_dsa_active() -> bool: - from sglang.srt.runtime_context import get_server_args + from sglang.srt.runtime_context import get_parallel, get_server_args - sa = get_server_args() + # `_is_dsa_model_arch` is set nowhere in the tree, so this predicate is + # inert today (the getattr default makes it False). Kept verbatim rather + # than "fixed" here, because deciding what it should name is the CP path's + # call; the ratchet exempts it with that reason. return bool( - getattr(sa, "enable_prefill_cp", False) - and getattr(sa, "_is_dsa_model_arch", False) + get_parallel().enable_prefill_cp + and getattr(get_server_args(), "_is_dsa_model_arch", False) ) @@ -288,7 +291,7 @@ def get_cp_strategy() -> Optional[ContextParallelStrategy]: server_args = get_server_args() except ValueError: return None - if server_args is not None and getattr(server_args, "enable_prefill_cp", False): + if server_args is not None and get_parallel().enable_prefill_cp: init_cp_strategy(server_args) return _STRATEGY diff --git a/python/sglang/srt/layers/dp_attention.py b/python/sglang/srt/layers/dp_attention.py index eb8f4e7df..a89a839d0 100644 --- a/python/sglang/srt/layers/dp_attention.py +++ b/python/sglang/srt/layers/dp_attention.py @@ -27,7 +27,11 @@ from sglang.srt.distributed import ( from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) -from sglang.srt.runtime_context import get_flags, get_server_args +from sglang.srt.runtime_context import ( + configured_attn_cp_size, + configured_moe_dp_size, + get_flags, +) from sglang.srt.utils import get_bool_env_var, is_hip if TYPE_CHECKING: @@ -986,8 +990,7 @@ def is_enable_moe_cp_allgather() -> bool: (``parallel_state.py``), so the live sizes are equal and the comparison would always be false. """ - server_args = get_server_args() - return server_args.attn_cp_size > server_args.moe_dp_size + return configured_attn_cp_size() > configured_moe_dp_size() def moe_cp_all_gather_into_tensor(output: torch.Tensor, input: torch.Tensor): diff --git a/python/sglang/srt/layers/k3_ar_fusion.py b/python/sglang/srt/layers/k3_ar_fusion.py index c15eb6c11..86556c27f 100644 --- a/python/sglang/srt/layers/k3_ar_fusion.py +++ b/python/sglang/srt/layers/k3_ar_fusion.py @@ -67,16 +67,17 @@ def _get_state() -> Optional[_State]: if get_device_sm() not in (100, 103): return None - server_args = ctx.get_server_args() - if server_args.enable_symm_mem or server_args.moe_a2a_backend != "none": + symm_mem = ctx.get_exec().comm.enable_symm_mem + a2a_backend = ctx.get_exec().moe.moe_a2a_backend + if symm_mem or a2a_backend != "none": logger.info( "K3 all-reduce fusion auto-probe: skipping " "(enable_symm_mem=%s, moe_a2a_backend=%s; under symm-mem the " "allocator contexts conflict, and under EP a2a the model's " "symm-pool allocation contract does not hold on every AR " "call-site. Set SGLANG_K3_AR_FUSION=1 to force.)", - server_args.enable_symm_mem, - server_args.moe_a2a_backend, + symm_mem, + a2a_backend, ) return None from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import ( @@ -131,11 +132,11 @@ _BUFS: List[_Buffer] = [] @cache_once def _max_buffer_rows() -> int: """Rows to reserve per buffer: the largest batch the server args allow.""" - server_args = ctx.get_server_args() - chunked = server_args.chunked_prefill_size + schedule = ctx.get_schedule() + chunked = schedule.chunked_prefill_size if chunked is not None and chunked > 0: return int(chunked) - return int(server_args.max_prefill_tokens or 0) + return int(schedule.max_prefill_tokens or 0) def _create_buffer( diff --git a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py index c4462e6c6..851cf57b5 100644 --- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py +++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_cutedsl.py @@ -12,7 +12,10 @@ from sglang.srt.layers.moe.moe_runner.base import ( register_fused_func, ) from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import ( + cutedsl_moe_max_num_tokens, + get_parallel, +) from sglang.srt.utils.common import log_info_on_rank0, print_warning_once if TYPE_CHECKING: @@ -256,14 +259,11 @@ def ensure_cutedsl_wrapper(layer: torch.nn.Module) -> None: "Install with: pip install flashinfer" ) from e - from sglang.srt.runtime_context import get_server_args - assert layer.intermediate_size_per_partition > 0, ( f"CuteDSL MoE: intermediate_size_per_partition must be > 0, " f"got {layer.intermediate_size_per_partition}. Check EP/TP configuration." ) - server_args = get_server_args() # CuteDSL wrapper preallocates CG buffers used by any captured graph # that routes through this MoE — decode and prefill alike. use_cuda_graph = not cuda_graph_fully_disabled() @@ -277,9 +277,7 @@ def ensure_cutedsl_wrapper(layer: torch.nn.Module) -> None: else: # Standard allgather path: the MoE sees up to dp_size local forwards # gathered together, so scale the per-rank forward bound by dp_size. - max_num_tokens = ( - get_parallel().dp_size * server_args.cutedsl_moe_max_num_tokens() - ) + max_num_tokens = get_parallel().dp_size * cutedsl_moe_max_num_tokens() top_k = layer.top_k if layer.top_k is not None else layer.moe_runner_config.top_k # inference_mode(False) ensures the wrapper's pre-allocated CUDA-graph # buffers are normal tensors. This call typically happens inside diff --git a/python/sglang/srt/layers/utils/cp_utils.py b/python/sglang/srt/layers/utils/cp_utils.py index 553c51708..08abfcd3e 100644 --- a/python/sglang/srt/layers/utils/cp_utils.py +++ b/python/sglang/srt/layers/utils/cp_utils.py @@ -15,7 +15,10 @@ from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.moe import get_moe_a2a_backend 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, get_server_args +from sglang.srt.runtime_context import ( + get_parallel, + uses_mla_backend, +) @dataclass @@ -69,8 +72,7 @@ def is_prefill_cp_in_seq_split(): def is_mla_prefill_cp_enabled() -> bool: - sa = get_server_args() - return get_parallel().enable_prefill_context_parallel and sa.use_mla_backend() + return get_parallel().enable_prefill_context_parallel and uses_mla_backend() def mla_use_prefill_cp(forward_batch, mla_enable_prefill_cp=None): diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 78a534bff..112378567 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -7,6 +7,7 @@ from sglang.srt.runtime_context import ( get_schedule, get_serving, get_spec, + mamba_cache_chunk_size, mamba_extra_buffer_enabled, mamba_extra_buffer_lazy_enabled, ) @@ -118,7 +119,7 @@ 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 from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.server_args import ServerArgs @@ -2616,21 +2617,20 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): self, req: Req, ) -> _MambaRadixCacheV2TrackEntry: - server_args = get_server_args() - mamba_cache_chunk_size = server_args.mamba_cache_chunk_size + chunk_size = mamba_cache_chunk_size() def _force_track_h(i: int) -> int: - assert i % mamba_cache_chunk_size == 0 + assert i % chunk_size == 0 # There are 3 cases for mamba_track_seqlen passed to mamba_track_seqlens_cpu: - # 1) aligned with mamba_cache_chunk_size-> retrieve from last_recurrent_state + # 1) aligned with chunk_size-> retrieve from last_recurrent_state # a) is the last position -> retrieve from last_recurrent_state # b) is NOT the last position -> retrieve from h - # 2) unaligned with mamba_cache_chunk_size -> retrieve from h + # 2) unaligned with chunk_size -> retrieve from h # Currently, the math calculation only supports case 1a and 2. So for 1b, we need to add 1 # to force the math calculation to retrieve the correct mamba state from h. return i + 1 - mask = req.extend_range.length >= mamba_cache_chunk_size + mask = req.extend_range.length >= chunk_size track_index = req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx].item() mamba_track_seqlen = -1 if mask: @@ -2647,18 +2647,16 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): # mamba radix cache to track which seqlen this mamba state should store at. mamba_track_seqlen_aligned = ( len(req.prefix_indices) - + (req.extend_range.length // mamba_cache_chunk_size) - * mamba_cache_chunk_size + + (req.extend_range.length // chunk_size) * chunk_size ) - # mamba_track_fla_chunk_aligned is the aligned seqlen based on mamba_cache_chunk_size + # mamba_track_fla_chunk_aligned is the aligned seqlen based on chunk_size # If mamba_track_fla_chunk_aligned != mamba_track_seqlen_aligned, which can be true when - # page_size > mamba_cache_chunk_size, we need to force the math calculation to retrieve the correct mamba state from h + # page_size > chunk_size, we need to force the math calculation to retrieve the correct mamba state from h # by _force_track_h() mamba_track_fla_chunk_aligned = ( len(req.prefix_indices) - + (req.extend_range.length // mamba_cache_chunk_size) - * mamba_cache_chunk_size + + (req.extend_range.length // chunk_size) * chunk_size ) if mamba_track_fla_chunk_aligned != mamba_track_seqlen_aligned: # We want to track mamba_track_seqlen_aligned, and it's not the last position, @@ -2679,7 +2677,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): # is within the current extend batch. branching_seqlen_aligned_mask = ( req.mamba_branching_seqlen - len(req.prefix_indices) - ) % mamba_cache_chunk_size == 0 + ) % chunk_size == 0 if ( req.mamba_branching_seqlen > len(req.prefix_indices) and req.mamba_branching_seqlen < mamba_track_seqlen 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 73998604e..c1e17365e 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -37,8 +37,8 @@ from sglang.srt.runtime_context import ( get_exec, get_memory, get_observability, - get_server_args, mamba_extra_buffer_lazy_enabled, + max_speculative_num_draft_tokens, ) from sglang.srt.speculative.base_spec_worker import BaseSpecWorker from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer @@ -1163,7 +1163,6 @@ class SchedulerBatchResultProcessor: keep_written_by_this_step = ( crossed and planned_pos == req.mamba_next_track_idx ) - server_args = get_server_args() other_idx = 1 - req.mamba_next_track_idx # Recompute the in-flight verify's plan (kv_committed_len is # frozen since its prepare, so the recompute is exact). @@ -1172,7 +1171,7 @@ class SchedulerBatchResultProcessor: ].item() == -1 and mamba_lazy_spec_in_window( req, get_exec().mamba.mamba_track_interval, - server_args.max_speculative_num_draft_tokens, + max_speculative_num_draft_tokens(), ) if ( planned_pos is None diff --git a/python/sglang/srt/mem_cache/mamba_radix_cache.py b/python/sglang/srt/mem_cache/mamba_radix_cache.py index 6daf9b16c..bd9538a93 100644 --- a/python/sglang/srt/mem_cache/mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/mamba_radix_cache.py @@ -50,7 +50,9 @@ from sglang.srt.mem_cache.multi_ended_allocator import ( ) from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.mem_cache.utils import split_node_hash_value -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import ( + mamba_cache_chunk_size, +) if TYPE_CHECKING: from sglang.srt.managers.schedule_batch import Req @@ -450,7 +452,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): ) self.req_to_token_pool: HybridReqToTokenPool = params.req_to_token_pool self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator - self.mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size + self.mamba_cache_chunk_size = mamba_cache_chunk_size() self.page_size = params.page_size self.disable = params.disable diff --git a/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py b/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py index bf46238fc..1f96b6b1f 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py @@ -35,7 +35,10 @@ from sglang.srt.mem_cache.unified_cache.components.tree_component import ( TreeComponent, get_and_increase_time_counter, ) -from sglang.srt.runtime_context import get_exec, get_server_args +from sglang.srt.runtime_context import ( + get_exec, + mamba_cache_chunk_size, +) if TYPE_CHECKING: from sglang.srt.managers.schedule_batch import Req @@ -65,7 +68,7 @@ class MambaComponent(TreeComponent): params.page_size == 1 ), f"MambaComponent requires page_size=1 when mamba_extra_buffer is disabled, got {params.page_size}" super().__init__(cache, params) - self.mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size + self.mamba_cache_chunk_size = mamba_cache_chunk_size() self.mamba_max_states_per_path = get_exec().mamba.mamba_max_states_per_path # HiCache state self._mamba_pool_host = None # set to host mamba pool when HiCache enabled diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index 9338d3690..cb941a800 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -48,6 +48,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import ( register_memory_region, ) from sglang.srt.runtime_context import ( + configured_moe_dp_size, get_exec, get_model, get_parallel, @@ -1776,14 +1777,13 @@ class PreshardedModelLoader(DefaultModelLoader): return 1 parallel = get_parallel() - server_args = get_server_args() return { "tp": _safe(lambda: parallel.tp_size), "dp": _safe(lambda: parallel.moe_dp_size), "ep": _safe(lambda: parallel.moe_ep_size), "pp": _safe(lambda: parallel.pp_size), "moe_dense_tp_size": parallel.moe_dense_tp_size, - "moe_dp_size": server_args.moe_dp_size, + "moe_dp_size": configured_moe_dp_size(), "enable_dp_lm_head": parallel.enable_dp_lm_head, "enable_fp32_lm_head": get_exec().features.enable_fp32_lm_head, "quantization": model_config.quantization, diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 8c424e242..ce3c75852 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -193,12 +193,12 @@ from sglang.srt.models.deepseek_common.utils import ( is_wint4afp8_or_wint4a16_config, ) from sglang.srt.runtime_context import ( + attention_backends, get_device, get_exec, get_forward, get_model, get_parallel, - get_server_args, get_spec, ) from sglang.srt.speculative.spec_info import SpeculativeAlgorithm @@ -1999,8 +1999,7 @@ class DeepseekV2AttentionMLA( # Determine attention backend name for current forward batch: prefer the # name stamped per-runner on the backend object, else resolve from server args. backend = get_attn_backend() - server_args = get_server_args() - default_prefill_str, default_decode_str = server_args.get_attention_backends() + default_prefill_str, default_decode_str = attention_backends() prefill_backend_str = ( backend.prefill_attention_backend_str or default_prefill_str ) diff --git a/python/sglang/srt/models/kimi_k25.py b/python/sglang/srt/models/kimi_k25.py index 80a1ae4b8..31e78d992 100644 --- a/python/sglang/srt/models/kimi_k25.py +++ b/python/sglang/srt/models/kimi_k25.py @@ -46,10 +46,9 @@ from sglang.srt.multimodal.mm_utils import ( run_dp_sharded_mrope_vision_model, ) from sglang.srt.runtime_context import ( + configured_tp_size, get_exec, get_mm, - get_parallel, - get_server_args, ) from sglang.srt.utils import add_prefix, is_cuda, is_npu @@ -731,13 +730,10 @@ class KimiK25ForConditionalGeneration(nn.Module): acknowledges the entire TP group so the bounded IPC pool remains recyclable. """ - parallel = get_parallel() - server_args = get_server_args() - # Match MmItemMemoryPool.try_to_recycle(), which waits for the - # server TP size rather than the attention subgroup size. - ipc_consumer_count = max( - getattr(server_args, "tp_size", parallel.attn_tp_size), 1 - ) + # Same source as MmItemMemoryPool.try_to_recycle(), which waits on + # configured_tp_size(): the live world size agrees once dist is up, + # but a refcount that disagrees with the waiter would strand items. + ipc_consumer_count = max(configured_tp_size(), 1) device_index = device.index if device.type == "cuda" and device_index is None: device_index = torch.cuda.current_device() diff --git a/python/sglang/srt/models/kimi_k3.py b/python/sglang/srt/models/kimi_k3.py index 3adb1831a..d1495c9ce 100644 --- a/python/sglang/srt/models/kimi_k3.py +++ b/python/sglang/srt/models/kimi_k3.py @@ -111,7 +111,12 @@ from sglang.srt.multimodal.kimi_k3_image_processing import ( to_chw_uint8, ) from sglang.srt.multimodal.mm_utils import materialize_multimodal_features -from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args +from sglang.srt.runtime_context import ( + configured_tp_size, + get_exec, + get_parallel, + get_server_args, +) from sglang.srt.utils import is_blackwell_supported, is_hip, make_layers from sglang.srt.utils.common import ( BumpAllocator, @@ -3082,7 +3087,10 @@ class KimiK3ForConditionalGeneration(nn.Module): def materialize_item_features(image_indices: List[int]) -> torch.Tensor: """Materialize only the images assigned to this vision-DP rank.""" - ipc_consumer_count = max(get_parallel().tp_size, 1) + # Same source as MmItemMemoryPool.try_to_recycle(), which waits on + # configured_tp_size(): the live world size agrees once dist is up, + # but a refcount that disagrees with the waiter would strand items. + ipc_consumer_count = max(configured_tp_size(), 1) device_index = device.index if device.type == "cuda" and device_index is None: device_index = torch.cuda.current_device() diff --git a/python/sglang/srt/models/minimax_m2.py b/python/sglang/srt/models/minimax_m2.py index 3dd4e3e9f..6fb388ee1 100644 --- a/python/sglang/srt/models/minimax_m2.py +++ b/python/sglang/srt/models/minimax_m2.py @@ -85,7 +85,7 @@ from sglang.srt.runtime_context import ( get_forward, get_parallel, get_schedule, - get_server_args, + process_model_config, ) # get_bool_env_var is defined in sglang.srt.utils.common, not sglang.srt.distributed. @@ -431,10 +431,9 @@ class MiniMaxM2QKRMSNorm: props = torch.cuda.get_device_properties(device) # probe the maximum tokens for one prefill - server_args = get_server_args() max_tokens = get_schedule().chunked_prefill_size if max_tokens is None: - max_tokens = server_args.model_config.context_len + max_tokens = process_model_config().context_len max_tokens = max(max_tokens, get_schedule().max_prefill_tokens) logger.info(f"[AR] Using CustomAllReduceV2 for MiniMaxM2 with {max_tokens = }") ALIGN = 512 diff --git a/python/sglang/srt/models/sarvam_moe.py b/python/sglang/srt/models/sarvam_moe.py index 087cf2e5f..e54dd6e79 100644 --- a/python/sglang/srt/models/sarvam_moe.py +++ b/python/sglang/srt/models/sarvam_moe.py @@ -61,8 +61,10 @@ from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha imp DeepseekMHAForwardMixin, ) from sglang.srt.runtime_context import ( + attention_backends, get_exec, get_forward, + get_memory, get_model, get_parallel, get_server_args, @@ -609,18 +611,12 @@ class SarvamMoEMLAAttention(nn.Module): return k def _set_current_attention_backend(self, forward_batch: ForwardBatch) -> None: - if self._server_args is None: - self._server_args = get_server_args() - if forward_batch.forward_mode.is_decode_or_idle(): - self.current_attention_backend = ( - self._server_args.decode_attention_backend - or self._server_args.attention_backend - ) - else: - self.current_attention_backend = ( - self._server_args.prefill_attention_backend - or self._server_args.attention_backend - ) + prefill_backend, decode_backend = attention_backends() + self.current_attention_backend = ( + decode_backend + if forward_batch.forward_mode.is_decode_or_idle() + else prefill_backend + ) def _maybe_fp8_bmm( self, @@ -684,7 +680,7 @@ class SarvamMoEMLAAttention(nn.Module): ) self._set_current_attention_backend(forward_batch) - can_use_prefix_cache = not self._server_args.disable_radix_cache + can_use_prefix_cache = not get_memory().disable_radix_cache do_prefix_merge = has_extend_prefix and can_use_prefix_cache if do_prefix_merge and forward_batch.num_prefix_chunks is None: diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index c780ddb61..386f7c5ac 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -47,6 +47,7 @@ test-only ``override(**kw)``. from __future__ import annotations import dataclasses +import functools import os import sys from contextlib import contextmanager @@ -827,6 +828,10 @@ class RuntimeContext: server_args, "enable_torch_compile", False ) self._server_args = server_args + # The adaptive draft-token bound memoizes on the config *path*, so a new + # publication that reuses the path must not inherit the bound computed + # from the file's previous contents. + _adaptive_draft_token_bound.cache_clear() # Snapshot resolved config into the namespace bags (the single source of # truth for config reads). Driven by NS(...) metadata; a mock/partial # config with no NS markers yields an empty tree (no bags projected). @@ -1376,6 +1381,7 @@ def reset_context() -> None: """ _CONTEXT._server_args = None _CONTEXT._config_bags = None + _adaptive_draft_token_bound.cache_clear() _CONTEXT._overrides_log = [] _CONTEXT._publish_role = None _CONTEXT.parallel._config = None @@ -1406,3 +1412,144 @@ def mamba_extra_buffer_lazy_enabled() -> bool: get_memory().disable_radix_cache is False and get_exec().mamba.mamba_radix_cache_strategy == "extra_buffer_lazy" ) + + +# --- Derived config accessors ------------------------------------------------ +# +# A few values are computed from several config fields plus the HF config, so +# they are ``ServerArgs`` members rather than namespace leaves. Business code +# must not reach for the startup record to get them: these accessors are the +# named home, and this module — which owns the slot — is the only place that +# reads it. Each one keeps the member's exact semantics, including which model +# config it derives from (always the process's, i.e. the target's). + + +def mamba_cache_chunk_size() -> int: + """The caching point granularity for mamba state: ``max(the model's mamba + chunk size, page_size)``. Cached on the config after the first call.""" + return get_server_args().mamba_cache_chunk_size + + +def max_speculative_num_draft_tokens() -> int | None: + """The largest draft-token count speculative decoding may use. + + All three inputs are ``spec`` leaves, so this derives from the bags and + follows a post-publish override; ``ServerArgs.max_speculative_num_draft_tokens`` + is the pre-publish equivalent. Adaptive spec resolves the count from its + candidate-step table instead of the flat field. + """ + spec = get_spec() + if spec.speculative_num_draft_tokens is None: + return None + if not spec.speculative_adaptive: + return spec.speculative_num_draft_tokens + # The adaptive branch parses a JSON config, and this is called per decode + # batch (`spec_prepare_for_decode`), so memoize on the inputs -- keyed, not + # cached once, so a post-publish override still recomputes. + return _adaptive_draft_token_bound(spec.speculative_adaptive_config) + + +@functools.lru_cache(maxsize=8) +def _adaptive_draft_token_bound(cfg_path: str | None) -> int: + from sglang.srt.speculative.adaptive_spec_params import ( + resolve_candidate_steps_from_config, + ) + + candidate_steps = resolve_candidate_steps_from_config(cfg_path=cfg_path) + # Adaptive spec requires topk=1 today, so each runtime state needs + # steps + 1 draft-token slots (mirrors the ServerArgs member). + return max(candidate_steps) + 1 + + +def uses_mla_backend() -> bool: + """Whether this process's model runs the MLA attention path.""" + return get_server_args().use_mla_backend() + + +def attention_backends() -> tuple: + """The configured ``(prefill, decode)`` backend pair, split fields falling + back to ``attention_backend``. + + All three inputs are ``exec.kernel`` leaves, so this derives from the bags + and follows a post-publish override; ``ServerArgs.get_attention_backends`` + is the pre-publish equivalent the resolution pipeline uses. A built runner + stamps its own resolved pair (``ModelRunner.prefill_attention_backend_str``); + read that when there is a runner in hand. + """ + from sglang.srt.arg_groups.overrides import attention_backends_of + + # All three leaves live in the same bag, so the resolution pipeline's own + # helper applies directly -- one definition of the fallback rule. + return attention_backends_of(get_exec().kernel) + + +def process_model_config(): + """The process's ``ModelConfig`` (built once from the published config).""" + return get_server_args().get_model_config() + + +def cutedsl_moe_max_num_tokens() -> int: + """The CuteDSL A2A per-rank token budget. + + Every input is a published leaf (``spec``, ``schedule``, ``exec.graph``), so + this derives from the bags and follows a post-publish override; + ``ServerArgs.cutedsl_moe_max_num_tokens`` is the pre-publish equivalent the + resolution pipeline uses. Max over the prefill bound, the piecewise-prefill + capture, and the decode/verify bound. + """ + from sglang.srt.model_executor.cuda_graph_config import Backend + + spec = get_spec() + num_tokens_per_req = ( + (spec.speculative_num_draft_tokens or 1) if spec.speculative_algorithm else 1 + ) + prefill_tokens = get_schedule().max_prefill_tokens + cg_config = get_exec().graph.cuda_graph_config + if cg_config is not None and cg_config.prefill.backend == Backend.TC_PIECEWISE: + prefill_tokens = max(prefill_tokens, cg_config.prefill.max_bs or 0) + decode_max_bs = (cg_config.decode.max_bs if cg_config is not None else 0) or 0 + return max(prefill_tokens, decode_max_bs * num_tokens_per_req) + + +# --- Configured (not live) parallel sizes ------------------------------------ +# +# ``get_parallel()`` shadows these names with the LIVE topology, which is the +# right answer almost everywhere. A handful of call sites need what was +# *configured* instead — before the groups exist, in a process that has none, +# or where the live value is deliberately aliased to another dimension. Each +# accessor below names that intent so no business call site has to reach for +# the startup record; the per-site reasons live in the read ratchet. +# +# They read the published leaf rather than the record: the bag is what +# ``override`` writes, and once the instance holds only the user's raw input +# the record would answer with what was *typed* instead of what resolution +# produced. Going through the bag directly is what gets past the live property +# that shadows these four names on ``get_parallel()``. + + +def _configured_parallel(name: str): + # The bag itself, not ParallelContext, whose live property shadows these + # four names. Read through the parallel slot the way the leaf accessor + # does — ``parallel`` is deliberately outside the per-role namespace table + # (every process reads topology config), so this must not route through + # ``config_bag()``'s role check, which would record or reject the read. + config = _CONTEXT.parallel._config + if config is None: + raise ValueError("config namespace 'parallel' not published") + return getattr(config, name) + + +def configured_tp_size() -> int: + return _configured_parallel("tp_size") + + +def configured_pp_size() -> int: + return _configured_parallel("pp_size") + + +def configured_moe_dp_size() -> int: + return _configured_parallel("moe_dp_size") + + +def configured_attn_cp_size() -> int: + return _configured_parallel("attn_cp_size") diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 536773628..3fa0eb22d 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -40,7 +40,12 @@ from sglang.srt.arg_groups.argparse_actions import ( DeprecatedStoreTrueAction, LoRAPathAction, ) -from sglang.srt.arg_groups.overrides import resolved_view +from sglang.srt.arg_groups.overrides import ( + attention_backends_of, + mamba_extra_buffer_lazy_of, + mamba_extra_buffer_of, + resolved_view, +) from sglang.srt.configs.embedding_model_spec import BCGPrefillPolicy from sglang.srt.configs.linear_attn_model_registry import get_linear_attn_spec_by_arch from sglang.srt.connector import ConnectorType @@ -8689,17 +8694,7 @@ class ServerArgs: return attention_backends_of(resolved_view(self)) def get_attention_backends(self): - prefill_attention_backend_str = ( - self.prefill_attention_backend - if self.prefill_attention_backend - else self.attention_backend - ) - decode_attention_backend_str = ( - self.decode_attention_backend - if self.decode_attention_backend - else self.attention_backend - ) - return prefill_attention_backend_str, decode_attention_backend_str + return attention_backends_of(self) def use_mla_backend(self): from sglang.srt.configs.model_config import AttentionArch @@ -8715,16 +8710,10 @@ class ServerArgs: ) def enable_mamba_extra_buffer(self) -> bool: - return ( - self.disable_radix_cache is False - and self.mamba_radix_cache_strategy in ("extra_buffer", "extra_buffer_lazy") - ) + return mamba_extra_buffer_of(self) def enable_mamba_extra_buffer_lazy(self) -> bool: - return ( - self.disable_radix_cache is False - and self.mamba_radix_cache_strategy == "extra_buffer_lazy" - ) + return mamba_extra_buffer_lazy_of(self) @cached_property def max_speculative_num_draft_tokens(self) -> Optional[int]: diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index 1c9044292..19a35c66b 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -48,10 +48,10 @@ from sglang.srt.mem_cache.allocation import ( ) from sglang.srt.runtime_context import ( get_exec, - get_server_args, get_spec, mamba_extra_buffer_enabled, mamba_extra_buffer_lazy_enabled, + max_speculative_num_draft_tokens, ) from sglang.srt.utils import ( is_cpu, @@ -1030,12 +1030,11 @@ def spec_prepare_for_decode(batch: ScheduleBatch) -> None: """eagle/ngram share a stateless free function; dflash keeps stateful prep on its draft input -- the dispatcher routes. """ - server_args = get_server_args() if mamba_extra_buffer_lazy_enabled(): # Scheduler phase (outside forward isolation). batch.mamba_lazy_spec_prepare( get_exec().mamba.mamba_track_interval, - server_args.max_speculative_num_draft_tokens, + max_speculative_num_draft_tokens(), ) if batch.spec_algorithm.is_dflash_family(): batch.spec_info.prepare_for_decode(batch) diff --git a/python/sglang/srt/utils/cuda_ipc_transport_utils.py b/python/sglang/srt/utils/cuda_ipc_transport_utils.py index 7a72f75c8..acc30e1ed 100644 --- a/python/sglang/srt/utils/cuda_ipc_transport_utils.py +++ b/python/sglang/srt/utils/cuda_ipc_transport_utils.py @@ -9,7 +9,9 @@ import numpy as np import torch from sglang.srt.environ import envs -from sglang.srt.runtime_context import get_server_args +from sglang.srt.runtime_context import ( + configured_tp_size, +) from sglang.srt.utils.stale_shm_cleanup import make_shm_name logger = logging.getLogger(__name__) @@ -130,7 +132,7 @@ class MmItemMemoryChunk: def try_to_recycle(self) -> bool: try: - tp_num = get_server_args().tp_size + tp_num = configured_tp_size() except Exception: logger.info( "server_args has not been published yet, skip this turn's recycle" diff --git a/test/registered/unit/layers/attention/test_gdn_prefill_backend_policy.py b/test/registered/unit/layers/attention/test_gdn_prefill_backend_policy.py index d10e0ae0f..63240dac7 100644 --- a/test/registered/unit/layers/attention/test_gdn_prefill_backend_policy.py +++ b/test/registered/unit/layers/attention/test_gdn_prefill_backend_policy.py @@ -18,6 +18,7 @@ from sglang.srt.layers.attention.linear.kernels.gdn_flashinfer import ( ) from sglang.srt.layers.attention.linear.kernels.gdn_triton import TritonGDNKernel from sglang.srt.layers.attention.linear.utils import LinearAttnKernelBackend +from sglang.srt.runtime_context import get_context from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=5, suite="base-a-test-cpu") @@ -151,12 +152,12 @@ class TestFlashInferGDNPrefillBackendPolicy(unittest.TestCase): track_ssm_h_dst=torch.empty(4), ) - with patch( - "sglang.srt.layers.attention.linear.kernels.gdn_flashinfer." - "get_server_args", - return_value=SimpleNamespace(mamba_cache_chunk_size=64), - ): - maybe_build_flashinfer_checkpoint_plan(forward_batch, metadata, "cpu") + # The chunk size is a derived config member; seed its private cache on a + # published config rather than patching an import binding. + override = get_context().override_server_args(_mamba_cache_chunk_size=64) + override.install() + self.addCleanup(override.restore) + maybe_build_flashinfer_checkpoint_plan(forward_batch, metadata, "cpu") torch.testing.assert_close( metadata.state_checkpoint_cu_starts, diff --git a/test/registered/unit/model_loader/test_presharded_loader.py b/test/registered/unit/model_loader/test_presharded_loader.py index 8a1391723..d1ff489e5 100644 --- a/test/registered/unit/model_loader/test_presharded_loader.py +++ b/test/registered/unit/model_loader/test_presharded_loader.py @@ -786,13 +786,6 @@ class TestShardConfig(unittest.TestCase): # `_collect_shard_config` is the exact failure mode that left # moe_dense_tp_size / LM-head flags out of the cache key before. loader = object.__new__(PreshardedModelLoader) - server_args = SimpleNamespace( - moe_dp_size=2, - enable_fp32_lm_head=True, - ep_num_redundant_experts=4, - enable_eplb=True, - init_expert_location="trivial", - ) model_config = SimpleNamespace(quantization="fp8", dtype=torch.bfloat16) required = { "tp", @@ -819,8 +812,8 @@ class TestShardConfig(unittest.TestCase): enable_dp_lm_head=True, ) with mock.patch( - "sglang.srt.model_loader.loader.get_server_args", - return_value=server_args, + "sglang.srt.model_loader.loader.configured_moe_dp_size", + return_value=2, ), mock.patch( "sglang.srt.model_loader.loader.get_parallel", return_value=parallel, diff --git a/test/registered/unit/models/test_kimi_k25.py b/test/registered/unit/models/test_kimi_k25.py index 425b2c13c..feec9e90a 100644 --- a/test/registered/unit/models/test_kimi_k25.py +++ b/test/registered/unit/models/test_kimi_k25.py @@ -39,7 +39,7 @@ from sglang.srt.multimodal.processors.kimi_k25 import ( _resize_bicubic_if_needed, _resize_images_by_source_shape, ) -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_context, get_parallel from sglang.srt.utils.cuda_ipc_transport_utils import ( DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY, CudaIpcTensorTransportProxy, @@ -317,7 +317,11 @@ def test_dp_helper_supports_moonvit3d_packed_embeddings_on_tp1(): tower = _MoonViT3dTower() pixel_values = torch.randn(4, 2) - with get_parallel().override(tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0): + # The IPC consumer count asks for the *configured* TP size (matching + # MmItemMemoryPool.try_to_recycle), so the double publishes one too. + with get_context().override_server_args(tp_size=1), get_parallel().override( + tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0 + ): output = run_dp_sharded_mrope_vision_model( tower, pixel_values, [[1, 2, 2]], rope_type="rope_2d_packed" ) @@ -331,7 +335,11 @@ def test_dp_helper_can_lazily_load_kimi_features_on_tp1(): pixel_values = torch.randn(4, 2) loader = Mock(return_value=pixel_values) - with get_parallel().override(tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0): + # The IPC consumer count asks for the *configured* TP size (matching + # MmItemMemoryPool.try_to_recycle), so the double publishes one too. + with get_context().override_server_args(tp_size=1), get_parallel().override( + tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0 + ): output = run_dp_sharded_mrope_vision_model( tower, None, @@ -479,11 +487,10 @@ def test_kimi_non_dp_keeps_grid_thws_on_the_host(): model.mm_projector = _IdentityProjector() items = [_image_item(torch.randn(4, 2), [[1, 2, 2]])] - with get_parallel().override( + # The IPC consumer count asks for the *configured* TP size (matching + # MmItemMemoryPool.try_to_recycle), so the double publishes one too. + with get_context().override_server_args(tp_size=1), get_parallel().override( tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0 - ), patch( - "sglang.srt.models.kimi_k25.get_server_args", - return_value=SimpleNamespace(tp_size=1), ): model.get_image_feature(items) diff --git a/test/registered/unit/models/test_kimi_k3_vision.py b/test/registered/unit/models/test_kimi_k3_vision.py index 4505baf62..d80508048 100644 --- a/test/registered/unit/models/test_kimi_k3_vision.py +++ b/test/registered/unit/models/test_kimi_k3_vision.py @@ -464,15 +464,17 @@ def test_kimi_k3_encoder_dp_defers_feature_materialization(monkeypatch): ] sharded_embeddings = torch.randn(2, 2) + # The IPC consumer count asks for the *configured* TP size (matching + # MmItemMemoryPool.try_to_recycle), so publish it; the live topology the + # sharding helper reads is forced through the context's own override. with mock_patch( "sglang.srt.multimodal.mm_utils.run_dp_sharded_mrope_vision_model", return_value=sharded_embeddings, - ) as run_dp, mock_patch( - "sglang.srt.models.kimi_k3.get_parallel", - return_value=SimpleNamespace(tp_size=1, attn_tp_size=1), + ) as run_dp, get_context().override_server_args(tp_size=1), get_parallel().override( + tp_size=1, attn_tp_size=1 ): output = model.get_image_feature(items) - # Exercise the loader while the runtime topology is patched. + # Exercise the loader while the runtime topology is forced. loader_in_scope = run_dp.call_args.kwargs["load_local_pixel_values"] local = loader_in_scope([1]) both = loader_in_scope([0, 1]) @@ -545,12 +547,13 @@ def test_kimi_k3_preprocesses_only_dp_owner_images(monkeypatch): calls.append([int(image[0, 0, 0]) for image in images]) return torch.tensor([[float(calls[-1][0]), 0.0]]), torch.tensor([[1, 1, 1]]) + # Configured TP size (the IPC consumer count) comes from the published + # bags; the live topology is forced through the context's own override. with mock_patch( "sglang.srt.multimodal.mm_utils.run_dp_sharded_mrope_vision_model", return_value=torch.zeros(1, 2), - ) as run_dp, mock_patch( - "sglang.srt.models.kimi_k3.get_parallel", - return_value=SimpleNamespace(tp_size=1, attn_tp_size=1), + ) as run_dp, get_context().override_server_args(tp_size=1), get_parallel().override( + tp_size=1, attn_tp_size=1 ), mock_patch( "sglang.srt.multimodal.processors.kimi_k25._gpu_preprocess_images", side_effect=fake_preprocess, diff --git a/test/registered/unit/spec/test_ngram_mamba_verify_update.py b/test/registered/unit/spec/test_ngram_mamba_verify_update.py index c9f70fa91..a43eb5d51 100644 --- a/test/registered/unit/spec/test_ngram_mamba_verify_update.py +++ b/test/registered/unit/spec/test_ngram_mamba_verify_update.py @@ -197,9 +197,6 @@ class TestNgramMambaVerifyUpdate(CustomTestCase): with patch( "sglang.srt.speculative.spec_utils.mambaish_config", return_value={"some": "config"}, - ), patch( - "sglang.srt.speculative.spec_utils.get_server_args", - return_value=MagicMock(mamba_track_interval=256), ), patch( "sglang.srt.speculative.spec_utils.get_exec", return_value=MagicMock(mamba=MagicMock(mamba_track_interval=256)), diff --git a/test/registered/unit/test_context_accessor_shadowing.py b/test/registered/unit/test_context_accessor_shadowing.py new file mode 100644 index 000000000..abb673342 --- /dev/null +++ b/test/registered/unit/test_context_accessor_shadowing.py @@ -0,0 +1,191 @@ +"""No local may shadow a ``runtime_context`` accessor it also calls. + +A mechanical sweep that rewrites ``self.server_args.mamba_cache_chunk_size`` +into ``mamba_cache_chunk_size()`` turns + + mamba_cache_chunk_size = self.server_args.mamba_cache_chunk_size + +into ``mamba_cache_chunk_size = mamba_cache_chunk_size()``, which is a +self-referential local: the name is local for the whole function, so the call +raises ``UnboundLocalError`` the first time that line runs. Five of these +shipped in one sweep and only one had unit coverage — a mamba model on the +radix-cache-v2 path found it at request time. + +This scans for the shape directly: a function-scope assignment whose target +name is an imported accessor. +""" + +import ast +import unittest +from pathlib import Path + +import sglang +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +_PACKAGE_ROOT = Path(next(iter(sglang.__path__))) +_CONTEXT_MODULE = "sglang.srt.runtime_context" + + +def _module_level_accessor_imports(tree: ast.AST) -> set[str]: + """Accessors imported at module scope — visible in every function. + + A *function-local* import is visible only inside its own scope, so it is + collected per function in the scan below: charging it file-wide would flag + an unrelated sibling function that binds the same name, where no shadowing + can occur. + """ + names: set[str] = set() + stack = list(tree.body) + while stack: + stmt = stack.pop() + if isinstance( + stmt, (ast.FunctionDef, ast.AsyncFunctionDef, ast.Lambda, ast.ClassDef) + ): + continue + if isinstance(stmt, ast.ImportFrom) and stmt.module == _CONTEXT_MODULE: + for alias in stmt.names: + names.add(alias.asname or alias.name) + stack.extend(ast.iter_child_nodes(stmt)) + return names + + +def _bound_names(target): + """Every name a binding target introduces, unpacking included. + + ``a, (b, c) = ...`` and ``for x, y in ...`` bind through Tuple/List/Starred + nodes, so a check that only accepts a bare ``ast.Name`` misses them. + """ + if isinstance(target, ast.Name): + yield target.id + elif isinstance(target, ast.Starred): + yield from _bound_names(target.value) + elif isinstance(target, (ast.Tuple, ast.List)): + for element in target.elts: + yield from _bound_names(element) + + +def _own_scope_statements(node) -> tuple: + """This function's OWN scope: its statements, plus the (name, lineno) of + each nested ``def``/``class`` — the definition's *name* is a binding in + this scope (an earlier accessor call raises UnboundLocalError just like an + assignment), while its *body* is the nested scope's own and descending into + it would misattribute bindings.""" + own_scope = [] + nested_def_bindings = [] + pending = list(node.body) + while pending: + stmt = pending.pop() + if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): + nested_def_bindings.append((stmt.name, stmt.lineno)) + continue + if isinstance(stmt, ast.Lambda): + continue + own_scope.append(stmt) + pending.extend(ast.iter_child_nodes(stmt)) + return own_scope, nested_def_bindings + + +def _child_functions(body) -> list: + """Function defs directly beneath this scope — descending through plain + statements and class bodies (a method closes over the enclosing function's + names, not the class's), but never into another function.""" + funcs = [] + pending = list(body) + while pending: + stmt = pending.pop() + if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef)): + funcs.append(stmt) + continue + if isinstance(stmt, ast.Lambda): + continue + pending.extend(ast.iter_child_nodes(stmt)) + return funcs + + +def _shadowing_assignments(tree: ast.AST, module_accessors: set[str]): + """Function-local bindings whose name shadows an accessor visible in that + scope -- every statement form that binds a local, not just ``=``. + + Python decides a name is local from *any* binding in the function, so a + loop variable, a ``with ... as``, a walrus, a comprehension target, or an + ``except ... as`` all shadow the accessor for the whole function body, + exactly like an assignment does. + + Visibility follows lexical scope: module-level imports reach every + function; a function-local import reaches its own scope and nested + functions (closure), but NOT an unrelated sibling — charging it file-wide + would flag bindings where no shadowing occurs. A function-scope *re-import* + of the accessor is itself fine: it binds the name to the same callable, so + calls after it behave identically (and the module is full of deliberate + local imports). + """ + stack = [(fn, module_accessors) for fn in _child_functions(tree.body)] + while stack: + node, inherited = stack.pop() + own_scope, nested_def_bindings = _own_scope_statements(node) + local_imports = { + alias.asname or alias.name + for stmt in own_scope + if isinstance(stmt, ast.ImportFrom) and stmt.module == _CONTEXT_MODULE + for alias in stmt.names + } + visible = inherited | local_imports + # ``def get_exec(): ...`` nested in the function binds the name in + # THIS scope, exactly like an assignment would. + for name, lineno in nested_def_bindings: + if name in visible: + yield node.name, name, lineno + for inner in own_scope: + targets = [] + if isinstance(inner, ast.Assign): + targets = inner.targets + elif isinstance(inner, (ast.AnnAssign, ast.AugAssign)): + targets = [inner.target] + elif isinstance(inner, (ast.For, ast.AsyncFor, ast.comprehension)): + targets = [inner.target] + elif isinstance(inner, ast.NamedExpr): + targets = [inner.target] + elif isinstance(inner, (ast.With, ast.AsyncWith)): + targets = [i.optional_vars for i in inner.items if i.optional_vars] + elif isinstance(inner, ast.ExceptHandler) and inner.name: + targets = [ast.Name(id=inner.name, ctx=ast.Store())] + for target in targets: + for name in _bound_names(target): + if name in visible: + yield node.name, name, getattr(inner, "lineno", node.lineno) + for nested in _child_functions(node.body): + stack.append((nested, visible)) + + +class TestNoAccessorShadowing(CustomTestCase): + def test_no_local_shadows_a_context_accessor(self): + offenders = [] + for path in sorted(_PACKAGE_ROOT.rglob("*.py")): + rel = path.relative_to(_PACKAGE_ROOT).as_posix() + if rel.startswith("srt/runtime_context.py"): + continue + source = path.read_text() + # A file that never names the module cannot import an accessor + # from it, at module scope or inside any function. + if _CONTEXT_MODULE not in source: + continue + try: + tree = ast.parse(source) + except SyntaxError: + continue + module_accessors = _module_level_accessor_imports(tree) + for func, name, lineno in _shadowing_assignments(tree, module_accessors): + offenders.append(f"{rel}:{lineno}: {func}() binds {name!r}") + self.assertFalse( + offenders, + "locals shadow a runtime_context accessor imported in the same " + "module; the name is local for the whole function, so any call to " + "the accessor there raises UnboundLocalError:\n" + "\n".join(offenders), + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/test_global_config_read_ratchet.py b/test/registered/unit/test_global_config_read_ratchet.py index a6b6a11e8..c79e8f1a4 100644 --- a/test/registered/unit/test_global_config_read_ratchet.py +++ b/test/registered/unit/test_global_config_read_ratchet.py @@ -6,22 +6,24 @@ startup record. Config decisions read the namespace accessors instead including post-publish overrides, and per-runner values come from the runner that owns them. -Two shapes count as a read: the direct ``get_server_args().field``, and the -alias ``sa = get_server_args()`` followed by ``sa.field`` in the same function. -A whole-object pass (``def f(server_args)``) is not a global read and is not -counted — there the caller decided which instance to hand over. +Business code no longer reads the published record for a config value at all: +the baselines are zero for both shapes, over the whole package minus the two +modules that own the slot. -What legitimately remains: +Where the remaining reads live (``runtime_context.py``, exempt by module): -- **Derived APIs.** ``@property`` and method members of ``ServerArgs`` - (``mamba_cache_chunk_size``, ``get_model_config()``, - ``enable_mamba_extra_buffer()``, …) are computed from several fields plus the - HF config, so they are not namespace leaves and ``ServerArgs`` is their only - home. Exempt by name below. +- **Derived members.** ``@property`` / method members of ``ServerArgs`` + (``mamba_cache_chunk_size``, ``max_speculative_num_draft_tokens``, + ``use_mla_backend()``, ``get_attention_backends()``, ``get_model_config()``, + ``cutedsl_moe_max_num_tokens()``) are computed from several fields plus the HF + config, so they are not namespace leaves and ``ServerArgs`` is their only + home. ``runtime_context`` exposes each one as a named accessor + (``mamba_cache_chunk_size()`` …) and is the only module that reads the slot + for them. - **Config-intent reads of live-shadowed sizes.** ``get_parallel()`` shadows - ``tp/pp/dcp/attn_cp/moe_dp_size`` with the live topology, so a config-intent - read of one has nowhere else to go. Each exempt site needs an answer the live - property cannot give: + ``tp/pp/dcp/attn_cp/moe_dp_size`` with the live topology, and a few call sites + need what was *configured*: the ``configured_*_size()`` accessors. Their + reasons, per call site: - ``dsa_indexer.pp_size`` gates ``pp_size > 1 and not get_pp_group()...``, and the short circuit is the point: with PP off the group is never touched, which @@ -35,8 +37,23 @@ What legitimately remains: sizes are equal there and a live comparison is always false. - ``model_loader/loader.py`` reports both: the same dict carries the live ``moe_dp_size`` under ``"dp"``, so this entry is the configured intent. -- The alias-form baseline is not zero yet. Lowering it is the next slice; the - failure message lists the sites whenever the count moves. + +What the ratchet sees, syntactically: ``get_server_args().field``, +``sa = get_server_args()`` followed by ``sa.field`` (function-local, module-level, +or parked on an instance attribute -- ``self._sa = get_server_args()`` read from +another method of the same class), function-local copies of an alias to a +fixpoint (``cfg = sa`` then ``cfg.field``), and the dynamic form of each -- +``getattr(, "field")`` -- since a string-named read reaches the same +slot. What it cannot see is a name computed at runtime (``getattr(sa, name)``) +or indirection deeper than a local name copy (through a container, an +attribute of another object, a cross-scope copy); the census tool in the +context repo is what audits those. + +A whole-object pass (``def f(server_args)``) is not a global read and is not +counted -- there the caller decided which instance to hand over. An optional +parameter that falls back to the global (``f(server_args=None)``) hides one, +so those fallbacks were removed; the ratchet cannot see them and the census +tool in the context repo is what audits that shape. """ from sglang.test.ci.ci_register import register_cpu_ci @@ -54,48 +71,92 @@ from sglang.test.test_utils import CustomTestCase # scanned so a new one cannot appear there unnoticed. _PACKAGE_ROOT = Path(next(iter(sglang.__path__))) -_DERIVED_MEMBERS = frozenset( - { - "cutedsl_moe_max_num_tokens", - "enable_mamba_extra_buffer", - "enable_mamba_extra_buffer_lazy", - "get_attention_backends", - "get_model_config", - "mamba_cache_chunk_size", - "max_speculative_num_draft_tokens", - "model_config", - "use_mla_backend", - } -) +# The modules that own the slot: runtime_context publishes it and exposes the +# named accessors for the derived members, server_args/arg_groups ARE the +# resolution pipeline. +_SLOT_OWNERS = ("srt/runtime_context.py", "srt/server_args.py", "srt/arg_groups/") -_CONFIG_INTENT_SIZES = frozenset( - { - ("srt/layers/attention/dsa/dsa_indexer.py", "pp_size"), - ("srt/layers/dp_attention.py", "attn_cp_size"), - ("srt/layers/dp_attention.py", "moe_dp_size"), - ("srt/model_loader/loader.py", "moe_dp_size"), - ("srt/utils/cuda_ipc_transport_utils.py", "tp_size"), - } -) +# Every call site of a ``configured_*_size()`` accessor, with the reason the +# live topology cannot answer there. The test below asserts this map is exactly +# the set of call sites, so the reasons cannot drift away from the code. +_CONFIGURED_SIZE_CALL_SITES = { + ("srt/layers/attention/dsa/dsa_indexer.py", "configured_pp_size"): ( + "gates `pp_size > 1 and not get_pp_group()...`; the short circuit is the " + "point, since with PP off the group is never touched, which is what lets " + "the Indexer be constructed before distributed init" + ), + ("srt/layers/dp_attention.py", "configured_attn_cp_size"): ( + "compared against the configured moe_dp_size below" + ), + ("srt/layers/dp_attention.py", "configured_moe_dp_size"): ( + "the configuration this predicate detects (attn_cp_size > moe_dp_size) is " + "the one where initialize_model_parallel aliases _MOE_DP to _ATTN_CP, so " + "the live sizes are equal there and a live comparison is always false" + ), + ("srt/model_loader/loader.py", "configured_moe_dp_size"): ( + "the same dict already carries the live moe_dp_size under 'dp'; this entry " + "is the configured intent" + ), + ("srt/utils/cuda_ipc_transport_utils.py", "configured_tp_size"): ( + "runs in the tokenizer process, which has no parallel groups at all" + ), + ("srt/models/kimi_k25.py", "configured_tp_size"): ( + "the IPC refcount has to name the same number the recycler waits on, and " + "that waiter (MmItemMemoryPool.try_to_recycle) reads configured_tp_size() " + "because it runs in the tokenizer process; a refcount taken from the live " + "attention subgroup would strand items in the bounded pool" + ), + ("srt/models/kimi_k3.py", "configured_tp_size"): ( + "same as kimi_k25: the IPC refcount must agree with the recycler's waiter" + ), +} + +# A dynamic read whose name is set nowhere in the tree, so the predicate it +# feeds is inert (the ``getattr`` default decides it). Converting it would mean +# choosing what it should have named, which is the CP path's call, not this +# sweep's -- so it is listed here rather than silently counted or "fixed". +_INERT_DYNAMIC_READS = frozenset({("srt/layers/cp/base.py", "_is_dsa_model_arch")}) _DIRECT_BASELINE = 0 _ALIAS_BASELINE = 0 def _is_global_call(node) -> bool: - return ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Name) - and node.func.id == "get_server_args" - ) + """``get_server_args()`` however it is spelled: bare, or module-qualified + (``ctx.get_server_args()``), which an ast.Name check alone would miss.""" + if not isinstance(node, ast.Call): + return False + func = node.func + if isinstance(func, ast.Name): + return func.id == "get_server_args" + return isinstance(func, ast.Attribute) and func.attr == "get_server_args" -def _collect(rel: str, tree: ast.AST): - """The (direct, alias) field reads in one module.""" +def _collect(rel: str, tree: ast.AST, inert: frozenset = frozenset()): + """The (direct, alias) field reads in one module. + + ``inert`` names the fields listed in ``_INERT_DYNAMIC_READS`` for this file; + they are dropped here, at the point the read is recognized, so the filter + matches on the field name rather than on the rendered message. + """ direct, alias = [], [] def counted(attr: str) -> bool: - return attr not in _DERIVED_MEMBERS and (rel, attr) not in _CONFIG_INTENT_SIZES + return attr not in inert + + def _getattr_name(node): + """``getattr(, "field")`` names a field just as ``.field`` does; + matching only ast.Attribute would let a dynamic read walk past.""" + if not ( + isinstance(node, ast.Call) + and isinstance(node.func, ast.Name) + and node.func.id == "getattr" + and len(node.args) >= 2 + and isinstance(node.args[1], ast.Constant) + and isinstance(node.args[1].value, str) + ): + return None + return node.args[1].value for node in ast.walk(tree): if ( @@ -105,17 +166,54 @@ def _collect(rel: str, tree: ast.AST): ): direct.append(f"{rel}:{node.lineno}: get_server_args().{node.attr}") + name = _getattr_name(node) + if name is not None and _is_global_call(node.args[0]) and counted(name): + direct.append(f"{rel}:{node.lineno}: getattr(get_server_args(), {name!r})") + if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): continue params = {a.arg for a in list(node.args.args) + list(node.args.kwonlyargs)} bound = {} for inner in ast.walk(node): - if isinstance(inner, ast.Assign) and _is_global_call(inner.value): - for target in inner.targets: - if isinstance(target, ast.Name) and target.id not in params: - bound.setdefault(target.id, inner.lineno) + # ``sa = get_server_args()`` and its annotated form + # ``sa: ServerArgs = get_server_args()``. + if isinstance(inner, (ast.Assign, ast.AnnAssign)) and _is_global_call( + getattr(inner, "value", None) + ): + targets = ( + inner.targets if isinstance(inner, ast.Assign) else [inner.target] + ) + for target in targets: + if not isinstance(target, ast.Name): + continue + # A parameter reassigned from the global is the + # optional-injection shape (``f(server_args=None)`` then + # ``server_args = get_server_args()``): the reads that + # follow are global reads wearing a parameter's name, so + # they count from the bind on. + bound.setdefault(target.id, inner.lineno) if not bound: continue + # A copy of an alias reaches the same record (``cfg = sa`` after + # ``sa = get_server_args()``), so follow Name-to-Name assignments to a + # fixpoint. Deeper indirection (through containers, attributes of + # other objects, cross-scope copies) stays census-tool territory. + changed = True + while changed: + changed = False + for inner in ast.walk(node): + if not isinstance(inner, (ast.Assign, ast.AnnAssign)): + continue + value = getattr(inner, "value", None) + if not (isinstance(value, ast.Name) and value.id in bound): + continue + targets = ( + inner.targets if isinstance(inner, ast.Assign) else [inner.target] + ) + for target in targets: + if isinstance(target, ast.Name) and target.id not in bound: + bound[target.id] = inner.lineno + changed = True for inner in ast.walk(node): if ( isinstance(inner, ast.Attribute) @@ -128,6 +226,170 @@ def _collect(rel: str, tree: ast.AST): f"{rel}:{inner.lineno}: {inner.value.id}.{inner.attr} " f"(bound from get_server_args() at line {bound[inner.value.id]})" ) + name = _getattr_name(inner) + if ( + name is not None + and isinstance(inner.args[0], ast.Name) + and inner.args[0].id in bound + and inner.lineno >= bound[inner.args[0].id] + and counted(name) + ): + alias.append( + f"{rel}:{inner.lineno}: getattr({inner.args[0].id}, {name!r}) " + f"(bound from get_server_args() at line {bound[inner.args[0].id]})" + ) + # A module-level alias is visible to every function in the file, so it needs + # its own pass -- the per-function scan above deliberately does not reach + # across scopes. + module_bound = {} + module_stack = list(tree.body) + while module_stack: + stmt = module_stack.pop() + # A module-level bind can sit inside an `if` / `try` / `with`, so the + # walk descends into those bodies -- but not into a nested function or + # class, whose binds are that scope's own. + if isinstance( + stmt, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef, ast.Lambda) + ): + continue + module_stack.extend(ast.iter_child_nodes(stmt)) + if isinstance(stmt, (ast.Assign, ast.AnnAssign)) and _is_global_call( + getattr(stmt, "value", None) + ): + targets = stmt.targets if isinstance(stmt, ast.Assign) else [stmt.target] + for target in targets: + if isinstance(target, ast.Name): + module_bound.setdefault(target.id, stmt.lineno) + if module_bound: + # Shadowing is per lexical scope: a function with its own `sa` hides the + # module alias *inside that function only*. Aggregating the names + # file-wide would suppress every read in the module, including the + # top-level ones and the ones in functions that do resolve to the alias. + parents = {} + scope_binds = {} + stack = [tree] + while stack: + node = stack.pop() + enclosing = parents.get(id(node)) + for child in ast.iter_child_nodes(node): + parents[id(child)] = ( + node + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) + else enclosing + ) + stack.append(child) + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): + names = { + a.arg for a in list(node.args.args) + list(node.args.kwonlyargs) + } + # Only this scope's own stores: a nested function's local `sa` + # shadows the alias inside *that* function, not in its parent. + pending = list(node.body) + while pending: + inner = pending.pop() + if isinstance( + inner, + ( + ast.FunctionDef, + ast.AsyncFunctionDef, + ast.Lambda, + ast.ClassDef, + ), + ): + continue + if isinstance(inner, ast.Name) and isinstance(inner.ctx, ast.Store): + names.add(inner.id) + pending.extend(ast.iter_child_nodes(inner)) + scope_binds[id(node)] = names + + def _shadowed(node, name): + scope = parents.get(id(node)) + while scope is not None: + if name in scope_binds.get(id(scope), ()): + return True + scope = parents.get(id(scope)) + return False + + for node in ast.walk(tree): + base = attr = None + if ( + isinstance(node, ast.Attribute) + and isinstance(node.value, ast.Name) + and node.value.id in module_bound + ): + base, attr = node.value.id, node.attr + shown = f"{base}.{attr}" + else: + attr_name = _getattr_name(node) + if ( + attr_name is not None + and isinstance(node.args[0], ast.Name) + and node.args[0].id in module_bound + ): + base, attr = node.args[0].id, attr_name + shown = f"getattr({base}, {attr!r})" + if base and not _shadowed(node, base) and counted(attr): + alias.append( + f"{rel}:{node.lineno}: {shown} " + f"(module-level bind from get_server_args() at line " + f"{module_bound[base]})" + ) + + # An alias parked on an instance attribute (``self._sa = get_server_args()`` + # in one method, ``self._sa.field`` in another) reaches the same slot and + # crosses function scopes, so it is collected per class rather than per + # function. + for node in ast.walk(tree): + if not isinstance(node, ast.ClassDef): + continue + attr_bound = {} + for inner in ast.walk(node): + if isinstance(inner, (ast.Assign, ast.AnnAssign)) and _is_global_call( + getattr(inner, "value", None) + ): + targets = ( + inner.targets if isinstance(inner, ast.Assign) else [inner.target] + ) + for target in targets: + if ( + isinstance(target, ast.Attribute) + and isinstance(target.value, ast.Name) + and target.value.id in ("self", "cls") + ): + attr_bound.setdefault( + (target.value.id, target.attr), inner.lineno + ) + if not attr_bound: + continue + + def _bound_attr(value): + """``self._sa`` when that attribute was bound from the global.""" + if ( + isinstance(value, ast.Attribute) + and isinstance(value.value, ast.Name) + and (value.value.id, value.attr) in attr_bound + ): + return (value.value.id, value.attr) + return None + + for inner in ast.walk(node): + key = shown = None + if isinstance(inner, ast.Attribute): + key = _bound_attr(inner.value) + if key is not None and counted(inner.attr): + shown = f"{key[0]}.{key[1]}.{inner.attr}" + else: + name = _getattr_name(inner) + if name is not None: + key = _bound_attr(inner.args[0]) + if key is not None and counted(name): + shown = f"getattr({key[0]}.{key[1]}, {name!r})" + if shown is not None: + alias.append( + f"{rel}:{inner.lineno}: {shown} " + f"(attribute bind from get_server_args() at line " + f"{attr_bound[key]})" + ) return direct, alias @@ -135,11 +397,14 @@ def _field_reads(): direct, alias = [], [] for path in sorted(_PACKAGE_ROOT.rglob("*.py")): rel = path.relative_to(_PACKAGE_ROOT).as_posix() + if rel.startswith(_SLOT_OWNERS): + continue try: tree = ast.parse(path.read_text()) except SyntaxError: continue - module_direct, module_alias = _collect(rel, tree) + inert = frozenset(name for path_, name in _INERT_DYNAMIC_READS if path_ == rel) + module_direct, module_alias = _collect(rel, tree, inert) direct += module_direct alias += module_alias return direct, alias @@ -167,5 +432,90 @@ class TestGlobalConfigReadRatchet(CustomTestCase): self._check("alias-form", alias, _ALIAS_BASELINE) +class TestConfiguredSizeCallSites(CustomTestCase): + """The configured-vs-live exceptions are enumerated, with reasons. + + ``configured_*_size()`` answers what the user asked for where + ``get_parallel()`` would answer what the process ended up with. Each such + exception is listed above with why the live property cannot serve it, and + this case fails if the code and that list disagree. + + The unit is **(file, accessor)**, not the individual call: a second + `configured_pp_size()` in a file already registered for it collapses into + the same entry, so the reason has to cover the file's use of that accessor + rather than one line. A new file, or a new accessor in a listed file, is + what this catches -- in either call form (bare or module-qualified). + """ + + def test_the_call_sites_match_the_documented_set(self): + found = set() + for path in sorted(_PACKAGE_ROOT.rglob("*.py")): + rel = path.relative_to(_PACKAGE_ROOT).as_posix() + if rel.startswith(_SLOT_OWNERS): + continue + try: + tree = ast.parse(path.read_text()) + except SyntaxError: + continue + for node in ast.walk(tree): + if not isinstance(node, ast.Call): + continue + func = node.func + name = ( + func.id + if isinstance(func, ast.Name) + else (func.attr if isinstance(func, ast.Attribute) else None) + ) + if name and name.startswith("configured_") and name.endswith("_size"): + found.add((rel, name)) + documented = set(_CONFIGURED_SIZE_CALL_SITES) + self.assertEqual( + documented, + found, + "configured-size call sites drifted from their documented reasons.\n" + f" undocumented: {sorted(found - documented)}\n" + f" stale entries: {sorted(documented - found)}", + ) + + +class TestNoRenamedAccessorImports(CustomTestCase): + """The scanners above match ``get_server_args`` and ``configured_*_size`` + by their literal names, so an ``import ... as`` rename would walk a read + straight past both the zero baseline and the call-site registry. Renaming + these accessors buys nothing (the names are already short and unambiguous), + so it is banned outright — which is exactly what makes literal-name + matching sound.""" + + def test_the_scanned_accessors_are_never_import_renamed(self): + offenders = [] + for path in sorted(_PACKAGE_ROOT.rglob("*.py")): + rel = path.relative_to(_PACKAGE_ROOT).as_posix() + try: + tree = ast.parse(path.read_text()) + except SyntaxError: + continue + for node in ast.walk(tree): + if not isinstance(node, (ast.ImportFrom, ast.Import)): + continue + for imported in node.names: + if imported.asname is None or imported.asname == imported.name: + continue + base = imported.name.rsplit(".", 1)[-1] + if base == "get_server_args" or ( + base.startswith("configured_") and base.endswith("_size") + ): + offenders.append( + f"{rel}:{node.lineno}: {imported.name} as " + f"{imported.asname}" + ) + self.assertFalse( + offenders, + "get_server_args / configured_*_size imported under another name; " + "the read ratchet and the configured-size registry match these " + "accessors by their literal names, so a rename silently escapes " + "both:\n" + "\n".join(offenders), + ) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index 9890c504e..3d519b5cc 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -5,7 +5,10 @@ from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=5, suite="base-a-test-cpu") import dataclasses +import json import os +import shutil +import tempfile import unittest from unittest.mock import patch @@ -20,8 +23,10 @@ from sglang.srt.runtime_context import ( get_flags, get_parallel, get_server_args, + max_speculative_num_draft_tokens, reset_context, ) +from sglang.srt.server_args import ServerArgs from sglang.test.test_utils import CustomTestCase _PS = "sglang.srt.distributed.parallel_state" @@ -378,6 +383,14 @@ class _FakeResolvedArgs: sampling_backend: A[ str | None, Arg(help="s", resolvable=True), NS("exec.kernel") ] = None + attention_backend: A[str | None, Arg(help="ab"), NS("exec.kernel")] = None + prefill_attention_backend: A[str | None, Arg(help="pab"), NS("exec.kernel")] = None + decode_attention_backend: A[str | None, Arg(help="dab"), NS("exec.kernel")] = None + disable_radix_cache: A[bool, Arg(help="drc"), NS("memory")] = False + mamba_radix_cache_strategy: A[str, Arg(help="mrcs"), NS("exec.mamba")] = "auto" + speculative_num_draft_tokens: A[int | None, Arg(help="d"), NS("spec")] = None + speculative_adaptive: A[bool, Arg(help="a"), NS("spec")] = False + speculative_adaptive_config: A[str | None, Arg(help="c"), NS("spec")] = None _resolved_overrides: list = dataclasses.field(default_factory=list) @@ -965,5 +978,180 @@ class TestPublishLifecycle(_IsolatedServerArgs): self.assertFalse(get_flags().capture.enable_torch_compile) +class TestDerivedPredicatesAgreeAcrossTiers(_IsolatedServerArgs): + """One definition per predicate, checked rather than asserted in prose. + + Each of these exists twice by construction -- once over a config-shaped + object (the resolution pipeline's `*_of` helper, which `ServerArgs` + delegates to) and once over the published bags. The pair must agree on + every input, or a decision made before publish differs from the same + decision made after it. + """ + + _STRATEGIES = ("auto", "no_buffer", "extra_buffer", "extra_buffer_lazy") + + def test_mamba_extra_buffer_matches_the_member(self): + from sglang.srt.runtime_context import ( + mamba_extra_buffer_enabled, + mamba_extra_buffer_lazy_enabled, + ) + + for disable_radix_cache in (False, True): + for strategy in self._STRATEGIES: + with self.subTest(radix=disable_radix_cache, strategy=strategy): + args = _FakeResolvedArgs( + disable_radix_cache=disable_radix_cache, + mamba_radix_cache_strategy=strategy, + ) + get_context().set_server_args(args) + self.assertEqual( + ServerArgs.enable_mamba_extra_buffer(args), + mamba_extra_buffer_enabled(), + ) + self.assertEqual( + ServerArgs.enable_mamba_extra_buffer_lazy(args), + mamba_extra_buffer_lazy_enabled(), + ) + + def test_attention_backends_match_the_member(self): + from sglang.srt.runtime_context import attention_backends + + backends = (None, "fa3", "triton") + for base in backends: + for prefill in backends: + for decode in backends: + with self.subTest(base=base, prefill=prefill, decode=decode): + args = _FakeResolvedArgs( + attention_backend=base, + prefill_attention_backend=prefill, + decode_attention_backend=decode, + ) + get_context().set_server_args(args) + self.assertEqual( + ServerArgs.get_attention_backends(args), + attention_backends(), + ) + + +class TestAdaptiveDraftBoundLifecycle(_IsolatedServerArgs): + """The adaptive draft-token bound is memoized on the config path, so the + memo has to end with the publication it was computed under. + + Without that, a process that republishes with the same adaptive-config path + -- the file having been rewritten in between -- keeps the previous bound and + under-allocates the draft-token buffers sized from it. + """ + + def _write_config(self, steps): + path = os.path.join(tempfile.mkdtemp(prefix="adaptive_cfg_"), "adaptive.json") + self.addCleanup(shutil.rmtree, os.path.dirname(path), ignore_errors=True) + with open(path, "w") as handle: + json.dump({"1": {"candidate_steps": steps}}, handle) + return path + + def test_republishing_recomputes_the_bound(self): + path = self._write_config([2]) + get_context().set_server_args( + _FakeResolvedArgs( + speculative_num_draft_tokens=3, + speculative_adaptive=True, + speculative_adaptive_config=path, + ) + ) + self.assertEqual(max_speculative_num_draft_tokens(), 3) + + with open(path, "w") as handle: + json.dump({"1": {"candidate_steps": [4]}}, handle) + # Same path, new contents: the memo must not survive the republish. + get_context().set_server_args( + _FakeResolvedArgs( + speculative_num_draft_tokens=3, + speculative_adaptive=True, + speculative_adaptive_config=path, + ) + ) + self.assertEqual(max_speculative_num_draft_tokens(), 5) + + def test_reset_clears_the_bound(self): + path = self._write_config([2]) + get_context().set_server_args( + _FakeResolvedArgs( + speculative_num_draft_tokens=3, + speculative_adaptive=True, + speculative_adaptive_config=path, + ) + ) + self.assertEqual(max_speculative_num_draft_tokens(), 3) + reset_context() + with open(path, "w") as handle: + json.dump({"1": {"candidate_steps": [6]}}, handle) + get_context().set_server_args( + _FakeResolvedArgs( + speculative_num_draft_tokens=3, + speculative_adaptive=True, + speculative_adaptive_config=path, + ) + ) + self.assertEqual(max_speculative_num_draft_tokens(), 7) + + +class TestNamedAccessorsCallWhatTheyWrap(CustomTestCase): + """A named accessor must *call* a member that is a method. + + `return get_server_args().x` hands back a bound method when `x` is defined + with `def`; the failure then lands far away, in whatever arithmetic the + caller does with it. Checked statically so accessors that need a real model + config are covered too. + """ + + def test_accessors_that_wrap_methods_call_them(self): + import ast + import functools + import inspect + + import sglang.srt.runtime_context as rc + from sglang.srt.server_args import ServerArgs + + tree = ast.parse(inspect.getsource(rc)) + wrong = [] + for node in tree.body: + if not isinstance(node, ast.FunctionDef): + continue + for inner in ast.walk(node): + if not (isinstance(inner, ast.Return) and inner.value is not None): + continue + value = inner.value + called = isinstance(value, ast.Call) + target = value.func if called else value + if not ( + isinstance(target, ast.Attribute) + and isinstance(target.value, ast.Call) + and isinstance(target.value.func, ast.Name) + and target.value.func.id == "get_server_args" + ): + continue + member = getattr(ServerArgs, target.attr, None) + # A `property` / `functools.cached_property` member is already + # evaluated by the attribute access, so it is named here to keep + # the failure message from calling it "not a method" -- the fix + # for those is the opposite one. + kind = ( + "a property" + if isinstance(member, (property, functools.cached_property)) + else "not a method" + ) + if inspect.isfunction(member) and not called: + wrong.append( + f"{node.name}(): returns ServerArgs.{target.attr} without " + "calling it, so callers get a bound method" + ) + if not inspect.isfunction(member) and called: + wrong.append( + f"{node.name}(): calls ServerArgs.{target.attr}, which is " + f"{kind} -- the attribute access already produced the value" + ) + self.assertEqual([], wrong, "\n".join(wrong)) + + if __name__ == "__main__": unittest.main()