config: the record is not an object that gets passed around (#36622)

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-08-27 12:57:10 -07:00
committed by GitHub
co-authored by Claude Opus 5
parent fd40a331bf
commit 6ff2a20ccf
23 changed files with 234 additions and 125 deletions
+6 -10
View File
@@ -265,12 +265,8 @@ def run_post_process_pass(server_args: Any, fn: Callable[..., dict]) -> None:
def _apply_fields(server_args: Any, fields: Dict[str, Any]) -> None: def _apply_fields(server_args: Any, fields: Dict[str, Any]) -> None:
"""Write record fields past the guard that forbids post-resolution writes. """Write fields on behalf of the pipeline (bypasses the strict bare-
assignment guard that protects post-resolution mutation)."""
Resolution declares, so nothing in the pipeline calls this. It exists for
``RuntimeContext.override_server_args``, the launch stand-in tests use: there
the caller's values are both the operator's input and resolution's answer.
"""
object.__setattr__(server_args, "_internal_write", True) object.__setattr__(server_args, "_internal_write", True)
try: try:
for field, value in fields.items(): for field, value in fields.items():
@@ -299,7 +295,7 @@ def declare_resolution(server_args: Any, source: str, **fields: Any) -> None:
stash = getattr(server_args, "_resolved_overrides", None) stash = getattr(server_args, "_resolved_overrides", None)
if stash is None: if stash is None:
stash = [] stash = []
object.__setattr__(server_args, "_resolved_overrides", stash) server_args._resolved_overrides = stash
stash.append((source, dict(fields))) stash.append((source, dict(fields)))
@@ -333,12 +329,12 @@ def declare_late_resolution(server_args: Any, source: str, **fields: Any) -> Non
log = getattr(server_args, "_runtime_mutations", None) log = getattr(server_args, "_runtime_mutations", None)
if log is None: if log is None:
log = [] log = []
object.__setattr__(server_args, "_runtime_mutations", log) server_args._runtime_mutations = log
log.append((source, dict(fields))) log.append((source, dict(fields)))
stash = getattr(server_args, "_resolved_overrides", None) stash = getattr(server_args, "_resolved_overrides", None)
if stash is None: if stash is None:
stash = [] stash = []
object.__setattr__(server_args, "_resolved_overrides", stash) server_args._resolved_overrides = stash
stash.append((source, dict(fields))) stash.append((source, dict(fields)))
@@ -376,7 +372,7 @@ def declare_direct_writes(
stash = getattr(server_args, "_resolved_overrides", None) stash = getattr(server_args, "_resolved_overrides", None)
if stash is None: if stash is None:
stash = [] stash = []
object.__setattr__(server_args, "_resolved_overrides", stash) server_args._resolved_overrides = stash
# A resolver reached this way can also declare properly -- the in-tree # A resolver reached this way can also declare properly -- the in-tree
# implementations of these hooks do. Those fields are already explained, and # implementations of these hooks do. Those fields are already explained, and
# recording them again would attribute them to the wrapper and bury an # recording them again would attribute them to the wrapper and bury an
@@ -26,7 +26,6 @@ from sglang.srt.runtime_context import (
get_schedule, get_schedule,
get_serving, get_serving,
) )
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.common import ceil_align from sglang.srt.utils.common import ceil_align
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -44,7 +43,6 @@ class DecodeKVCacheOffloadManager:
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator, token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator,
tp_group: torch.distributed.ProcessGroup, tp_group: torch.distributed.ProcessGroup,
tree_cache: BasePrefixCache, tree_cache: BasePrefixCache,
server_args: ServerArgs,
) -> None: ) -> None:
self.req_to_token_pool = req_to_token_pool self.req_to_token_pool = req_to_token_pool
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
@@ -64,7 +62,6 @@ class DecodeKVCacheOffloadManager:
self.decode_host_mem_pool = build_kv_host_pool( self.decode_host_mem_pool = build_kv_host_pool(
kv_pool=kv_cache, kv_pool=kv_cache,
page_size=self.page_size, page_size=self.page_size,
server_args=server_args,
use_mla=isinstance(kv_cache, MLATokenToKVPool), use_mla=isinstance(kv_cache, MLATokenToKVPool),
) )
+1 -1
View File
@@ -10,11 +10,11 @@ import torch
from sglang.srt.distributed import get_world_group, parallel_state from sglang.srt.distributed import get_world_group, parallel_state
from sglang.srt.distributed.utils import get_global_tcp_store from sglang.srt.distributed.utils import get_global_tcp_store
from sglang.srt.eplb.expert_location import broadcast_global_expert_location_metadata from sglang.srt.eplb.expert_location import broadcast_global_expert_location_metadata
from sglang.srt.managers.schedule_batch import ServerArgs
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_exec, get_exec,
get_parallel, get_parallel,
) )
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import is_cpu, is_cuda from sglang.srt.utils import is_cpu, is_cuda
if TYPE_CHECKING: if TYPE_CHECKING:
+4 -4
View File
@@ -228,13 +228,13 @@ async def init_multi_tokenizer() -> ServerArgs:
server_args: ServerArgs server_args: ServerArgs
port_args: PortArgs port_args: PortArgs
publish(server_args, role="tokenizer")
# API key authentication is not supported in multi-tokenizer mode # API key authentication is not supported in multi-tokenizer mode
assert ( assert (
server_args.api_key is None get_serving().api_key is None
), "API key is not supported in multi-tokenizer mode" ), "API key is not supported in multi-tokenizer mode"
publish(server_args, role="tokenizer")
# Create a new ipc name for the current process # Create a new ipc name for the current process
port_args.tokenizer_ipc_name = ( port_args.tokenizer_ipc_name = (
f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}" f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}"
@@ -2754,7 +2754,7 @@ def _start_native_grpc_server_for_runtime(
host=get_serving().host, host=get_serving().host,
port=grpc_port, port=grpc_port,
runtime_handle=runtime_handle, runtime_handle=runtime_handle,
worker_threads=server_args.grpc_worker_threads, worker_threads=get_serving().grpc_worker_threads,
) )
logger.info(f"Native gRPC server started on {get_serving().host}:{grpc_port}") logger.info(f"Native gRPC server started on {get_serving().host}:{grpc_port}")
return grpc_handle return grpc_handle
+1 -1
View File
@@ -187,7 +187,7 @@ class NativeMmHost:
import_processors("sglang.srt.multimodal.processors") import_processors("sglang.srt.multimodal.processors")
if mm_process_pkg := envs.SGLANG_EXTERNAL_MM_PROCESSOR_PACKAGE.get(): if mm_process_pkg := envs.SGLANG_EXTERNAL_MM_PROCESSOR_PACKAGE.get():
import_processors(mm_process_pkg, overwrite=True) import_processors(mm_process_pkg, overwrite=True)
self._processor = processor or get_processor_wrapper(self.server_args) self._processor = processor or get_processor_wrapper()
def resolve_native_spec(self) -> Optional[NativeMmSpec]: def resolve_native_spec(self) -> Optional[NativeMmSpec]:
"""The :class:`NativeMmSpec` for this model, or ``None`` when it has no """The :class:`NativeMmSpec` for this model, or ``None`` when it has no
+4 -14
View File
@@ -140,7 +140,6 @@ from sglang.srt.observability.req_time_stats import (
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import flatten_nested_list from sglang.srt.utils import flatten_nested_list
from sglang.srt.utils.token_sequence_matcher import TokenSequenceMatcher from sglang.srt.utils.token_sequence_matcher import TokenSequenceMatcher
@@ -1960,7 +1959,6 @@ def release_req(
*, *,
req: Req, req: Req,
remaing_req_count: int, remaing_req_count: int,
server_args: ServerArgs,
req_to_token_pool: ReqToTokenPool, req_to_token_pool: ReqToTokenPool,
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator, token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator,
tree_cache: BasePrefixCache, tree_cache: BasePrefixCache,
@@ -1997,7 +1995,6 @@ def release_req(
def retract_all( def retract_all(
*, *,
reqs: List[Req], reqs: List[Req],
server_args: ServerArgs,
req_to_token_pool: ReqToTokenPool, req_to_token_pool: ReqToTokenPool,
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator, token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator,
tree_cache: BasePrefixCache, tree_cache: BasePrefixCache,
@@ -2008,7 +2005,6 @@ def retract_all(
release_req( release_req(
req=reqs[idx], req=reqs[idx],
remaing_req_count=len(reqs) - idx, remaing_req_count=len(reqs) - idx,
server_args=server_args,
req_to_token_pool=req_to_token_pool, req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=token_to_kv_pool_allocator, token_to_kv_pool_allocator=token_to_kv_pool_allocator,
tree_cache=tree_cache, tree_cache=tree_cache,
@@ -2877,9 +2873,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
evict_from_tree_cache(self.tree_cache, num_tokens) evict_from_tree_cache(self.tree_cache, num_tokens)
return self.token_to_kv_pool_allocator.available_size() >= num_tokens return self.token_to_kv_pool_allocator.available_size() >= num_tokens
def retract_decode( def retract_decode(self) -> Tuple[List[Req], float, List[Req]]:
self, server_args: ServerArgs
) -> Tuple[List[Req], float, List[Req]]:
"""Retract the decoding requests when there is not enough memory.""" """Retract the decoding requests when there is not enough memory."""
sorted_indices = self._get_decode_retraction_order(self.reqs) sorted_indices = self._get_decode_retraction_order(self.reqs)
sorted_indices = beam_retraction_order(sorted_indices, self.reqs) sorted_indices = beam_retraction_order(sorted_indices, self.reqs)
@@ -2913,12 +2907,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.token_to_kv_pool_allocator, self.token_to_kv_pool_allocator,
) )
# Aborting, so a host backup to resume from would be wasted. # Aborting, so a host backup to resume from would be wasted.
self.release_req( self.release_req(idx, len(sorted_indices), offload_kv=False)
idx, len(sorted_indices), server_args, offload_kv=False
)
continue continue
# release memory and don't insert into the tree because we need the space instantly # release memory and don't insert into the tree because we need the space instantly
if self.release_req(idx, len(sorted_indices), server_args): if self.release_req(idx, len(sorted_indices)):
retracted_reqs.append(req) retracted_reqs.append(req)
else: else:
# The retraction host pool could not hold the backup and the # The retraction host pool could not hold the backup and the
@@ -2953,7 +2945,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.req_to_token_pool, self.req_to_token_pool,
self.token_to_kv_pool_allocator, self.token_to_kv_pool_allocator,
) )
self.release_req(last_idx, 0, server_args, offload_kv=False) self.release_req(last_idx, 0, offload_kv=False)
logger.warning( logger.warning(
"retract_decode: aborted last request %s due to OOM", last_req.rid "retract_decode: aborted last request %s due to OOM", last_req.rid
) )
@@ -3012,13 +3004,11 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self, self,
idx: int, idx: int,
remaing_req_count: int, remaing_req_count: int,
server_args: ServerArgs,
offload_kv: bool = True, offload_kv: bool = True,
) -> bool: ) -> bool:
return release_req( return release_req(
req=self.reqs[idx], req=self.reqs[idx],
remaing_req_count=remaing_req_count, remaing_req_count=remaing_req_count,
server_args=server_args,
req_to_token_pool=self.req_to_token_pool, req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tree_cache=self.tree_cache, tree_cache=self.tree_cache,
@@ -1498,7 +1498,7 @@ class PrefillAdder:
) )
release_counter += 1 release_counter += 1
self.running_batch.release_req( self.running_batch.release_req(
i, len(self.running_batch.reqs) - release_counter, server_args i, len(self.running_batch.reqs) - release_counter
) )
else: else:
keep_indices.append(i) keep_indices.append(i)
+3 -8
View File
@@ -598,7 +598,6 @@ class Scheduler(
else self.tp_cpu_group else self.tp_cpu_group
), ),
tree_cache=self.tree_cache, tree_cache=self.tree_cache,
server_args=self.server_args,
) )
else: else:
self.decode_offload_manager = None self.decode_offload_manager = None
@@ -3680,9 +3679,7 @@ class Scheduler(
if mamba_allocator is not None if mamba_allocator is not None
else None else None
) )
retracted_reqs, new_token_ratio, reqs_to_abort = batch.retract_decode( retracted_reqs, new_token_ratio, reqs_to_abort = batch.retract_decode()
self.server_args
)
new_available_tokens = self.token_to_kv_pool_allocator.available_size() new_available_tokens = self.token_to_kv_pool_allocator.available_size()
new_token_gained = new_available_tokens - old_available_tokens new_token_gained = new_available_tokens - old_available_tokens
mamba_num_gained = ( mamba_num_gained = (
@@ -4595,9 +4592,8 @@ class Scheduler(
if envs.SGLANG_EXPOSE_OWN_ENV_VARS.get(): if envs.SGLANG_EXPOSE_OWN_ENV_VARS.get():
ret["env_vars"] = exportable_env_vars() ret["env_vars"] = exportable_env_vars()
# These fields are not msgpack-serializable (a config object and a bound # A bound signal handler is not msgpack-serializable, and no reader
# signal handler); no reader consumes them. # consumes it.
ret.pop("model_config", None)
ret.pop("custom_sigquit_handler", None) ret.pop("custom_sigquit_handler", None)
return GetInternalStateReqOutput(internal_state=msgspec_to_builtins(ret)) return GetInternalStateReqOutput(internal_state=msgspec_to_builtins(ret))
@@ -4950,7 +4946,6 @@ class Scheduler(
# discarded. Non-decode modes ignore offload_kv (they never offload). # discarded. Non-decode modes ignore offload_kv (they never offload).
retract_all( retract_all(
reqs=retract_reqs, reqs=retract_reqs,
server_args=self.server_args,
req_to_token_pool=self.req_to_token_pool, req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tree_cache=self.tree_cache, tree_cache=self.tree_cache,
@@ -490,7 +490,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
import_processors("sglang.srt.multimodal.processors") import_processors("sglang.srt.multimodal.processors")
if mm_process_pkg := envs.SGLANG_EXTERNAL_MM_PROCESSOR_PACKAGE.get(): if mm_process_pkg := envs.SGLANG_EXTERNAL_MM_PROCESSOR_PACKAGE.get():
import_processors(mm_process_pkg, overwrite=True) import_processors(mm_process_pkg, overwrite=True)
_processor = get_processor_wrapper(server_args) _processor = get_processor_wrapper()
transport_mode = determine_tensor_transport_mode() transport_mode = determine_tensor_transport_mode()
# We want to parallelize the image pre-processing so we create an executor for it # We want to parallelize the image pre-processing so we create an executor for it
@@ -3602,7 +3602,7 @@ async def print_exception_wrapper(func):
sys.exit(1) sys.exit(1)
def get_processor_wrapper(server_args): def get_processor_wrapper():
return get_processor( return get_processor(
get_serving().tokenizer_path, get_serving().tokenizer_path,
tokenizer_mode=get_serving().tokenizer_mode, tokenizer_mode=get_serving().tokenizer_mode,
@@ -97,7 +97,6 @@ def build_kv_host_pool(
*, *,
kv_pool: Any, kv_pool: Any,
page_size: int, page_size: int,
server_args: ServerArgs,
use_mla: bool, use_mla: bool,
override_kv_cache_dim: Optional[int] = None, override_kv_cache_dim: Optional[int] = None,
host_size: Optional[float] = None, host_size: Optional[float] = None,
@@ -179,7 +178,6 @@ def build_pool_entry(
def build_kv_only_group( def build_kv_only_group(
*, *,
page_size: int, page_size: int,
server_args: ServerArgs,
kv_pool: Any, kv_pool: Any,
full_layer_mapping: dict[int, int], full_layer_mapping: dict[int, int],
use_mla: bool, use_mla: bool,
@@ -192,7 +190,6 @@ def build_kv_only_group(
kv_host_pool = build_kv_host_pool( kv_host_pool = build_kv_host_pool(
kv_pool=kv_pool, kv_pool=kv_pool,
page_size=page_size, page_size=page_size,
server_args=server_args,
use_mla=use_mla, use_mla=use_mla,
override_kv_cache_dim=override_kv_cache_dim, override_kv_cache_dim=override_kv_cache_dim,
host_size=host_size, host_size=host_size,
@@ -223,7 +220,6 @@ def build_kv_only_group(
def build_hybrid_swa_group( def build_hybrid_swa_group(
*, *,
page_size: int, page_size: int,
server_args: ServerArgs,
full_kv_pool: Any, full_kv_pool: Any,
swa_kv_pool: Any, swa_kv_pool: Any,
full_layer_mapping: dict[int, int], full_layer_mapping: dict[int, int],
@@ -241,7 +237,6 @@ def build_hybrid_swa_group(
kv_host_pool = build_kv_host_pool( kv_host_pool = build_kv_host_pool(
kv_pool=full_kv_pool, kv_pool=full_kv_pool,
page_size=page_size, page_size=page_size,
server_args=server_args,
use_mla=use_mla, use_mla=use_mla,
host_size=kv_host_size, host_size=kv_host_size,
pool_label="full", pool_label="full",
@@ -249,7 +244,6 @@ def build_hybrid_swa_group(
swa_host_pool = build_kv_host_pool( swa_host_pool = build_kv_host_pool(
kv_pool=swa_kv_pool, kv_pool=swa_kv_pool,
page_size=page_size, page_size=page_size,
server_args=server_args,
use_mla=use_mla, use_mla=use_mla,
host_size=swa_host_size, host_size=swa_host_size,
mtp_draft_device_pools=mtp_swa_device_pools, mtp_draft_device_pools=mtp_swa_device_pools,
@@ -310,7 +304,6 @@ def build_kv_only_stack(
transfer_layer_num = len(full_layer_mapping) transfer_layer_num = len(full_layer_mapping)
host_pool_group = build_kv_only_group( host_pool_group = build_kv_only_group(
page_size=params.page_size, page_size=params.page_size,
server_args=server_args,
kv_pool=kv_pool, kv_pool=kv_pool,
full_layer_mapping=full_layer_mapping, full_layer_mapping=full_layer_mapping,
use_mla=use_mla, use_mla=use_mla,
@@ -371,7 +364,6 @@ def build_hybrid_swa_stack(
host_pool_group = build_hybrid_swa_group( host_pool_group = build_hybrid_swa_group(
page_size=params.page_size, page_size=params.page_size,
server_args=server_args,
full_kv_pool=full_kv_pool, full_kv_pool=full_kv_pool,
swa_kv_pool=swa_kv_pool, swa_kv_pool=swa_kv_pool,
full_layer_mapping=full_layer_mapping, full_layer_mapping=full_layer_mapping,
@@ -723,7 +715,6 @@ def build_hybrid_mamba_stack(
kv_host_pool = build_kv_host_pool( kv_host_pool = build_kv_host_pool(
kv_pool=kv_pool, kv_pool=kv_pool,
page_size=params.page_size, page_size=params.page_size,
server_args=server_args,
use_mla=use_mla, use_mla=use_mla,
host_size=kv_host_size, host_size=kv_host_size,
mtp_draft_device_pools=mtp_draft_device_pools, mtp_draft_device_pools=mtp_draft_device_pools,
@@ -826,7 +817,6 @@ def build_hybrid_mamba_swa_stack(
kv_host_pool = build_kv_host_pool( kv_host_pool = build_kv_host_pool(
kv_pool=full_kv_pool, kv_pool=full_kv_pool,
page_size=page_size, page_size=page_size,
server_args=server_args,
use_mla=False, use_mla=False,
host_size=kv_host_size, host_size=kv_host_size,
pool_label="full", pool_label="full",
@@ -834,7 +824,6 @@ def build_hybrid_mamba_swa_stack(
swa_host_pool = build_kv_host_pool( swa_host_pool = build_kv_host_pool(
kv_pool=swa_kv_pool, kv_pool=swa_kv_pool,
page_size=page_size, page_size=page_size,
server_args=server_args,
use_mla=False, use_mla=False,
host_size=swa_host_size, host_size=swa_host_size,
pool_label="swa", pool_label="swa",
@@ -925,7 +914,6 @@ def build_anchor_sidecar_stack(
kv_host_pool = build_kv_host_pool( kv_host_pool = build_kv_host_pool(
kv_pool=kv_pool, kv_pool=kv_pool,
page_size=params.page_size, page_size=params.page_size,
server_args=server_args,
use_mla=use_mla, use_mla=use_mla,
override_kv_cache_dim=override_kv_cache_dim, override_kv_cache_dim=override_kv_cache_dim,
mtp_draft_device_pools=mtp_draft_device_pools, mtp_draft_device_pools=mtp_draft_device_pools,
@@ -1013,7 +1001,6 @@ def build_full_draft_pools(
*, *,
draft_kv_pool: Any, draft_kv_pool: Any,
tree_cache: Any, tree_cache: Any,
server_args: ServerArgs,
) -> tuple[list[SidecarPoolSpec], list[PoolEntry]]: ) -> tuple[list[SidecarPoolSpec], list[PoolEntry]]:
"""Build draft KV/DSA sidecars whose indices follow target full KV.""" """Build draft KV/DSA sidecars whose indices follow target full KV."""
from sglang.srt.mem_cache.memory_pool import ( from sglang.srt.mem_cache.memory_pool import (
@@ -1089,7 +1076,6 @@ def build_swa_draft_pools(
*, *,
draft_kv_pool: Any, draft_kv_pool: Any,
tree_cache: Any, tree_cache: Any,
server_args: ServerArgs,
) -> tuple[list[SidecarPoolSpec], list[PoolEntry]]: ) -> tuple[list[SidecarPoolSpec], list[PoolEntry]]:
"""Build a draft SWA sidecar whose indices follow target SWA.""" """Build a draft SWA sidecar whose indices follow target SWA."""
draft_swa_pool = draft_kv_pool.swa_kv_pool draft_swa_pool = draft_kv_pool.swa_kv_pool
@@ -1143,7 +1129,6 @@ def build_hicache_draft_sidecars(
*, *,
draft_device_pools: tuple[Any, ...], draft_device_pools: tuple[Any, ...],
tree_cache: Any, tree_cache: Any,
server_args: ServerArgs,
) -> tuple[list[SidecarPoolSpec], list[PoolEntry]]: ) -> tuple[list[SidecarPoolSpec], list[PoolEntry]]:
"""Compose the full and SWA draft-sidecar paths.""" """Compose the full and SWA draft-sidecar paths."""
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
@@ -1158,7 +1143,6 @@ def build_hicache_draft_sidecars(
return builder( return builder(
draft_kv_pool=draft_kv_pool, draft_kv_pool=draft_kv_pool,
tree_cache=tree_cache, tree_cache=tree_cache,
server_args=server_args,
) )
@@ -1561,7 +1545,6 @@ class _MiniMaxSparseStrategy(StackStrategy):
): ):
host_pool_group, cache_controller = build_minimax_sparse_hicache_stack( host_pool_group, cache_controller = build_minimax_sparse_hicache_stack(
params=params, params=params,
server_args=server_args,
sparse_pool=kvcache, sparse_pool=kvcache,
load_cache_event=load_cache_event, load_cache_event=load_cache_event,
storage_backend=storage_backend, storage_backend=storage_backend,
@@ -1754,7 +1737,6 @@ def attach_hybrid_pool_to_unified_cache(
def build_minimax_sparse_hicache_stack( def build_minimax_sparse_hicache_stack(
*, *,
params: CacheInitParams, params: CacheInitParams,
server_args: ServerArgs,
sparse_pool: Any, sparse_pool: Any,
load_cache_event, load_cache_event,
storage_backend: Optional[str], storage_backend: Optional[str],
@@ -1787,7 +1769,6 @@ def build_minimax_sparse_hicache_stack(
kv_host_pool = build_kv_host_pool( kv_host_pool = build_kv_host_pool(
kv_pool=main_pool, kv_pool=main_pool,
page_size=params.page_size, page_size=params.page_size,
server_args=server_args,
use_mla=False, use_mla=False,
) )
entries = [ entries = [
@@ -1892,7 +1873,6 @@ def attach_hybrid_minimax_sparse_pool_to_hiradix_cache(
else: else:
host_pool_group, cache_controller = build_minimax_sparse_hicache_stack( host_pool_group, cache_controller = build_minimax_sparse_hicache_stack(
params=params, params=params,
server_args=server_args,
sparse_pool=sparse_pool, sparse_pool=sparse_pool,
load_cache_event=load_cache_event, load_cache_event=load_cache_event,
storage_backend=get_memory().hicache_storage_backend, storage_backend=get_memory().hicache_storage_backend,
@@ -66,7 +66,6 @@ def maybe_register_hicache_draft(
*, *,
tree_cache, tree_cache,
draft_plan: HiCacheDraftPlan, draft_plan: HiCacheDraftPlan,
server_args: ServerArgs,
) -> None: ) -> None:
from sglang.srt.speculative.base_spec_worker import HiCacheDraftMode from sglang.srt.speculative.base_spec_worker import HiCacheDraftMode
@@ -85,7 +84,6 @@ def maybe_register_hicache_draft(
specs, entries = build_hicache_draft_sidecars( specs, entries = build_hicache_draft_sidecars(
draft_device_pools=draft_plan.device_pools, draft_device_pools=draft_plan.device_pools,
tree_cache=tree_cache, tree_cache=tree_cache,
server_args=server_args,
) )
for spec, entry in zip(specs, entries, strict=True): for spec, entry in zip(specs, entries, strict=True):
tree_cache.register_sidecar_pool(spec, entry) tree_cache.register_sidecar_pool(spec, entry)
@@ -322,7 +320,6 @@ def build_kv_cache(
maybe_register_hicache_draft( maybe_register_hicache_draft(
tree_cache=tree_cache, tree_cache=tree_cache,
draft_plan=hicache_draft_plan, draft_plan=hicache_draft_plan,
server_args=server_args,
) )
if retraction_backup == "host_pool": if retraction_backup == "host_pool":
@@ -652,7 +652,6 @@ class ModelRunner:
prepare_moe_topk( prepare_moe_topk(
model=self.model, model=self.model,
model_config=self.model_config, model_config=self.model_config,
server_args=self.server_args,
moe_ep_size=self.ps.moe_ep_size, moe_ep_size=self.ps.moe_ep_size,
moe_ep_rank=self.ps.moe_ep_rank, moe_ep_rank=self.ps.moe_ep_rank,
) )
@@ -1187,7 +1186,6 @@ class ModelRunner:
# before configure_kv_cache_dtype.) # before configure_kv_cache_dtype.)
load_kv_cache_scales( load_kv_cache_scales(
model=self.model, model=self.model,
server_args=self.server_args,
kv_cache_dtype=get_model().kv_cache_dtype, kv_cache_dtype=get_model().kv_cache_dtype,
) )
@@ -105,9 +105,7 @@ def maybe_trigger_remote_instance_nccl_send_group(
t.start() t.start()
def load_kv_cache_scales( def load_kv_cache_scales(*, model, kv_cache_dtype: str) -> None:
*, model, server_args: ServerArgs, kv_cache_dtype: str
) -> None:
"""``kv_cache_dtype`` is the caller's resolved value. Required rather than """``kv_cache_dtype`` is the caller's resolved value. Required rather than
defaulted: a fallback to ``server_args`` would be a hidden global read for defaulted: a fallback to ``server_args`` would be a hidden global read for
any future caller that forgets to pass one.""" any future caller that forgets to pass one."""
@@ -18,7 +18,6 @@ from sglang.srt.utils import get_bool_env_var, is_hip, log_info_on_rank0
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -29,7 +28,6 @@ def prepare_moe_topk(
*, *,
model, model,
model_config: ModelConfig, model_config: ModelConfig,
server_args: ServerArgs,
moe_ep_size: int, moe_ep_size: int,
moe_ep_rank: int, moe_ep_rank: int,
) -> None: ) -> None:
+16 -21
View File
@@ -1279,6 +1279,8 @@ class ServerArgs:
"defaults to --port + 10000.", "defaults to --port + 10000.",
NS("serving"), NS("serving"),
] = None ] = None
# Env-only (SGLANG_GRPC_WORKER_THREADS); a field so the projection sees it.
grpc_worker_threads: A[Optional[int], Arg(no_cli=True), NS("serving")] = None
sidecar: A[ sidecar: A[
Optional[str], Optional[str],
"Start a locally managed sidecar against the native gRPC server. " "Start a locally managed sidecar against the native gRPC server. "
@@ -3696,7 +3698,7 @@ class ServerArgs:
except BaseException: except BaseException:
# The handlers that ran already declared, and they are not # The handlers that ran already declared, and they are not
# idempotent over their own output. # idempotent over their own output.
object.__setattr__(self, "_resolution_failed", True) self._resolution_failed = True
raise raise
# Set here too, because the dummy/absent-model path returns before the # Set here too, because the dummy/absent-model path returns before the
# end of the pipeline that normally sets it: the gate is about whether # end of the pipeline that normally sets it: the gate is about whether
@@ -3747,8 +3749,8 @@ class ServerArgs:
# Everything outside the fields, enumerated from the instance: the raw # Everything outside the fields, enumerated from the instance: the raw
# snapshot, the stash, and what resolution memoized -- including the # snapshot, the stash, and what resolution memoized -- including the
# `get_model_config()` cache, which a resolved copy can no longer fill # `get_model_config()` memo, which the copy carries over rather than
# (the read-only guard refuses the write). # rebuild.
field_names = {field.name for field in dataclasses.fields(self)} field_names = {field.name for field in dataclasses.fields(self)}
for name, value in vars(self).items(): for name, value in vars(self).items():
if name in field_names or name == "_resolution_finished": if name in field_names or name == "_resolution_finished":
@@ -4472,7 +4474,10 @@ class ServerArgs:
# Native gRPC tuning knob is env-only; --grpc-port (CLI) enables the # Native gRPC tuning knob is env-only; --grpc-port (CLI) enables the
# native server, falling back to SGLANG_GRPC_PORT. # native server, falling back to SGLANG_GRPC_PORT.
self.grpc_worker_threads = envs.SGLANG_GRPC_WORKER_THREADS.get() self._declare(
"_handle_deprecated_args",
grpc_worker_threads=envs.SGLANG_GRPC_WORKER_THREADS.get(),
)
grpc_port_env = envs.SGLANG_GRPC_PORT.get() grpc_port_env = envs.SGLANG_GRPC_PORT.get()
if cfg.grpc_port is None and grpc_port_env is not None: if cfg.grpc_port is None and grpc_port_env is not None:
@@ -4496,10 +4501,10 @@ class ServerArgs:
"--grpc-port / SGLANG_GRPC_PORT " "--grpc-port / SGLANG_GRPC_PORT "
f"({cfg.grpc_port}) must be between 1 and 65535" f"({cfg.grpc_port}) must be between 1 and 65535"
) )
if self.grpc_worker_threads < 1: if cfg.grpc_worker_threads is not None and cfg.grpc_worker_threads < 1:
raise ValueError( raise ValueError(
"SGLANG_GRPC_WORKER_THREADS " "SGLANG_GRPC_WORKER_THREADS "
f"({self.grpc_worker_threads}) must be >= 1" f"({cfg.grpc_worker_threads}) must be >= 1"
) )
# Native gRPC is incompatible with launch paths it doesn't wire into. # Native gRPC is incompatible with launch paths it doesn't wire into.
@@ -7271,7 +7276,6 @@ class ServerArgs:
"_handle_dwdp", "_handle_dwdp",
ep_size=cfg.dwdp_size, ep_size=cfg.dwdp_size,
) )
self.moe_ep_size = cfg.dwdp_size
self._declare( self._declare(
"_handle_dwdp", "_handle_dwdp",
moe_dp_size=1, moe_dp_size=1,
@@ -7290,7 +7294,7 @@ class ServerArgs:
logger.info( logger.info(
f"DWDP enabled: dwdp_size={cfg.dwdp_size}, " f"DWDP enabled: dwdp_size={cfg.dwdp_size}, "
f"auto-forced dp_size={cfg.dp_size}, moe_ep_size={self.moe_ep_size}, " f"auto-forced dp_size={cfg.dp_size}, ep_size={cfg.dwdp_size}, "
f"moe_dense_tp_size=1, moe_a2a_backend=none, " f"moe_dense_tp_size=1, moe_a2a_backend=none, "
f"dp_attention_local_control_broadcast=True, " f"dp_attention_local_control_broadcast=True, "
f"enable_dp_lm_head=True, SCHEDULER_SKIP_ALL_GATHER=True, " f"enable_dp_lm_head=True, SCHEDULER_SKIP_ALL_GATHER=True, "
@@ -10098,7 +10102,7 @@ class ServerArgs:
cfg = resolving_view(self) cfg = resolving_view(self)
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
memo = getattr(self, "model_config", None) memo = getattr(self, "_model_config", None)
if memo is not None: if memo is not None:
# The key is the path this record carried when the cache was # The key is the path this record carried when the cache was
# filled. The GGUF and ModelScope handlers declare a different # filled. The GGUF and ModelScope handlers declare a different
@@ -10112,7 +10116,7 @@ class ServerArgs:
return memo return memo
model_config = ModelConfig.from_server_args(self) model_config = ModelConfig.from_server_args(self)
self.model_config = model_config self._model_config = model_config
self._model_config_built_from = cfg.model_path self._model_config_built_from = cfg.model_path
if model_config.is_hybrid_swa: if model_config.is_hybrid_swa:
logger.info( logger.info(
@@ -10147,7 +10151,6 @@ class ServerArgs:
if ( if (
getattr(self, "_resolution_finished", False) getattr(self, "_resolution_finished", False)
and not getattr(self, "_internal_write", False) and not getattr(self, "_internal_write", False)
and name not in _CACHE_SLOTS
and (not name.startswith("_") or name in _underscore_field_names()) and (not name.startswith("_") or name in _underscore_field_names())
): ):
raise AttributeError( raise AttributeError(
@@ -10227,7 +10230,7 @@ class ServerArgs:
# is supported. # is supported.
result = max(candidate_steps) + 1 result = max(candidate_steps) + 1
if getattr(self, "_resolution_finished", False): if getattr(self, "_resolution_finished", False):
object.__setattr__(self, "_max_speculative_num_draft_tokens", result) self._max_speculative_num_draft_tokens = result
return result return result
@property @property
@@ -10851,7 +10854,7 @@ class ServerArgs:
result = json.loads(self.modelexpress_config) result = json.loads(self.modelexpress_config)
else: else:
result = self.modelexpress_config result = self.modelexpress_config
object.__setattr__(self, "_mx_config_cache", result) self._mx_config_cache = result
return result return result
@property @property
@@ -11058,14 +11061,6 @@ def m3_fp8_attn_gemm_enabled(args) -> bool:
) )
# Caches, which the read-only guard lets through: a value the record derived
# from itself is not resolved configuration, and a key that can invalidate on a
# resolved record needs the refill to be storable there. Only the public-named
# ones are listed -- a cache key spelled with a leading underscore is already
# exempt.
_CACHE_SLOTS = frozenset({"model_config"})
# NOTE: The process-wide ServerArgs is owned by the runtime context # NOTE: The process-wide ServerArgs is owned by the runtime context
# (sglang.srt.runtime_context). The two functions below are LEGACY shims kept # (sglang.srt.runtime_context). The two functions below are LEGACY shims kept
# for the existing call-sites; they publish/read the same live object by # for the existing call-sites; they publish/read the same live object by
@@ -40,7 +40,7 @@ def _make_dspark_server_args(
server_args.speculative_algorithm = "DSPARK" server_args.speculative_algorithm = "DSPARK"
server_args.speculative_draft_model_path = None server_args.speculative_draft_model_path = None
server_args.speculative_dspark_block_size = 5 server_args.speculative_dspark_block_size = 5
server_args.model_config = SimpleNamespace(hf_config=hf_config) server_args._model_config = SimpleNamespace(hf_config=hf_config)
return server_args return server_args
@@ -72,7 +72,7 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
def test_supported_multimodal_model_upgrades_default_to_tc_piecewise(self): def test_supported_multimodal_model_upgrades_default_to_tc_piecewise(self):
args = ServerArgs(model_path="dummy") args = ServerArgs(model_path="dummy")
args.model_config = SimpleNamespace( args._model_config = SimpleNamespace(
is_multimodal_piecewise_cuda_graph_supported=True, is_multimodal_piecewise_cuda_graph_supported=True,
is_multimodal_breakable_cuda_graph_supported=False, is_multimodal_breakable_cuda_graph_supported=False,
) )
@@ -101,7 +101,7 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
args = ServerArgs(model_path="dummy") args = ServerArgs(model_path="dummy")
# trtllm_mla skips the tc_piecewise upgrade and keeps breakable, which # trtllm_mla skips the tc_piecewise upgrade and keeps breakable, which
# now serves MLA by falling back to the flashinfer MLA impl for extend. # now serves MLA by falling back to the flashinfer MLA impl for extend.
args.model_config = SimpleNamespace( args._model_config = SimpleNamespace(
is_multimodal_piecewise_cuda_graph_supported=True, is_multimodal_piecewise_cuda_graph_supported=True,
is_multimodal=False, is_multimodal=False,
is_multimodal_breakable_cuda_graph_supported=False, is_multimodal_breakable_cuda_graph_supported=False,
@@ -169,7 +169,7 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
def test_embedding_gemma_forces_breakable_prefill(self): def test_embedding_gemma_forces_breakable_prefill(self):
args = ServerArgs(model_path="dummy") args = ServerArgs(model_path="dummy")
args.model_config = SimpleNamespace( args._model_config = SimpleNamespace(
is_embedding_gemma=True, is_embedding_gemma=True,
is_multimodal=False, is_multimodal=False,
context_len=2048, context_len=2048,
@@ -183,7 +183,7 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
args.chunked_prefill_size = 2048 args.chunked_prefill_size = 2048
with ( with (
patch.object(args, "get_model_config", return_value=args.model_config), patch.object(args, "get_model_config", return_value=args._model_config),
patch("sglang.srt.server_args.is_cuda", return_value=True), patch("sglang.srt.server_args.is_cuda", return_value=True),
): ):
args._handle_model_capability_adjustments() args._handle_model_capability_adjustments()
@@ -202,7 +202,7 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
def test_encoder_embedding_model_enables_embedding_mode_without_flag(self): def test_encoder_embedding_model_enables_embedding_mode_without_flag(self):
args = ServerArgs(model_path="dummy") args = ServerArgs(model_path="dummy")
args.is_embedding = False args.is_embedding = False
args.model_config = SimpleNamespace( args._model_config = SimpleNamespace(
embedding_model_spec=resolve_embedding_model_spec( embedding_model_spec=resolve_embedding_model_spec(
["BertModel"], ["BertModel"],
is_embedding_requested=False, is_embedding_requested=False,
@@ -212,7 +212,7 @@ class TestMultimodalPiecewiseCudaGraph(CustomTestCase):
hf_config=SimpleNamespace(architectures=["BertModel"]), hf_config=SimpleNamespace(architectures=["BertModel"]),
) )
with patch.object(args, "get_model_config", return_value=args.model_config): with patch.object(args, "get_model_config", return_value=args._model_config):
args._handle_model_capability_adjustments() args._handle_model_capability_adjustments()
self.assertTrue(resolution_result(args, "is_embedding")) self.assertTrue(resolution_result(args, "is_embedding"))
@@ -116,7 +116,6 @@ class TestDecodeRetractionBackup(unittest.TestCase):
mode=HiCacheDraftMode.SIDECAR, mode=HiCacheDraftMode.SIDECAR,
device_pools=(draft_pool,), device_pools=(draft_pool,),
), ),
server_args=server_args,
) )
self.assertIn(PoolName.DRAFT, cache.host_pool_group.entry_map) self.assertIn(PoolName.DRAFT, cache.host_pool_group.entry_map)
cache.validate_retraction_host_capacity() cache.validate_retraction_host_capacity()
@@ -86,7 +86,6 @@ class TestDraftSidecarPoolDispatch(CustomTestCase):
specs, entries = build_full_draft_pools( specs, entries = build_full_draft_pools(
draft_kv_pool=draft_kv_pool, draft_kv_pool=draft_kv_pool,
tree_cache=None, tree_cache=None,
server_args=None,
) )
self.assertEqual(specs, []) self.assertEqual(specs, [])
@@ -124,7 +123,6 @@ class TestDraftSidecarPoolDispatch(CustomTestCase):
specs, entries = build_full_draft_pools( specs, entries = build_full_draft_pools(
draft_kv_pool=draft_kv_pool, draft_kv_pool=draft_kv_pool,
tree_cache=tree_cache, tree_cache=tree_cache,
server_args=server_args,
) )
self.assertEqual(build_host_pool.call_args.kwargs["host_to_device_ratio"], 1.0) self.assertEqual(build_host_pool.call_args.kwargs["host_to_device_ratio"], 1.0)
@@ -110,7 +110,7 @@ class TestTheModelConfigCache(CustomTestCase):
server_args = self._resolved( server_args = self._resolved(
model_path=_OBJECT_STORE_URI, load_format="runai_streamer" model_path=_OBJECT_STORE_URI, load_format="runai_streamer"
) )
cached = server_args.__dict__["model_config"] cached = server_args.__dict__["_model_config"]
self.assertIsInstance(cached, ModelConfig) self.assertIsInstance(cached, ModelConfig)
# The record still carries the URI the operator typed, and the # The record still carries the URI the operator typed, and the
# configuration carries the directory it read the metadata from. # configuration carries the directory it read the metadata from.
@@ -162,7 +162,7 @@ class TestTheModelConfigCache(CustomTestCase):
invalidates it.""" invalidates it."""
server_args = ServerArgs(model_path=self._checkpoint(), device="cuda") server_args = ServerArgs(model_path=self._checkpoint(), device="cuda")
stand_in = SimpleNamespace(model_path="somewhere/else") stand_in = SimpleNamespace(model_path="somewhere/else")
server_args.model_config = stand_in server_args._model_config = stand_in
self.assertIs(server_args.get_model_config(), stand_in) self.assertIs(server_args.get_model_config(), stand_in)
@@ -0,0 +1,83 @@
"""The record grows no attribute the projection cannot see.
A publicly-named attribute that is not a dataclass field is invisible to every
other guard here: the namespace coverage walks fields, the projection walks
fields, and the read ratchets watch field reads. Three of them accumulated that
way -- a `ModelConfig` cache, an `moe_ep_size` that only a log line read, and an
env-derived `grpc_worker_threads` that one entry point read across the boundary.
Leading-underscore names are the record's own bookkeeping and stay: the
read-only guard classifies writability by that spelling, so a private name is
already outside the config tier by construction.
"""
import ast
import dataclasses
import pathlib
import unittest
import sglang
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=4, suite="base-a-test-cpu")
def _self_written_attributes() -> set:
"""Names `ServerArgs` writes on itself, by either spelling."""
source = (
pathlib.Path(next(iter(sglang.__path__))) / "srt" / "server_args.py"
).read_text(encoding="utf-8-sig")
tree = ast.parse(source)
cls = next(
node
for node in tree.body
if isinstance(node, ast.ClassDef) and node.name == "ServerArgs"
)
written = set()
for node in ast.walk(cls):
if isinstance(node, ast.Assign):
for target in node.targets:
if (
isinstance(target, ast.Attribute)
and isinstance(target.value, ast.Name)
and target.value.id == "self"
):
written.add(target.attr)
if (
isinstance(node, ast.Call)
and getattr(node.func, "attr", None) == "__setattr__"
and getattr(getattr(node.func, "value", None), "id", None) == "object"
and len(node.args) >= 2
and isinstance(node.args[1], ast.Constant)
):
written.add(node.args[1].value)
return written
class TestNoPublicNonFieldSlot(CustomTestCase):
def test_every_public_attribute_is_a_field(self):
written = _self_written_attributes()
self.assertGreater(
len(written),
5,
f"only {len(written)} self-writes found; the scan is broken, not the "
"record",
)
fields = {field.name for field in dataclasses.fields(ServerArgs)}
stray = sorted(
name for name in written if not name.startswith("_") and name not in fields
)
self.assertEqual(
[],
stray,
"these are written on the record under a public name but are not "
"fields, so the projection cannot see them and no other guard "
"watches them: make each a field, or give it the leading underscore "
f"that says it is the record's own bookkeeping: {stray}",
)
if __name__ == "__main__":
unittest.main()
@@ -283,7 +283,7 @@ class TestImageProcessorBackend(CustomTestCase):
class TestMultimodalFeatureTransport(CustomTestCase): class TestMultimodalFeatureTransport(CustomTestCase):
@staticmethod @staticmethod
def _set_model_type(server_args, *, is_multimodal): def _set_model_type(server_args, *, is_multimodal):
server_args.model_config = SimpleNamespace(is_multimodal=is_multimodal) server_args._model_config = SimpleNamespace(is_multimodal=is_multimodal)
@patch("sglang.srt.server_args.is_cuda", return_value=True) @patch("sglang.srt.server_args.is_cuda", return_value=True)
def test_cuda_ipc_is_explicit_and_bounded(self, _mock_is_cuda): def test_cuda_ipc_is_explicit_and_bounded(self, _mock_is_cuda):
@@ -920,8 +920,8 @@ class TestFa4PageSizeAutoForce(CustomTestCase):
# use_mla_backend() (mocked) and is_sm100_supported() (mocked), not a # use_mla_backend() (mocked) and is_sm100_supported() (mocked), not a
# real model_config. Pre-set the attribute so get_model_config returns # real model_config. Pre-set the attribute so get_model_config returns
# early without touching ModelConfig.from_server_args. # early without touching ModelConfig.from_server_args.
args.model_config = MagicMock() args._model_config = MagicMock()
args.model_config.hf_config.dual_chunk_attention_config = None args._model_config.hf_config.dual_chunk_attention_config = None
return args return args
@patch("sglang.srt.arg_groups.overrides.is_sm100_supported", return_value=True) @patch("sglang.srt.arg_groups.overrides.is_sm100_supported", return_value=True)
@@ -1809,7 +1809,7 @@ class TestCudaGraphConfigDataclassAccess(CustomTestCase):
class TestCudaGraphDisaggregationRoles(CustomTestCase): class TestCudaGraphDisaggregationRoles(CustomTestCase):
def _handled_args(self, **overrides): def _handled_args(self, **overrides):
args = ServerArgs(model_path="dummy", **overrides) args = ServerArgs(model_path="dummy", **overrides)
args.model_config = SimpleNamespace( args._model_config = SimpleNamespace(
hf_config=SimpleNamespace(architectures=["LlamaForCausalLM"]), hf_config=SimpleNamespace(architectures=["LlamaForCausalLM"]),
is_piecewise_cuda_graph_disabled_model=False, is_piecewise_cuda_graph_disabled_model=False,
is_multimodal=False, is_multimodal=False,
@@ -1882,7 +1882,7 @@ class TestPrefillCudaGraphLoRACompatibility(CustomTestCase):
def _handled_args(self, **overrides): def _handled_args(self, **overrides):
args = ServerArgs(model_path="dummy", **overrides) args = ServerArgs(model_path="dummy", **overrides)
args.model_config = SimpleNamespace( args._model_config = SimpleNamespace(
hf_config=SimpleNamespace(architectures=["LlamaForCausalLM"]), hf_config=SimpleNamespace(architectures=["LlamaForCausalLM"]),
is_piecewise_cuda_graph_disabled_model=False, is_piecewise_cuda_graph_disabled_model=False,
is_multimodal=False, is_multimodal=False,
@@ -1915,7 +1915,7 @@ class TestPrefillCudaGraphLoRACompatibility(CustomTestCase):
# Pin the tc_piecewise LoRA rule itself, with the hardware rule # Pin the tc_piecewise LoRA rule itself, with the hardware rule
# neutralized so this runs on CPU-only CI. # neutralized so this runs on CPU-only CI.
args = ServerArgs(model_path="dummy", enable_lora=True) args = ServerArgs(model_path="dummy", enable_lora=True)
args.model_config = SimpleNamespace( args._model_config = SimpleNamespace(
hf_config=SimpleNamespace(architectures=["LlamaForCausalLM"]), hf_config=SimpleNamespace(architectures=["LlamaForCausalLM"]),
is_piecewise_cuda_graph_disabled_model=False, is_piecewise_cuda_graph_disabled_model=False,
is_multimodal=False, is_multimodal=False,
@@ -1945,7 +1945,7 @@ class TestBreakableCudaGraphMultimodalAllowlist(CustomTestCase):
def _handled_args(self, *, architectures, is_multimodal, allowlisted): def _handled_args(self, *, architectures, is_multimodal, allowlisted):
args = ServerArgs(model_path="dummy") args = ServerArgs(model_path="dummy")
args.model_config = SimpleNamespace( args._model_config = SimpleNamespace(
hf_config=SimpleNamespace(architectures=architectures), hf_config=SimpleNamespace(architectures=architectures),
is_piecewise_cuda_graph_disabled_model=False, is_piecewise_cuda_graph_disabled_model=False,
is_multimodal=is_multimodal, is_multimodal=is_multimodal,
@@ -2129,7 +2129,7 @@ class TestDeepEPv2Args(CustomTestCase):
def _args(self, **overrides): def _args(self, **overrides):
server_args = ServerArgs(model_path="dummy", moe_a2a_backend="deepep_v2") server_args = ServerArgs(model_path="dummy", moe_a2a_backend="deepep_v2")
server_args.model_config = SimpleNamespace( server_args._model_config = SimpleNamespace(
hf_config=SimpleNamespace(architectures=["DeepseekV4ForCausalLM"]) hf_config=SimpleNamespace(architectures=["DeepseekV4ForCausalLM"])
) )
# The dummy path does not initialize phase configs. # The dummy path does not initialize phase configs.
@@ -2152,7 +2152,7 @@ class TestDeepEPv2Args(CustomTestCase):
"Qwen3MoeForCausalLM", "Qwen3MoeForCausalLM",
): ):
args = self._args(moe_runner_backend="deep_gemm") args = self._args(moe_runner_backend="deep_gemm")
args.model_config.hf_config.architectures = [architecture] args._model_config.hf_config.architectures = [architecture]
args._handle_a2a_moe() args._handle_a2a_moe()
def test_unvalidated_and_missing_architectures_rejected(self): def test_unvalidated_and_missing_architectures_rejected(self):
@@ -2163,7 +2163,7 @@ class TestDeepEPv2Args(CustomTestCase):
None, None,
): ):
args = self._args(moe_runner_backend="deep_gemm") args = self._args(moe_runner_backend="deep_gemm")
args.model_config.hf_config.architectures = architectures args._model_config.hf_config.architectures = architectures
with self.assertRaisesRegex(ValueError, "not validated"): with self.assertRaisesRegex(ValueError, "not validated"):
args._handle_a2a_moe() args._handle_a2a_moe()
@@ -2188,7 +2188,7 @@ class TestDeepEPv2Args(CustomTestCase):
moe_runner_backend="deep_gemm", moe_runner_backend="deep_gemm",
rl_on_policy_target="fsdp", rl_on_policy_target="fsdp",
) )
args.model_config.hf_config.architectures = ["Qwen3MoeForCausalLM"] args._model_config.hf_config.architectures = ["Qwen3MoeForCausalLM"]
with ( with (
envs.SGLANG_VLM_CACHE_SIZE_MB.override(envs.SGLANG_VLM_CACHE_SIZE_MB.get()), envs.SGLANG_VLM_CACHE_SIZE_MB.override(envs.SGLANG_VLM_CACHE_SIZE_MB.get()),
envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.override( envs.SGLANG_ENABLE_DETERMINISTIC_INFERENCE.override(
@@ -2514,7 +2514,7 @@ class TestGrpcServerArgs(CustomTestCase):
with envs.SGLANG_GRPC_WORKER_THREADS.override(8): with envs.SGLANG_GRPC_WORKER_THREADS.override(8):
sa._handle_deprecated_args() sa._handle_deprecated_args()
self.assertEqual(resolution_result(sa, "grpc_port"), 50051) self.assertEqual(resolution_result(sa, "grpc_port"), 50051)
self.assertEqual(sa.grpc_worker_threads, 8) self.assertEqual(resolution_result(sa, "grpc_worker_threads"), 8)
def test_env_grpc_port_enables_native(self): def test_env_grpc_port_enables_native(self):
sa = self._args(port=30000) sa = self._args(port=30000)
@@ -2698,13 +2698,11 @@ class TestGrpcServerArgs(CustomTestCase):
fake_core = SimpleNamespace(start_server=MagicMock(return_value="handle")) fake_core = SimpleNamespace(start_server=MagicMock(return_value="handle"))
fake_bridge = SimpleNamespace(RuntimeHandle=MagicMock(return_value="rt")) fake_bridge = SimpleNamespace(RuntimeHandle=MagicMock(return_value="rt"))
# The host comes from the `serving` bag; `grpc_worker_threads` is not a override = get_context().override_server_args(
# field (resolution sets it from the environment), so it stays on the host="127.0.0.1", grpc_port=50051, grpc_worker_threads=4
# stand-in the call site is handed. )
override = get_context().override_server_args(host="127.0.0.1", grpc_port=50051) server_args = override.install()
override.install()
self.addCleanup(override.restore) self.addCleanup(override.restore)
server_args = SimpleNamespace(grpc_worker_threads=4)
with ( with (
patch( patch(
"sglang.srt.rust_extensions.load_rust_extension", "sglang.srt.rust_extensions.load_rust_extension",
@@ -2728,6 +2726,7 @@ class TestGrpcServerArgs(CustomTestCase):
self.assertEqual( self.assertEqual(
set(kwargs), {"host", "port", "runtime_handle", "worker_threads"} set(kwargs), {"host", "port", "runtime_handle", "worker_threads"}
) )
self.assertEqual(kwargs["worker_threads"], 4)
self.assertNotIn("max_prefill_tokens", kwargs) self.assertNotIn("max_prefill_tokens", kwargs)
@@ -0,0 +1,86 @@
"""A function does not take the record it never reads.
A `server_args` parameter that the body never names keeps a reference to the
whole record alive across a call boundary, and it reads as an invitation: the
next person to need one value takes it off the parameter that is already there,
instead of deciding where that value should come from. Removing one usually
uncovers the next -- the caller that only had a record to pass it along.
Class methods are exempt: a base class, an override, or one implementation of a
strategy carries the parameter for its contract, and the body of any single one
of them is not evidence. This walks module-level functions only.
"""
import ast
import pathlib
import unittest
import sglang
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=6, suite="base-a-test-cpu")
_PACKAGE_ROOT = pathlib.Path(next(iter(sglang.__path__)))
# The resolution pipeline builds the record, so a parameter there is the subject
# rather than a passenger. `multimodal_gen` has a different, same-named class
# outside this contract, as the mutation ratchet also records.
_EXCLUDED = ("srt/arg_groups", "srt/server_args.py", "multimodal_gen")
_BASELINE = 0
def _dead_parameters():
found = []
scanned = 0
for path in sorted(_PACKAGE_ROOT.rglob("*.py")):
rel = path.relative_to(_PACKAGE_ROOT).as_posix()
if rel.startswith(_EXCLUDED):
continue
source = path.read_text(encoding="utf-8-sig")
if "server_args" not in source:
continue
scanned += 1
try:
tree = ast.parse(source)
except SyntaxError:
continue
for node in tree.body:
if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
continue
taken = [a.arg for a in node.args.args] + [
a.arg for a in node.args.kwonlyargs
]
if "server_args" not in taken:
continue
named = any(
isinstance(inner, ast.Name) and inner.id == "server_args"
for inner in ast.walk(node)
if inner is not node
)
if not named:
found.append(f"{rel}:{node.lineno} {node.name}")
return found, scanned
class TestNoDeadServerArgsParameter(CustomTestCase):
def test_no_module_level_function_takes_a_record_it_ignores(self):
found, scanned = _dead_parameters()
self.assertGreater(
scanned,
50,
f"only {scanned} files mention server_args; the scan is broken, not "
"the tree",
)
self.assertEqual(
_BASELINE,
len(found),
"these functions take `server_args` and never name it; drop the "
"parameter and the argument at every call site, then check whether "
f"the caller still needs its own: {found}",
)
if __name__ == "__main__":
unittest.main()