config: constructing a config no longer resolves it (#35907)

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-08-23 01:18:53 -07:00
committed by GitHub
co-authored by Claude Opus 5
parent 4bc79a1b49
commit 64aa859da2
34 changed files with 947 additions and 107 deletions
+2 -2
View File
@@ -19,7 +19,6 @@ llm = Engine(
) )
""" """
import dataclasses
from argparse import ArgumentParser from argparse import ArgumentParser
from pathlib import Path from pathlib import Path
@@ -44,11 +43,12 @@ parser.add_argument(
def main(args): def main(args):
engine_args = ServerArgs.from_cli_args(args) engine_args = ServerArgs.from_cli_args(args)
engine_args.resolve_once()
model_path = engine_args.model_path model_path = engine_args.model_path
if not Path(model_path).is_dir(): if not Path(model_path).is_dir():
raise ValueError("model path must be a local directory") raise ValueError("model path must be a local directory")
# Create LLM instance from arguments # Create LLM instance from arguments
llm = Engine(**dataclasses.asdict(engine_args)) llm = Engine(server_args=engine_args)
llm.save_remote_model( llm.save_remote_model(
url=args.remote_model_save_url, draft_url=args.remote_draft_model_save_url url=args.remote_model_save_url, draft_url=args.remote_draft_model_save_url
) )
@@ -22,7 +22,6 @@ llm = Engine(
) )
""" """
import dataclasses
import os import os
import shutil import shutil
from argparse import ArgumentParser from argparse import ArgumentParser
@@ -49,11 +48,12 @@ parser.add_argument(
def main(args): def main(args):
engine_args = ServerArgs.from_cli_args(args) engine_args = ServerArgs.from_cli_args(args)
engine_args.resolve_once()
model_path = engine_args.model_path model_path = engine_args.model_path
if not Path(model_path).is_dir(): if not Path(model_path).is_dir():
raise ValueError("model path must be a local directory") raise ValueError("model path must be a local directory")
# Create LLM instance from arguments # Create LLM instance from arguments
llm = Engine(**dataclasses.asdict(engine_args)) llm = Engine(server_args=engine_args)
Path(args.output).mkdir(exist_ok=True) Path(args.output).mkdir(exist_ok=True)
llm.save_sharded_model( llm.save_sharded_model(
path=args.output, pattern=args.file_pattern, max_size=args.max_file_size path=args.output, pattern=args.file_pattern, max_size=args.max_file_size
@@ -1,5 +1,4 @@
import argparse import argparse
import dataclasses
from typing import Tuple from typing import Tuple
from transformers import AutoProcessor from transformers import AutoProcessor
@@ -45,7 +44,7 @@ def token_in_out_example(
model_override_args=server_args.json_model_override_args, model_override_args=server_args.json_model_override_args,
), ),
) )
backend = Engine(**dataclasses.asdict(server_args)) backend = Engine(server_args=server_args)
output = backend.generate( output = backend.generate(
input_ids=input_ids, input_ids=input_ids,
@@ -71,4 +70,5 @@ if __name__ == "__main__":
] ]
args = parser.parse_args(args=args) args = parser.parse_args(args=args)
server_args = ServerArgs.from_cli_args(args) server_args = ServerArgs.from_cli_args(args)
server_args.resolve_once()
token_in_out_example(server_args) token_in_out_example(server_args)
+5
View File
@@ -48,6 +48,11 @@ def _launch_server_target(launch_server_func: Callable, server_args: ServerArgs)
def launch_or_reuse_server(launch_server_func: Callable, server_args: ServerArgs): def launch_or_reuse_server(launch_server_func: Callable, server_args: ServerArgs):
# Resolve in the parent, before the fork. The pipeline probes the device
# (the default attention backend reads the CUDA capability), and a forked
# child cannot re-initialize CUDA once this process has.
server_args.resolve_once()
base_url = resolve_base_url("", server_args.host, server_args.port) base_url = resolve_base_url("", server_args.host, server_args.port)
# Reuse an already-running server instead of forking a second one onto the # Reuse an already-running server instead of forking a second one onto the
@@ -398,7 +398,7 @@ def _create_ray_engine_backend(server_args: ServerArgs):
placement_group=pg, placement_group=pg,
placement_group_bundle_index=0, placement_group_bundle_index=0,
), ),
).remote(**dataclasses.asdict(server_args)) ).remote(**dict(server_args._raw_input))
class _Proxy: class _Proxy:
"""Forwards method calls to the remote RayEngine actor.""" """Forwards method calls to the remote RayEngine actor."""
@@ -432,15 +432,18 @@ def throughput_test(
server_args: ServerArgs, server_args: ServerArgs,
bench_args: BenchArgs, bench_args: BenchArgs,
): ):
# A programmatic caller may hand over a freshly constructed record, and
# the backends below read the resolved paths and the raw snapshot.
server_args.resolve_once()
if bench_args.backend == "engine": if bench_args.backend == "engine":
if server_args.use_ray: if server_args.use_ray:
backend = _create_ray_engine_backend(server_args) backend = _create_ray_engine_backend(server_args)
else: else:
backend = Engine(**dataclasses.asdict(server_args)) backend = Engine(server_args=server_args)
if not backend: if not backend:
raise ValueError("Please provide valid engine arguments") raise ValueError("Please provide valid engine arguments")
elif bench_args.backend == "runtime": elif bench_args.backend == "runtime":
backend = Runtime(**dataclasses.asdict(server_args)) backend = Runtime(**dict(server_args._raw_input))
else: else:
raise ValueError('Please set backend to either "engine" or "runtime"') raise ValueError('Please set backend to either "engine" or "runtime"')
@@ -569,6 +572,7 @@ def cli_main():
raise e raise e
server_args = ServerArgs.from_cli_args(args) server_args = ServerArgs.from_cli_args(args)
server_args.resolve_once()
bench_args = BenchArgs.from_cli_args(args) bench_args = BenchArgs.from_cli_args(args)
logging.basicConfig( logging.basicConfig(
+4 -2
View File
@@ -989,8 +989,10 @@ def latency_test(
def main(server_args, bench_args): def main(server_args, bench_args):
# Post-init write to the legacy cuda_graph_max_bs_decode field would server_args.resolve_once()
# not propagate to cuda_graph_config; update the decode phase directly.
# The legacy cuda_graph_max_bs_decode field does not propagate; set the
# decode phase.
if server_args.cuda_graph_config is not None: if server_args.cuda_graph_config is not None:
server_args.cuda_graph_config[Phase.DECODE].max_bs = max(bench_args.batch_size) server_args.cuda_graph_config[Phase.DECODE].max_bs = max(bench_args.batch_size)
@@ -1255,6 +1255,7 @@ def cli_main():
args = parser.parse_args() args = parser.parse_args()
server_args = ServerArgs.from_cli_args(args) server_args = ServerArgs.from_cli_args(args)
server_args.resolve_once()
bench_args = BenchArgs.from_cli_args(args) bench_args = BenchArgs.from_cli_args(args)
run_benchmark(server_args, bench_args) run_benchmark(server_args, bench_args)
+6
View File
@@ -102,6 +102,10 @@ def launch_server_internal(server_args):
def launch_server_process_and_send_one_request( def launch_server_process_and_send_one_request(
server_args: ServerArgs, compile_args: CompileArgs server_args: ServerArgs, compile_args: CompileArgs
): ):
# Keeps the device probe out of the fork below, for a caller that reaches
# this without resolving first.
server_args.resolve_once()
proc = multiprocessing.Process(target=launch_server_internal, args=(server_args,)) proc = multiprocessing.Process(target=launch_server_internal, args=(server_args,))
proc.start() proc.start()
base_url = f"http://{server_args.host}:{server_args.port}" base_url = f"http://{server_args.host}:{server_args.port}"
@@ -177,6 +181,8 @@ def compile_server_args(args, compile_args: CompileArgs) -> ServerArgs:
args.watchdog_timeout = compile_args.timeout args.watchdog_timeout = compile_args.timeout
args.warmups = "compile-deep-gemm" args.warmups = "compile-deep-gemm"
server_args = ServerArgs.from_cli_args(args) server_args = ServerArgs.from_cli_args(args)
# `cuda_graph_config` is None until resolution parses it.
server_args.resolve_once()
server_args.cuda_graph_config[Phase.DECODE].backend = Backend.DISABLED server_args.cuda_graph_config[Phase.DECODE].backend = Backend.DISABLED
server_args.cuda_graph_config[Phase.PREFILL].backend = Backend.DISABLED server_args.cuda_graph_config[Phase.PREFILL].backend = Backend.DISABLED
print(f"Disable CUDA Graph and Torch Compile to save time...") print(f"Disable CUDA Graph and Torch Compile to save time...")
@@ -393,6 +393,10 @@ class Runtime:
if is_port_available(port): if is_port_available(port):
break break
self.server_args = ServerArgs(*args, log_level=log_level, port=port, **kwargs) self.server_args = ServerArgs(*args, log_level=log_level, port=port, **kwargs)
# The spawned server gets a copy of this record, and this object keeps
# reading it afterwards -- `get_tokenizer` wants the downloaded GGUF
# file and the rewritten ModelScope path, not what the caller typed.
self.server_args.resolve_once()
self.url = self.server_args.url() self.url = self.server_args.url()
self.generate_url = self.url + "/generate" self.generate_url = self.url + "/generate"
@@ -455,7 +459,7 @@ class Runtime:
from sglang.srt.utils.hf_transformers_utils import get_tokenizer from sglang.srt.utils.hf_transformers_utils import get_tokenizer
return get_tokenizer( return get_tokenizer(
self.server_args.tokenizer_path, self.server_args.tokenizer_path or self.server_args.model_path,
tokenizer_mode=self.server_args.tokenizer_mode, tokenizer_mode=self.server_args.tokenizer_mode,
trust_remote_code=self.server_args.trust_remote_code, trust_remote_code=self.server_args.trust_remote_code,
revision=self.server_args.revision, revision=self.server_args.revision,
+4
View File
@@ -15,6 +15,10 @@ suppress_noisy_warnings()
def run_server(server_args): def run_server(server_args):
"""Run the server based on the gRPC flags and server_args.encoder_only.""" """Run the server based on the gRPC flags and server_args.encoder_only."""
# The flags dispatched on below are decided by resolution (`--grpc-mode`
# folds into `smg_grpc_mode`), and `prepare_server_args` returns raw input.
server_args.resolve_once()
if server_args.encoder_only: if server_args.encoder_only:
# For encoder disaggregation # For encoder disaggregation
if server_args.smg_grpc_mode or server_args.grpc_mode: if server_args.smg_grpc_mode or server_args.grpc_mode:
@@ -287,11 +287,11 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
dist_timeout=self.server_args.dist_timeout, dist_timeout=self.server_args.dist_timeout,
) )
from sglang.srt.runtime_context import get_context from sglang.srt.runtime_context import get_context, publish
from sglang.srt.server_args import ServerArgs as SrtServerArgs from sglang.srt.server_args import ServerArgs as SrtServerArgs
if get_context()._server_args is None: if get_context()._server_args is None:
get_context().set_server_args(SrtServerArgs(model_path="dummy")) publish(SrtServerArgs(model_path="dummy"), role="diffusion_gpu_worker")
# set proc title # set proc title
if model_parallel_is_initialized(): if model_parallel_is_initialized():
@@ -20,7 +20,6 @@ from __future__ import annotations
import json import json
import unittest import unittest
from contextlib import contextmanager from contextlib import contextmanager
from types import SimpleNamespace
import torch import torch
@@ -67,7 +66,12 @@ def _srt_trace_server_args():
prev_server_args = srt_server_args_module.get_global_server_args() prev_server_args = srt_server_args_module.get_global_server_args()
except ValueError: # nothing published yet except ValueError: # nothing published yet
prev_server_args = None prev_server_args = None
set_global_server_args_for_scheduler(SimpleNamespace(trace_modules="request")) # publish resolves what it is handed, so a stand-in cannot go through it.
from sglang.srt.server_args import ServerArgs as SrtServerArgs
set_global_server_args_for_scheduler(
SrtServerArgs(model_path="dummy", trace_modules="request")
)
try: try:
yield yield
finally: finally:
+2
View File
@@ -1081,6 +1081,8 @@ class Engine(EngineScoreMixin, EngineBase):
# Configure global environment # Configure global environment
configure_logger(server_args) configure_logger(server_args)
server_args.resolve_once()
_set_envs_and_config(server_args) _set_envs_and_config(server_args)
# Defensive: ensure plugins loaded (may already be loaded by # Defensive: ensure plugins loaded (may already be loaded by
@@ -12,6 +12,10 @@ from sglang.srt.utils import MultiprocessingSerializer, kill_process_tree
def launch_server_process(server_args: ServerArgs) -> multiprocessing.Process: def launch_server_process(server_args: ServerArgs) -> multiprocessing.Process:
# Resolve here, not in the child: the pipeline probes the device, and a
# forked child cannot re-initialize CUDA if this process already has. The
# child's gate then finds nothing left to do.
server_args.resolve_once()
p = multiprocessing.Process(target=launch_server, args=(server_args,)) p = multiprocessing.Process(target=launch_server, args=(server_args,))
p.start() p.start()
+2 -2
View File
@@ -464,8 +464,8 @@ class RayEngine(Engine):
# Set dist_init_addr on server_args so PortArgs.init_new() can compute # Set dist_init_addr on server_args so PortArgs.init_new() can compute
# TCP addresses correctly (required for DP attention path). # TCP addresses correctly (required for DP attention path).
dp_server_args = dataclasses.replace( dp_server_args = server_args.replace_resolved(
server_args, "ray.dp_controller",
dist_init_addr=f"{rank0_node_ip}:{port_args.nccl_port}", dist_init_addr=f"{rank0_node_ip}:{port_args.nccl_port}",
) )
# Create the DP controller in-process. This blocks until all actors # Create the DP controller in-process. This blocks until all actors
+4 -5
View File
@@ -47,8 +47,6 @@ class SchedulerActor:
dp_rank: Optional[int], dp_rank: Optional[int],
dist_init_addr: Optional[str] = None, dist_init_addr: Optional[str] = None,
): ):
import dataclasses
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.scheduler import Scheduler, configure_scheduler_process from sglang.srt.managers.scheduler import Scheduler, configure_scheduler_process
from sglang.srt.utils.numa_utils import ( from sglang.srt.utils.numa_utils import (
@@ -56,10 +54,11 @@ class SchedulerActor:
numa_bind_to_node, numa_bind_to_node,
) )
# Override dist_init_addr if provided (for multi-node) # Override dist_init_addr if provided (for multi-node), through
# `replace_resolved` so the copy keeps the parent's resolution.
if dist_init_addr: if dist_init_addr:
server_args = dataclasses.replace( server_args = server_args.replace_resolved(
server_args, dist_init_addr=dist_init_addr "ray.scheduler_actor", dist_init_addr=dist_init_addr
) )
# Get actual GPU IDs from Ray runtime context # Get actual GPU IDs from Ray runtime context
+9 -8
View File
@@ -1034,10 +1034,11 @@ class _ServerArgsOverride:
self._prev_capture = ctx.flags.capture.enable_torch_compile self._prev_capture = ctx.flags.capture.enable_torch_compile
from sglang.srt.arg_groups.overrides import ( from sglang.srt.arg_groups.overrides import (
_apply_fields, _apply_fields,
declare_resolution, declare_late_resolution,
) )
server_args = ServerArgs(model_path="dummy") server_args = ServerArgs(model_path="dummy")
server_args.resolve_once()
# Underscore names seed private property caches (the strict guard # Underscore names seed private property caches (the strict guard
# exempts them); everything else must be a real config field. # exempts them); everything else must be a real config field.
unknown = {name for name in self._fields if not name.startswith("_")} - set( unknown = {name for name in self._fields if not name.startswith("_")} - set(
@@ -1047,23 +1048,19 @@ class _ServerArgsOverride:
raise ValueError( raise ValueError(
f"override_server_args: unknown ServerArgs field(s): {sorted(unknown)}" f"override_server_args: unknown ServerArgs field(s): {sorted(unknown)}"
) )
# Declared so the projection sees it. # Declared so the projection sees it; late, because the record is
# resolved already and not yet published.
# Underscore names are not fields at all (they seed private property # Underscore names are not fields at all (they seed private property
# caches), so they stay a direct write. # caches), so they stay a direct write.
declared = { declared = {
name: value for name, value in self._fields.items() if name[0] != "_" name: value for name, value in self._fields.items() if name[0] != "_"
} }
if declared: if declared:
declare_resolution(server_args, "override_server_args", **declared) declare_late_resolution(server_args, "override_server_args", **declared)
_apply_fields( _apply_fields(
server_args, server_args,
{name: value for name, value in self._fields.items() if name[0] == "_"}, {name: value for name, value in self._fields.items() if name[0] == "_"},
) )
# The dummy boundary skips materialization, which would leave the
# strict mutation guard unarmed on the published object — mark it
# materialized so bare post-publish writes raise like they do on a
# fully resolved config.
object.__setattr__(server_args, "_declarations_materialized", True)
ctx.set_server_args(server_args) ctx.set_server_args(server_args)
self._installed = True self._installed = True
return server_args return server_args
@@ -1205,6 +1202,9 @@ ROLE_NAMESPACE_SETS: dict[str, frozenset[str] | None] = {
"encoder": None, "encoder": None,
"expert_backup": None, "expert_backup": None,
"weight_cache_daemon": None, "weight_cache_daemon": None,
# The diffusion GPU worker runs a model and publishes a placeholder so
# shared SRT reads do not fail closed; declared full for that reason.
"diffusion_gpu_worker": None,
} }
@@ -1325,6 +1325,7 @@ def publish(server_args, *, role: str, hf_config: Any = None) -> RuntimeContext:
f"publish role {role!r} has no ROLE_NAMESPACE_SETS entry; declare " f"publish role {role!r} has no ROLE_NAMESPACE_SETS entry; declare "
"its namespace set (None for the full tree)." "its namespace set (None for the full tree)."
) )
server_args.resolve_once()
discarded = _CONTEXT.overrides_log() discarded = _CONTEXT.overrides_log()
_CONTEXT.set_server_args(server_args) _CONTEXT.set_server_args(server_args)
if discarded: if discarded:
+156 -21
View File
@@ -16,7 +16,9 @@
from __future__ import annotations from __future__ import annotations
import argparse import argparse
import copy
import dataclasses import dataclasses
import functools
import glob import glob
import importlib import importlib
import importlib.util import importlib.util
@@ -28,7 +30,6 @@ import random
import socket import socket
import tempfile import tempfile
import uuid import uuid
from functools import cached_property
from typing import Any, Callable, Dict, List, Literal, Optional, Union from typing import Any, Callable, Dict, List, Literal, Optional, Union
from sglang.kernels.ops.kv_canary.consts import RealKvHashMode from sglang.kernels.ops.kv_canary.consts import RealKvHashMode
@@ -3642,7 +3643,96 @@ class ServerArgs:
] = None ] = None
def __post_init__(self): def __post_init__(self):
self._run_resolution_pipeline() """Construction leaves the record at what the caller asked for.
Resolution is a separate act, entered through ``resolve_once``: the
launcher runs it once per engine, and every publishing process asks the
gate on the way in. A record that is only constructed -- a fixture, a
config being inspected, one being handed to a subprocess that will
resolve it itself -- stays raw.
"""
def resolve_once(self) -> None:
"""Run the resolution pipeline, unless this record has been through it.
Resolution is a deterministic function of the raw inputs -- two records
built from the same arguments declare the same things -- but the
handlers do not survive a second pass over their own output: DP
attention halves ``chunked_prefill_size`` again on every re-entry.
The publishing entry of every process calls this. In a child the record
arrived by pickle and brought its declarations along, so the child has
nothing left to derive and projects what the parent decided.
"""
if getattr(self, "_declarations_materialized", False):
return
if getattr(self, "_resolution_failed", False):
raise RuntimeError(
"resolution already failed on this ServerArgs; the handlers that "
"ran left their writes on the record, and a second pass would "
"read that partial output as fresh input. Build a new record "
"from the corrected arguments."
)
try:
self._run_resolution_pipeline()
except BaseException:
# The handlers that ran already wrote to the record, and they are
# not idempotent over their own output.
object.__setattr__(self, "_resolution_failed", True)
raise
# Set here too, because the dummy/absent-model path returns before the
# materialization that normally sets it: the gate is about whether the
# handlers ran, not how far they got.
self._declarations_materialized = True
def replace_resolved(self, source: str, **changes: Any) -> ServerArgs:
"""A copy of this record that stays resolved, and says what it changed.
`dataclasses.replace` builds a new instance, so the copy carries none of
what makes a record resolved: no raw snapshot, no declarations, no
materialization. The next publish therefore finds an unmaterialized
record and runs the pipeline over values it already decided -- DP
attention halves `chunked_prefill_size` a second time (8192 -> 4096 ->
2048) and the schedule conservativeness is scaled again (0.3 -> 0.09).
The Ray paths replace `dist_init_addr` on a resolved record, which is
how they hit it.
The change is appended to the stash rather than left on the field: the
projection reads the raw snapshot plus the declarations, so a field the
copy set on its own would publish the parent's raw value instead.
The carry is shallow. The containers are copied so the copy's own
declaration does not travel back into the parent, but everything inside
them -- the stash entries, the raw-input values, the memoized
`ModelConfig` -- is shared. That is fine for what this is for: a copy
that immediately crosses a process boundary (Ray actors, the gateway's
workers), where pickling severs the sharing. A caller that mutates the
copy's deep structure in-process mutates the parent's too.
"""
replacement = dataclasses.replace(self, **changes)
if not getattr(self, "_declarations_materialized", False):
# Not resolved yet: the copy goes through the gate itself.
return replacement
# Everything outside the fields, enumerated from the instance: the raw
# snapshot, the stash, and what resolution memoized -- including the
# `get_model_config()` cache, which a resolved copy can no longer fill
# (the read-only guard refuses the write).
field_names = {field.name for field in dataclasses.fields(self)}
for name, value in vars(self).items():
if name in field_names or name == "_declarations_materialized":
continue
if isinstance(value, (dict, list, set)):
value = copy.copy(value)
object.__setattr__(replacement, name, value)
stash = getattr(replacement, "_resolved_overrides", None)
if stash is None:
stash = []
object.__setattr__(replacement, "_resolved_overrides", stash)
if changes:
stash.append((source, dict(changes)))
object.__setattr__(replacement, "_declarations_materialized", True)
return replacement
def _declare(self, source: str, **fields: Any) -> None: def _declare(self, source: str, **fields: Any) -> None:
"""This record's handlers declaring their resolution writes. """This record's handlers declaring their resolution writes.
@@ -9522,8 +9612,17 @@ class ServerArgs:
# Lazy init to avoid circular import # Lazy init to avoid circular import
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
if hasattr(self, "model_config"): memo = getattr(self, "model_config", None)
return self.model_config if memo is not None:
# A configuration built before resolution describes the path the
# caller typed; the GGUF and ModelScope handlers declare a
# different `model_path`, and every later decision keyed on the
# architecture would read the wrong contents. Only a real
# `ModelConfig` is checked -- a fixture's stand-in stays untouched.
if not (
isinstance(memo, ModelConfig) and memo.model_path != self.model_path
):
return memo
self.model_config = ModelConfig.from_server_args(self) self.model_config = ModelConfig.from_server_args(self)
if self.model_config.is_hybrid_swa: if self.model_config.is_hybrid_swa:
logger.info( logger.info(
@@ -9556,9 +9655,9 @@ class ServerArgs:
# get_context().override(source, ...); a value one runner or worker # get_context().override(source, ...); a value one runner or worker
# owns travels as a constructor argument to it. # owns travels as a constructor argument to it.
if ( if (
not name.startswith("_") getattr(self, "_declarations_materialized", False)
and getattr(self, "_declarations_materialized", False)
and not getattr(self, "_internal_write", False) and not getattr(self, "_internal_write", False)
and (not name.startswith("_") or name in _underscore_field_names())
): ):
raise AttributeError( raise AttributeError(
f"server_args.{name} assigned after resolution; server_args is " f"server_args.{name} assigned after resolution; server_args is "
@@ -9599,30 +9698,47 @@ class ServerArgs:
def enable_mamba_extra_buffer_lazy(self) -> bool: def enable_mamba_extra_buffer_lazy(self) -> bool:
return mamba_extra_buffer_lazy_of(self) return mamba_extra_buffer_lazy_of(self)
@cached_property @property
def max_speculative_num_draft_tokens(self) -> Optional[int]: def max_speculative_num_draft_tokens(self) -> Optional[int]:
"""Return the maximum draft-token count speculative decoding may use.""" """Return the maximum draft-token count speculative decoding may use.
Memoized only once the record is resolved: an answer computed off a raw
record describes inputs resolution is about to rewrite (auto speculative
sizing fills `speculative_num_draft_tokens` in), and a cache filled that
early would keep answering with it.
"""
memo = self.__dict__.get("_max_speculative_num_draft_tokens")
if memo is not None:
return memo
if self.speculative_num_draft_tokens is None: if self.speculative_num_draft_tokens is None:
return None result = None
if not self.speculative_adaptive: elif not self.speculative_adaptive:
return self.speculative_num_draft_tokens result = self.speculative_num_draft_tokens
else:
from sglang.srt.speculative.adaptive_spec_params import (
resolve_candidate_steps_from_config,
)
from sglang.srt.speculative.adaptive_spec_params import ( candidate_steps = resolve_candidate_steps_from_config(
resolve_candidate_steps_from_config, cfg_path=self.speculative_adaptive_config,
) )
# TODO: adaptive spec currently requires topk=1, so each runtime
candidate_steps = resolve_candidate_steps_from_config( # state needs steps + 1 draft-token slots. Revisit this if topk>1
cfg_path=self.speculative_adaptive_config, # is supported.
) result = max(candidate_steps) + 1
# TODO: adaptive spec currently requires topk=1, so each runtime state if getattr(self, "_declarations_materialized", False):
# needs steps + 1 draft-token slots. Revisit this if topk>1 is supported. object.__setattr__(self, "_max_speculative_num_draft_tokens", result)
return max(candidate_steps) + 1 return result
@property @property
def mamba_cache_chunk_size(self) -> int: def mamba_cache_chunk_size(self) -> int:
# For mamba cache with extra buffer, the chunk size is the max of FLA_CHUNK_SIZE # For mamba cache with extra buffer, the chunk size is the max of FLA_CHUNK_SIZE
# (or mamba_chunk_size if it is defined in the model's config) and page_size. # (or mamba_chunk_size if it is defined in the model's config) and page_size.
# It is used to determine the caching point in a sequence during prefill. # It is used to determine the caching point in a sequence during prefill.
# A pre-seeded `_mamba_cache_chunk_size` (fixtures supply one so a dummy
# model never loads an HF config) is honored as-is; otherwise the memo
# is only kept once the record is resolved, because `page_size` below
# is resolution-written.
if not hasattr(self, "_mamba_cache_chunk_size"): if not hasattr(self, "_mamba_cache_chunk_size"):
try: try:
@@ -9639,6 +9755,8 @@ class ServerArgs:
assert ( assert (
max(chunk_size, page_size) % min(chunk_size, page_size) == 0 max(chunk_size, page_size) % min(chunk_size, page_size) == 0
), f"For SSM models, either chunk_size or page_size must be divisible by the other, got {chunk_size=}, {page_size=}" ), f"For SSM models, either chunk_size or page_size must be divisible by the other, got {chunk_size=}, {page_size=}"
if not getattr(self, "_declarations_materialized", False):
return max(chunk_size, page_size)
self._mamba_cache_chunk_size = max(chunk_size, page_size) self._mamba_cache_chunk_size = max(chunk_size, page_size)
return self._mamba_cache_chunk_size return self._mamba_cache_chunk_size
@@ -10331,6 +10449,23 @@ def m3_fp8_attn_gemm_enabled(args) -> bool:
# reference. Do not add new call-sites — the counts are ratcheted # reference. Do not add new call-sites — the counts are ratcheted
# (decrease-only) by test/registered/unit/test_legacy_global_ratchet.py. # (decrease-only) by test/registered/unit/test_legacy_global_ratchet.py.
# Imports are in-function so the two modules stay cycle-free at import time. # Imports are in-function so the two modules stay cycle-free at import time.
@functools.lru_cache(maxsize=1)
def _underscore_field_names() -> frozenset:
"""Real dataclass fields whose names start with an underscore.
The read-only guard exempts underscore names because they are the record's
own bookkeeping (the stash, the flags, the memoized model config). A *field*
that happens to start with an underscore is still resolved configuration --
`_speculative_draft_quantization_explicitly_set` is one -- and exempting it
by spelling would leave exactly one leaf writable on a read-only record.
"""
return frozenset(
field.name
for field in dataclasses.fields(ServerArgs)
if field.name.startswith("_")
)
def set_global_server_args_for_scheduler(server_args: ServerArgs): def set_global_server_args_for_scheduler(server_args: ServerArgs):
"""Legacy publish shim (role=scheduler) — prefer """Legacy publish shim (role=scheduler) — prefer
``runtime_context.publish(server_args, role=...)`` in new code.""" ``runtime_context.publish(server_args, role=...)`` in new code."""
+1
View File
@@ -315,5 +315,6 @@ if __name__ == "__main__":
parser.add_argument("--is-multimodal", action="store_true", default=False) parser.add_argument("--is-multimodal", action="store_true", default=False)
args = parser.parse_args() args = parser.parse_args()
server_args: ServerArgs = ServerArgs.from_cli_args(args) server_args: ServerArgs = ServerArgs.from_cli_args(args)
server_args.resolve_once()
main(args, server_args) main(args, server_args)
@@ -87,10 +87,23 @@ def launch_server_process(
server_args: ServerArgs, worker_port: int, dp_id: int server_args: ServerArgs, worker_port: int, dp_id: int
) -> mp.Process: ) -> mp.Process:
"""Launch a single server process with the given args and port.""" """Launch a single server process with the given args and port."""
server_args = copy.deepcopy(server_args) # This binding is installed against a released sglang, so it cannot call
server_args.port = worker_port # into helpers newer than that wheel. Copy first, then write through the
server_args.base_gpu_id = dp_id * server_args.tp_size # sanctioned channel if the record is resolved (a resolved record refuses
server_args.dp_size = 1 # plain assignment), else assign.
worker_args = copy.deepcopy(server_args)
changes = {
"port": worker_port,
"base_gpu_id": dp_id * server_args.tp_size,
"dp_size": 1,
}
late = getattr(worker_args, "_late_resolution", None)
if late is not None and getattr(worker_args, "_declarations_materialized", False):
late("sglang_router.launch_server_process", **changes)
else:
for field, value in changes.items():
setattr(worker_args, field, value)
server_args = worker_args
proc = mp.Process(target=run_server, args=(server_args, dp_id)) proc = mp.Process(target=run_server, args=(server_args, dp_id))
proc.start() proc.start()
@@ -170,6 +183,9 @@ def main():
args = parser.parse_args() args = parser.parse_args()
server_args = ServerArgs.from_cli_args(args) server_args = ServerArgs.from_cli_args(args)
# Older released wheels resolve in the constructor and have no gate.
if hasattr(server_args, "resolve_once"):
server_args.resolve_once()
router_args = RouterArgs.from_cli_args(args, use_router_prefix=True) router_args = RouterArgs.from_cli_args(args, use_router_prefix=True)
# Find available ports for workers # Find available ports for workers
@@ -184,6 +184,7 @@ def test_epd_rejection_reads_the_resolved_transfer_backend():
mamba_radix_cache_strategy="no_buffer", mamba_radix_cache_strategy="no_buffer",
disable_overlap_schedule=True, disable_overlap_schedule=True,
) )
resolved.resolve_once()
finally: finally:
os.environ.clear() os.environ.clear()
os.environ.update(environ_before) os.environ.update(environ_before)
@@ -179,13 +179,15 @@ assert not any(name == "mlx" or name.startswith("mlx.") for name in sys.modules)
with self.assertRaisesRegex(RuntimeError, "stable Torch 2.13.x"): with self.assertRaisesRegex(RuntimeError, "stable Torch 2.13.x"):
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
ServerArgs(model_path="dummy") # The check runs in the pipeline, and the pipeline runs
# at the gate -- still ahead of the dummy short circuit.
ServerArgs(model_path="dummy").resolve_once()
runtime.use_mlx.cache_clear() runtime.use_mlx.cache_clear()
runtime._validate_runtime.cache_clear() runtime._validate_runtime.cache_clear()
with mock.patch.object(mx, "__version__", "0.31.0"): with mock.patch.object(mx, "__version__", "0.31.0"):
with self.assertRaisesRegex(RuntimeError, "MLX >= 0.32.0"): with self.assertRaisesRegex(RuntimeError, "MLX >= 0.32.0"):
ServerArgs(model_path="dummy") ServerArgs(model_path="dummy").resolve_once()
finally: finally:
runtime.use_mlx.cache_clear() runtime.use_mlx.cache_clear()
runtime._validate_runtime.cache_clear() runtime._validate_runtime.cache_clear()
@@ -107,10 +107,11 @@ class TestTokenizerConfigUpdates(CustomTestCase):
def __deepcopy__(self, memo): def __deepcopy__(self, memo):
raise RuntimeError("refuses to be copied") raise RuntimeError("refuses to be copied")
manager = _manager(self) # Through the constructor: the field is raw input, and a resolved
# record refuses to be written.
manager = _manager(self, custom_sigquit_handler=Hostile())
manager.model_path = "dummy" manager.model_path = "dummy"
manager.served_model_name = "dummy" manager.served_model_name = "dummy"
manager.server_args.custom_sigquit_handler = Hostile()
self.assertIsNone(manager._dump_config_snapshot()) self.assertIsNone(manager._dump_config_snapshot())
@@ -118,11 +119,10 @@ class TestTokenizerConfigUpdates(CustomTestCase):
import dataclasses import dataclasses
import pickle import pickle
manager = _manager(self) # What --custom-sigquit-handler leaves on a real ServerArgs.
manager = _manager(self, custom_sigquit_handler=lambda *_: None)
manager.model_path = "dummy" manager.model_path = "dummy"
manager.served_model_name = "dummy" manager.served_model_name = "dummy"
# What --custom-sigquit-handler leaves on a real ServerArgs.
manager.server_args.custom_sigquit_handler = lambda *_: None
payload = { payload = {
"server_args": manager.server_args, "server_args": manager.server_args,
@@ -450,6 +450,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
node_rank=0, node_rank=0,
tokenizer_worker_num=1, tokenizer_worker_num=1,
check_server_args=MagicMock(), check_server_args=MagicMock(),
resolve_once=MagicMock(),
) )
scheduler_init_result = SimpleNamespace( scheduler_init_result = SimpleNamespace(
all_child_pids=[], all_child_pids=[],
@@ -259,7 +259,9 @@ class TestResolutionDeclarations(CustomTestCase):
json.dump(_MINI_CONFIG, handle) json.dump(_MINI_CONFIG, handle)
fields = {"random_seed": 42} fields = {"random_seed": 42}
fields.update(extra) fields.update(extra)
return ServerArgs(model_path=path, device="cuda", **fields) server_args = ServerArgs(model_path=path, device="cuda", **fields)
server_args.resolve_once()
return server_args
def test_converted_fields_are_not_assigned_bare(self): def test_converted_fields_are_not_assigned_bare(self):
bare = _bare_assignments() bare = _bare_assignments()
@@ -27,12 +27,15 @@ import copy
import dataclasses import dataclasses
import json import json
import os import os
import pathlib
import shutil import shutil
import tempfile import tempfile
import unittest import unittest
import unittest.mock
import torch import torch
import sglang
from sglang.srt.environ import EnvField, envs from sglang.srt.environ import EnvField, envs
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import is_cuda from sglang.srt.utils import is_cuda
@@ -187,13 +190,13 @@ _STICKY_ACROSS_RESOLUTIONS = frozenset({"mm_feature_transport"})
_NOT_COMPARABLE = frozenset({"random_seed"}) _NOT_COMPARABLE = frozenset({"random_seed"})
class TestResolutionIsReproducible(CustomTestCase): class _RestoresProcessState:
def _config_dir(self, config: dict = None) -> str: """Resolution leaves process state behind, so a case that resolves has to
config_dir = tempfile.mkdtemp(prefix="resolution_repro_") put it back. `_handle_multimodal_feature_transport` sets
self.addCleanup(shutil.rmtree, config_dir, ignore_errors=True) `SGLANG_USE_CUDA_IPC_TRANSPORT` and the same handler reads `is_set()` on the
with open(os.path.join(config_dir, "config.json"), "w") as handle: way in, so one resolution is visible to the next one in this process -- and
json.dump(config or _MINI_CONFIG, handle) `TestMultimodalFeatureTransport` is the case that notices.
return config_dir """
def _process_state(self): def _process_state(self):
"""What a resolution may leave behind: the environment and the """What a resolution may leave behind: the environment and the
@@ -236,13 +239,24 @@ class TestResolutionIsReproducible(CustomTestCase):
# to catch, turned into a pass. # to catch, turned into a pass.
unittest.TestCase._callTestMethod(self, method) unittest.TestCase._callTestMethod(self, method)
class TestResolutionIsReproducible(_RestoresProcessState, CustomTestCase):
def _config_dir(self, config: dict = None) -> str:
config_dir = tempfile.mkdtemp(prefix="resolution_repro_")
self.addCleanup(shutil.rmtree, config_dir, ignore_errors=True)
with open(os.path.join(config_dir, "config.json"), "w") as handle:
json.dump(config or _MINI_CONFIG, handle)
return config_dir
def _resolved(self, model_path: str, **kwargs) -> ServerArgs: def _resolved(self, model_path: str, **kwargs) -> ServerArgs:
# device="cuda" keeps the golden path host-independent: an # device="cuda" keeps the golden path host-independent: an
# accelerator-less runner resolves only the base platform, where # accelerator-less runner resolves only the base platform, where
# get_device() raises. # get_device() raises.
kwargs.setdefault("device", "cuda") kwargs.setdefault("device", "cuda")
kwargs.setdefault("random_seed", 42) kwargs.setdefault("random_seed", 42)
return ServerArgs(model_path=model_path, **kwargs) server_args = ServerArgs(model_path=model_path, **kwargs)
server_args.resolve_once()
return server_args
def _comparable(self, server_args: ServerArgs) -> dict: def _comparable(self, server_args: ServerArgs) -> dict:
"""The dataclass fields, and only those. """The dataclass fields, and only those.
@@ -385,6 +399,62 @@ class TestResolutionIsReproducible(CustomTestCase):
) )
self.assertEqual(self._comparable(first), snapshot) self.assertEqual(self._comparable(first), snapshot)
def test_the_gate_refuses_a_second_resolution(self):
"""A record that has been resolved is left exactly as it was.
Every publishing process calls the gate, and in a child the record
arrived already resolved -- so this is the property that keeps the
child agreeing with the parent. The handlers are not written to survive
a second pass over their own output (the DP-attention step derives the
chunked prefill size *from* the chunked prefill size), which is why the
gate refuses rather than re-deriving.
"""
for label, config, kwargs in _SHAPES:
with self.subTest(shape=label):
self._restore_process_state(self._pristine_state)
model_path = self._config_dir(config)
resolved = self._resolved(model_path, **kwargs)
snapshot = self._comparable(resolved)
declarations = list(getattr(resolved, "_resolved_overrides", []))
resolved.resolve_once()
self.assertEqual(self._comparable(resolved), snapshot)
self.assertEqual(
list(getattr(resolved, "_resolved_overrides", [])), declarations
)
def test_the_gate_closes_on_the_dummy_path_too(self):
"""The dummy model leaves the pipeline early, and the gate still shuts.
That exit is above the materialization the gate reads, so a dummy
record answered "not resolved yet" forever and every publish of one ran
the handlers again. Nothing about the early exit makes a second pass
safe -- the handlers above it declare and apply like any other -- and
the four that do run happening to be idempotent today is what the gate
exists to stop depending on. So this counts entries rather than
comparing values: the values agree either way.
"""
self._restore_process_state(self._pristine_state)
record = ServerArgs(model_path="dummy")
record.resolve_once()
entries = []
original = ServerArgs._run_resolution_pipeline
def counted(self):
entries.append(1)
return original(self)
with unittest.mock.patch.object(
ServerArgs, "_run_resolution_pipeline", counted
):
record.resolve_once()
self.assertEqual(
entries,
[],
"a resolved dummy record entered the pipeline again, so every "
"publish of one re-runs the handlers",
)
def test_the_declaration_provenance_is_reproducible(self): def test_the_declaration_provenance_is_reproducible(self):
model_path = self._config_dir() model_path = self._config_dir()
first = self._resolved(model_path) first = self._resolved(model_path)
@@ -402,17 +472,464 @@ class TestResolutionIsReproducible(CustomTestCase):
self.assertEqual(getattr(first, "_resolved_overrides", None), first_provenance) self.assertEqual(getattr(first, "_resolved_overrides", None), first_provenance)
class TestTheResolutionSeamHasOneCaller(CustomTestCase): class TestProgramsResolveBeforeReadingResolution(CustomTestCase):
"""The pipeline is entered from exactly one place. """A program that builds its own record resolves it before reading what
resolution decides.
Step 12 moves the call from ``__post_init__`` to ``publish`` so the record Construction is inert, so a program that builds a record and then reads a
stays raw; that is a one-line move only while the seam has a single caller. resolution-written field reads the CLI default. Two of these shipped past
A second entry point would also mean resolution could run twice on one the earlier censuses because those are rooted at the `sglang` package: the
instance, which the strict ``__setattr__`` guard turns into an model gateway's launcher sized its worker plan from a raw `dp_size`
``AttributeError`` rather than a silent re-resolve. (`--dwdp-size 4` launched one server instead of four) and a speculative
benchmark forwarded `--mem-fraction-static None` to the server it spawns.
So the universe here is the *repository*, not the package.
""" """
def test_only_post_init_runs_the_pipeline(self): # Entries that hand the record on instead of reading it. Reason required.
_EXEMPT: dict = {}
def _repo_root(self):
# <repo>/python/sglang/__init__.py -> <repo>
root = pathlib.Path(next(iter(sglang.__path__))).resolve().parents[1]
if root.name == "python":
root = root.parent
return root
def _written_fields(self):
"""Fields resolution declares, read out of the pipeline's own source.
Deliberately local: the chain ratchet has a wider derivation (it also
walks the model-override registries), but it arrives later in this
series, and a check that imports it would fail at this PR's boundary.
Coarser is fine here -- what this needs is the fields the entries below
actually read -- and the floor keeps it from drifting narrower.
"""
import ast
import dataclasses as _dataclasses
from sglang.srt.server_args import ServerArgs as _ServerArgs
srt = pathlib.Path(next(iter(sglang.__path__))).resolve() / "srt"
declarers = {"_declare", "declare_resolution", "declare_late_resolution"}
fields = set()
field_names = {field.name for field in _dataclasses.fields(_ServerArgs)}
for name in ("server_args.py", "arg_groups/overrides.py"):
tree = ast.parse((srt / name).read_text(encoding="utf-8-sig"))
for node in ast.walk(tree):
# Registry data: provider dict keys are field names as
# *data*, invisible to the keyword scan below. Filtered
# against the real field set.
if isinstance(node, ast.Dict):
fields |= {
key.value
for key in node.keys
if isinstance(key, ast.Constant)
and isinstance(key.value, str)
and key.value in field_names
}
if not isinstance(node, ast.Call):
continue
func = node.func
called = (
func.attr
if isinstance(func, ast.Attribute)
else getattr(func, "id", "")
)
if called in declarers or called == "update":
fields |= {
kw.arg
for kw in node.keywords
if kw.arg and (called != "update" or kw.arg in field_names)
}
return fields
def _candidates(self, root):
"""Source files that build a record, with the names they bind it to."""
import ast
skip = {".git", "build", "dist", "node_modules", ".venv", "target"}
found = {}
for path in sorted(root.rglob("*.py")):
parts = set(path.relative_to(root).parts)
if parts & skip:
continue
rel = path.relative_to(root).as_posix()
# Tests build raw records on purpose.
if rel.startswith("test/") or "/test/" in rel or "/tests/" in rel:
continue
try:
tree = ast.parse(path.read_text(encoding="utf-8-sig"))
except (SyntaxError, UnicodeDecodeError):
continue
# Which local names are *the srt record*, by import source: the
# diffusion runtime has a same-spelled `ServerArgs` with no
# resolution, so the spelling alone is not enough.
record_classes, record_helpers = set(), set()
for node in ast.walk(tree):
if not isinstance(node, ast.ImportFrom):
continue
for alias in node.names:
bound = alias.asname or alias.name
if node.module == "sglang" and alias.name == "ServerArgs":
record_classes.add(bound)
if node.module == "sglang.srt.server_args":
if alias.name == "ServerArgs":
record_classes.add(bound)
if alias.name == "prepare_server_args":
record_helpers.add(bound)
names = set()
for node in ast.walk(tree):
if isinstance(node, ast.Assign):
targets = node.targets
elif isinstance(node, (ast.AnnAssign, ast.NamedExpr)):
# `x: ServerArgs = ...` is an AnnAssign, not an Assign.
targets = [node.target]
else:
continue
call = getattr(node, "value", None)
if not isinstance(call, ast.Call):
continue
func = call.func
# `prepare_server_args(argv)` is the CLI launcher's way.
builds = (
isinstance(func, ast.Name)
and func.id in (record_classes | record_helpers)
) or (
isinstance(func, ast.Attribute)
and func.attr == "from_cli_args"
and isinstance(func.value, ast.Name)
and func.value.id in record_classes
)
if builds:
names |= {t.id for t in targets if isinstance(t, ast.Name)}
if names:
found[rel] = (tree, names, path)
return found
def test_every_program_that_builds_a_record_resolves_it(self):
import ast
root = self._repo_root()
candidates = self._candidates(root)
self.assertGreater(
len(candidates),
10,
f"only {len(candidates)} files build a record under {root}; either "
"this is not a source checkout or the scan broke",
)
written = self._written_fields()
self.assertGreater(len(written), 50, "the written-field set collapsed")
# What the escaped entries actually read: a narrower derivation goes
# quiet on exactly those.
for field in ("dp_size", "mem_fraction_static"):
self.assertIn(field, written)
offenders = []
for rel, (tree, names, path) in sorted(candidates.items()):
source = path.read_text(encoding="utf-8-sig")
if "resolve_once(" in source or "publish(" in source:
continue
reads = sorted(
{
f"{node.attr}:{node.lineno}"
for node in ast.walk(tree)
if isinstance(node, ast.Attribute)
and isinstance(node.value, ast.Name)
and node.value.id in names
and node.attr in written
}
)
if reads and rel not in self._EXEMPT:
offenders.append(f"{rel} reads {', '.join(reads[:4])}")
self.assertEqual(
offenders,
[],
"a program builds its own record and reads what resolution decides "
"without resolving it, so it reads the CLI default:\n "
+ "\n ".join(offenders),
)
self.assertEqual(
sorted(set(self._EXEMPT) - set(candidates)),
[],
"an exemption names a file that no longer builds a record",
)
class TestForksResolveFirst(CustomTestCase):
"""A process that forks a child to run the record resolves it first.
The pipeline probes the device (the default attention backend reads the CUDA
capability), and a forked child cannot initialize CUDA once its parent has.
Construction used to resolve, so the probe always happened in whoever built
the record; now it happens at the gate, and the gate must not be reached for
the first time inside a fork.
"""
# Sites inside the launcher: `_launch_subprocesses` resolves at its top, so
# every fork below it already has a resolved record.
_AFTER_LAUNCHER_RESOLVE = {
"srt/entrypoints/engine.py",
"srt/managers/data_parallel_controller.py",
"srt/disaggregation/encoder/grpc_server.py",
"srt/disaggregation/encoder/runtime.py",
"srt/elastic_ep/expert_backup_manager.py",
}
def test_every_fork_of_a_record_has_a_resolved_one(self):
import ast
package_root = pathlib.Path(next(iter(sglang.__path__))).resolve()
offenders, examined = [], 0
for path in sorted(package_root.rglob("*.py")):
rel = path.relative_to(package_root).as_posix()
if rel.startswith("test/") or "/test/" in rel:
continue
# The diffusion runtime has its own record with no gate.
if rel.startswith("multimodal_gen/"):
continue
try:
source = path.read_text(encoding="utf-8-sig")
tree = ast.parse(source)
except (SyntaxError, UnicodeDecodeError):
continue
for node in ast.walk(tree):
if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
continue
body = ast.get_source_segment(source, node) or ""
forks = [
call
for call in ast.walk(node)
if isinstance(call, ast.Call)
and (
(
isinstance(call.func, ast.Attribute)
and call.func.attr == "Process"
)
or (
isinstance(call.func, ast.Name)
and call.func.id == "Process"
)
)
and "server_args" in (ast.get_source_segment(source, call) or "")
]
if not forks:
continue
examined += 1
# `spawn` starts a fresh interpreter, so the child may probe.
if 'get_context("spawn")' in body or "'spawn'" in body:
continue
if "resolve_once(" in body or "publish(" in body:
continue
if rel in self._AFTER_LAUNCHER_RESOLVE:
continue
offenders.append(f"{rel}:{forks[0].lineno} {node.name}")
self.assertGreater(
examined, 5, f"only {examined} fork sites found; the scan broke"
)
self.assertEqual(
offenders,
[],
"these fork a child that will resolve the record, without resolving "
"it first -- the child cannot initialize CUDA if this process "
f"already has:\n " + "\n ".join(offenders),
)
class TestACopyStaysResolved(_RestoresProcessState, CustomTestCase):
"""A resolved record copied with `dataclasses.replace` loses what makes it
resolved, and the next publish resolves it a second time -- over values it
already decided. The Ray paths copy a resolved record to set
`dist_init_addr`, which is how they reach this.
"""
def _resolved(self):
config_dir = tempfile.mkdtemp(prefix="replace_resolved_")
self.addCleanup(shutil.rmtree, config_dir, ignore_errors=True)
with open(os.path.join(config_dir, "config.json"), "w") as handle:
json.dump(_MINI_CONFIG, handle)
# Two steps that are not repeatable on their own output.
server_args = ServerArgs(
model_path=config_dir,
device="cuda",
dp_size=2,
tp_size=2,
enable_dp_attention=True,
random_seed=42,
)
server_args.resolve_once()
return server_args
def test_a_bare_replace_would_resolve_a_second_time(self):
"""Why the helper exists. If this stops drifting, the pipeline became
idempotent and the helper's reason is gone -- read it again before
deleting either."""
parent = self._resolved()
bare = dataclasses.replace(parent, dist_init_addr="1.2.3.4:5000")
self.assertFalse(
getattr(bare, "_declarations_materialized", False),
"a bare replace carried the flag; then this test proves nothing",
)
bare.resolve_once()
self.assertEqual(
(bare.chunked_prefill_size, round(bare.schedule_conservativeness, 4)),
(
parent.chunked_prefill_size // 2,
round(parent.schedule_conservativeness * 0.3, 4),
),
"the second pass no longer drifts; this is the drift the copy "
"helper exists to avoid",
)
def test_replace_resolved_keeps_the_parents_resolution(self):
parent = self._resolved()
copy_ = parent.replace_resolved("ray.test", dist_init_addr="1.2.3.4:5000")
self.assertTrue(getattr(copy_, "_declarations_materialized", False))
drifted = {
field.name: (getattr(parent, field.name), getattr(copy_, field.name))
for field in dataclasses.fields(parent)
if field.name != "dist_init_addr"
and getattr(parent, field.name) != getattr(copy_, field.name)
}
self.assertEqual(
drifted,
{},
f"the copy differs from its parent beyond the change: {drifted}",
)
self.assertEqual(copy_.dist_init_addr, "1.2.3.4:5000")
def test_the_copy_carries_what_resolution_left_on_the_record(self):
"""Not just the stash and the flag.
`get_model_config()` memoizes on the record, and that cache is filled
during resolution. A copy that is marked resolved but arrives without it
cannot fill it -- the read-only guard refuses the cache write -- so the
first `get_model_config()` raises. That is what killed the Ray
schedulers, and it is why the carry is enumerated from the instance
rather than from a list of names.
"""
parent = self._resolved()
copy_ = parent.replace_resolved("ray.test", dist_init_addr="1.2.3.4:5000")
fields = {field.name for field in dataclasses.fields(parent)}
missing = sorted(
name
for name in vars(parent)
if name not in fields and name not in vars(copy_)
)
self.assertEqual(
missing,
[],
f"the copy did not carry what resolution left on the record: {missing}",
)
self.assertIsNotNone(copy_.get_model_config())
# Containers are copied, so the copy's declaration stays with it.
self.assertEqual(
len(parent._resolved_overrides) + 1, len(copy_._resolved_overrides)
)
def test_the_change_reaches_the_bags(self):
"""The projection reads the raw snapshot plus the declarations, so a
change the copy only wrote to the field would publish the parent's raw
value."""
from sglang.srt.runtime_context import (
get_parallel,
get_schedule,
publish,
reset_context,
)
parent = self._resolved()
copy_ = parent.replace_resolved("ray.test", dist_init_addr="1.2.3.4:5000")
self.addCleanup(reset_context)
reset_context()
publish(copy_, role="scheduler")
self.assertEqual(get_parallel().dist_init_addr, "1.2.3.4:5000")
self.assertEqual(
get_schedule().chunked_prefill_size,
parent.chunked_prefill_size,
"publishing the copy re-ran resolution; the bag disagrees with the "
"record the parent resolved",
)
def test_no_bare_replace_of_a_record_outside_the_helper(self):
"""`dataclasses.replace` on a record is the helper's job now.
Derived, not listed: any `dataclasses.replace` whose first argument is
named for a record. The helper's own call is the positive control -- if
the scan stops seeing it, the scan broke rather than the tree.
"""
import ast
# The repository, not the package: the gateway is outside `sglang/`.
package_root = pathlib.Path(next(iter(sglang.__path__))).resolve().parents[1]
if package_root.name == "python":
package_root = package_root.parent
helper = "python/sglang/srt/server_args.py"
bare, inside_helper = [], 0
def replaces_a_record(node, record_names):
if not (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Attribute)
and node.func.attr == "replace"
and isinstance(node.func.value, ast.Name)
and node.func.value.id == "dataclasses"
and node.args
):
return False
first = node.args[0]
name = (
first.id if isinstance(first, ast.Name) else getattr(first, "attr", "")
)
return name in record_names or "server_args" in name
for path in sorted(package_root.rglob("*.py")):
try:
tree = ast.parse(path.read_text(encoding="utf-8-sig"))
except SyntaxError:
continue
rel = path.relative_to(package_root).as_posix()
# `self` is a record only inside the record's own class body.
in_record_class = [
node
for node in ast.walk(tree)
if isinstance(node, ast.ClassDef) and node.name == "ServerArgs"
]
for scope, record_names in [(tree, set())] + [
(klass, {"self"}) for klass in in_record_class
]:
for node in ast.walk(scope):
if not replaces_a_record(node, record_names):
continue
if rel == helper and record_names:
inside_helper += 1
elif not record_names:
bare.append(f"{rel}:{node.lineno}")
self.assertEqual(
inside_helper,
1,
"the scan no longer finds `replace_resolved`'s own call; it broke",
)
self.assertEqual(
bare,
[],
"a record is copied with a bare `dataclasses.replace`, so the copy "
"loses the parent's resolution and the next publish resolves it "
"again: " + ", ".join(bare),
)
class TestTheResolutionSeamHasOneCaller(CustomTestCase):
"""The pipeline is entered from exactly one place, and that place decides
whether it runs at all.
``resolve_once`` is the gate: the handlers are not written to survive a
second pass over their own output, so a record must go through the pipeline
at most once. Keeping the pipeline itself down to a single caller is what
makes that gate impossible to bypass -- and what keeps the remaining move
(construction time to publish time) a matter of who calls the gate.
"""
def test_only_the_gate_runs_the_pipeline(self):
import ast import ast
from pathlib import Path from pathlib import Path
@@ -449,12 +966,93 @@ class TestTheResolutionSeamHasOneCaller(CustomTestCase):
# __post_init__, or another class growing a same-named __post_init__ # __post_init__, or another class growing a same-named __post_init__
# all show up here. # all show up here.
self.assertEqual( self.assertEqual(
[("srt/server_args.py", "ServerArgs.__post_init__")], [("srt/server_args.py", "ServerArgs.resolve_once")],
callers, callers,
"the resolution pipeline must be entered exactly once, from " "the resolution pipeline must be entered exactly once, from "
f"ServerArgs.__post_init__; found: {callers}", f"ServerArgs.resolve_once; found: {callers}",
) )
def test_the_gate_is_reached_from_the_launcher_and_from_publish(self):
"""Both entries go through the gate, so neither can resolve twice.
The launcher resolves the engine's record before reading any resolved
value from it; every publishing process asks the gate on the way in and
finds nothing left to do when the record arrived resolved.
"""
import ast
from pathlib import Path
import sglang
package_root = Path(next(iter(sglang.__path__)))
callers = []
for path in sorted(package_root.rglob("*.py")):
try:
tree = ast.parse(path.read_text())
except SyntaxError:
continue
for node in ast.walk(tree):
# `self.resolve_once()` at construction; publish looks the
# attribute up first, so it appears as a bare name call.
called = (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Attribute)
and node.func.attr == "resolve_once"
) or (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Name)
and node.func.id == "resolve_once"
)
if called:
callers.append(path.relative_to(package_root).as_posix())
machinery = {"srt/entrypoints/engine.py", "srt/runtime_context.py"}
self.assertEqual(
[
# Program entries: each builds a record from its own
# arguments and then reads effective configuration, or hands it
# to a fork that must not be the first to probe the device.
"benchmark/endpoint.py",
"benchmark/offline_throughput.py",
"benchmark/one_batch.py",
"benchmark/one_batch_server.py",
"compile_deep_gemm.py",
"lang/backend/runtime_endpoint.py",
"launch_server.py",
# The mechanism.
"srt/entrypoints/engine.py",
"srt/entrypoints/http_server_engine.py",
"srt/runtime_context.py",
],
sorted(set(callers)),
f"the resolution gate grew or lost a caller: {sorted(set(callers))}",
)
# The rule the list stands for: a caller that is not the mechanism
# resolves a record it built itself from argv. Anything else was handed
# one someone already resolved, or should publish.
for caller in sorted(set(callers) - machinery):
source = (package_root / caller).read_text()
# Either the module turned argv into the record -- the dataclass,
# the CLI classmethod, or the argv helper `launch_server.py` uses
# -- or it hands the record to a fork, which has to resolve first:
# the pipeline probes the device and a forked child cannot
# re-initialize CUDA. A worker handed a resolved record is neither.
builds_its_own = any(
spelling in source
for spelling in (
"ServerArgs(",
".from_cli_args(",
"prepare_server_args(",
)
) or ("Process(" in source and "server_args" in source)
# `assertTrue`, not `assertIn`: the container is a whole module.
self.assertTrue(
builds_its_own,
f"{caller} calls the resolution gate but does not build the "
"record it resolves; a record it was handed is already "
"resolved by whoever built it, and publish resolves what it "
"is handed",
)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -59,6 +59,8 @@ class TestPrepareServerArgs(CustomTestCase):
f"parser, got SystemExit({exc.code})" f"parser, got SystemExit({exc.code})"
) )
args.resolve_once()
self.assertTrue(args.enable_w4a4_mxfp4_megamoe) self.assertTrue(args.enable_w4a4_mxfp4_megamoe)
self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "1") self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "1")
self.assertEqual(os.environ["DG_USE_MXF4_KIND"], "1") self.assertEqual(os.environ["DG_USE_MXF4_KIND"], "1")
@@ -70,6 +72,9 @@ class TestPrepareServerArgs(CustomTestCase):
} }
with patch.dict(os.environ, deepgemm_env, clear=False): with patch.dict(os.environ, deepgemm_env, clear=False):
args = prepare_server_args(["--model-path", "dummy"]) args = prepare_server_args(["--model-path", "dummy"])
# Resolve, or the check that the environment stays untouched has
# nothing to be untouched by.
args.resolve_once()
self.assertFalse(args.enable_w4a4_mxfp4_megamoe) self.assertFalse(args.enable_w4a4_mxfp4_megamoe)
self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "0") self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "0")
@@ -77,12 +82,13 @@ class TestPrepareServerArgs(CustomTestCase):
def test_prefill_decode_interval(self): def test_prefill_decode_interval(self):
args = ServerArgs(model_path="dummy", prefill_decode_interval=16) args = ServerArgs(model_path="dummy", prefill_decode_interval=16)
args.resolve_once()
self.assertEqual(args.prefill_decode_interval, 16) self.assertEqual(args.prefill_decode_interval, 16)
with self.assertRaisesRegex( with self.assertRaisesRegex(
ValueError, "--prefill-decode-interval must be non-negative" ValueError, "--prefill-decode-interval must be non-negative"
): ):
ServerArgs(model_path="dummy", prefill_decode_interval=-1) ServerArgs(model_path="dummy", prefill_decode_interval=-1).resolve_once()
def test_dsv4_prefill_backend_cli_choices(self): def test_dsv4_prefill_backend_cli_choices(self):
parser = server_args_module.argparse.ArgumentParser() parser = server_args_module.argparse.ArgumentParser()
@@ -102,18 +108,23 @@ class TestPrepareServerArgs(CustomTestCase):
parser.parse_args(base_args + ["--dsv4-prefill-backend", "flashmla_kv"]) parser.parse_args(base_args + ["--dsv4-prefill-backend", "flashmla_kv"])
def test_return_hidden_states_mode_configuration(self): def test_return_hidden_states_mode_configuration(self):
disabled = ServerArgs(model_path="dummy") def _resolved(**kwargs):
server_args = ServerArgs(**kwargs)
server_args.resolve_once()
return server_args
disabled = _resolved(model_path="dummy")
self.assertFalse(disabled.enable_return_hidden_states) self.assertFalse(disabled.enable_return_hidden_states)
self.assertIsNone(disabled.return_hidden_states_mode) self.assertIsNone(disabled.return_hidden_states_mode)
last = ServerArgs( last = _resolved(
model_path="dummy", model_path="dummy",
return_hidden_states_mode="last", return_hidden_states_mode="last",
) )
self.assertTrue(last.enable_return_hidden_states) self.assertTrue(last.enable_return_hidden_states)
self.assertEqual(last.return_hidden_states_mode, "last") self.assertEqual(last.return_hidden_states_mode, "last")
legacy_full = ServerArgs( legacy_full = _resolved(
model_path="dummy", model_path="dummy",
enable_return_hidden_states=True, enable_return_hidden_states=True,
) )
@@ -128,14 +139,16 @@ class TestPrepareServerArgs(CustomTestCase):
"last", "last",
] ]
) )
parsed_last.resolve_once()
self.assertTrue(parsed_last.enable_return_hidden_states) self.assertTrue(parsed_last.enable_return_hidden_states)
self.assertEqual(parsed_last.return_hidden_states_mode, "last") self.assertEqual(parsed_last.return_hidden_states_mode, "last")
# The rejection is resolution's, not the constructor's.
with self.assertRaisesRegex( with self.assertRaisesRegex(
ValueError, ValueError,
"return_hidden_states_mode must be one of", "return_hidden_states_mode must be one of",
): ):
ServerArgs( _resolved(
model_path="dummy", model_path="dummy",
return_hidden_states_mode="lst", return_hidden_states_mode="lst",
) )
@@ -1300,9 +1313,11 @@ class TestSSLArgs(unittest.TestCase):
class TestHiCacheArgs(unittest.TestCase): class TestHiCacheArgs(unittest.TestCase):
def _make_args(self, **overrides) -> ServerArgs: def _make_args(self, **overrides) -> ServerArgs:
args = ServerArgs(model_path="dummy") # Not resolved: a dummy model path takes the pipeline's early return,
for key, value in overrides.items(): # so `_handle_hicache` would never run. Its one prerequisite (the
setattr(args, key, value) # host/device ratio default) is run by hand.
args = ServerArgs(model_path="dummy", **overrides)
args._handle_hicache_ratio_default()
return args return args
def _assert_hicache_fields( def _assert_hicache_fields(
@@ -1392,7 +1407,6 @@ class TestHiCacheArgs(unittest.TestCase):
attention_backend="fa3", attention_backend="fa3",
decode_attention_backend=None, decode_attention_backend=None,
) )
args._handle_hicache() args._handle_hicache()
self.assertEqual(args.hicache_io_backend, "kernel") self.assertEqual(args.hicache_io_backend, "kernel")
+9 -7
View File
@@ -259,15 +259,15 @@ class TestPublishInstallsSlot(_IsolatedPublish):
set_global_server_args_for_scheduler, set_global_server_args_for_scheduler,
) )
sa = ServerArgs(model_path="dummy") # __post_init__ early-returns sa = ServerArgs(model_path="dummy") # construction resolves nothing
# A dummy path short-circuits the pipeline, but the handlers ahead of self.assertFalse(hasattr(sa, "_resolved_overrides"))
# that point still declare; whatever they left in the stash is on the set_global_server_args_for_scheduler(sa)
# object by the time publish sees it. self.assertIs(get_server_args(), sa)
# Publishing is what resolved it; the handlers ahead of the dummy
# short-circuit still declare.
for source, declared in sa._resolved_overrides: for source, declared in sa._resolved_overrides:
for field, value in declared.items(): for field, value in declared.items():
self.assertEqual(getattr(sa, field), value, f"{source}: {field}") self.assertEqual(getattr(sa, field), value, f"{source}: {field}")
set_global_server_args_for_scheduler(sa)
self.assertIs(get_server_args(), sa)
class TestGoldenModelOverrides(_IsolatedPublish): class TestGoldenModelOverrides(_IsolatedPublish):
@@ -306,7 +306,9 @@ class TestGoldenModelOverrides(_IsolatedPublish):
self.addCleanup(shutil.rmtree, config_dir, ignore_errors=True) self.addCleanup(shutil.rmtree, config_dir, ignore_errors=True)
with open(os.path.join(config_dir, "config.json"), "w") as f: with open(os.path.join(config_dir, "config.json"), "w") as f:
json.dump(config, f) json.dump(config, f)
return ServerArgs(model_path=config_dir, **server_kwargs) server_args = ServerArgs(model_path=config_dir, **server_kwargs)
server_args.resolve_once()
return server_args
def _publish(self, server_args): def _publish(self, server_args):
from sglang.srt.server_args import ( from sglang.srt.server_args import (
@@ -105,6 +105,13 @@ _UNREAD_ENTRIES: dict = {
("multimodal_gen/test/unit/test_disagg_trace.py", "_srt_trace_server_args"): ( ("multimodal_gen/test/unit/test_disagg_trace.py", "_srt_trace_server_args"): (
"a trace fixture publishing its own context" "a trace fixture publishing its own context"
), ),
(
"multimodal_gen/runtime/managers/gpu_worker.py",
"init_device_and_model",
): (
"a worker installing a placeholder when its process has nothing "
"published; it reads its own config, not the srt bags"
),
} }
# `publish` itself and its named wrappers live here; a call inside them is the # `publish` itself and its named wrappers live here; a call inside them is the
+7 -4
View File
@@ -217,8 +217,8 @@ class TestServerArgsOwnership(_IsolatedServerArgs):
"""V2b: the context owns the slot; the legacy getters are identity shims.""" """V2b: the context owns the slot; the legacy getters are identity shims."""
def test_legacy_setter_publishes_into_context(self): def test_legacy_setter_publishes_into_context(self):
# Identity (not equality) is the contract; publish accepts any object. # Identity, not equality: the slot holds the very object published.
sentinel = object() sentinel = ServerArgs(model_path="dummy")
server_args_module.set_global_server_args_for_scheduler(sentinel) server_args_module.set_global_server_args_for_scheduler(sentinel)
self.assertIs(server_args_module.get_global_server_args(), sentinel) self.assertIs(server_args_module.get_global_server_args(), sentinel)
self.assertIs(get_server_args(), sentinel) self.assertIs(get_server_args(), sentinel)
@@ -240,13 +240,16 @@ class TestServerArgsOwnership(_IsolatedServerArgs):
self.assertEqual(str(cm.exception), "Global server args is not set yet!") self.assertEqual(str(cm.exception), "Global server args is not set yet!")
def test_republish_overwrite_allowed(self): def test_republish_overwrite_allowed(self):
first, second = object(), object() first = ServerArgs(model_path="dummy")
second = ServerArgs(model_path="dummy")
server_args_module.set_global_server_args_for_scheduler(first) server_args_module.set_global_server_args_for_scheduler(first)
server_args_module.set_global_server_args_for_scheduler(second) server_args_module.set_global_server_args_for_scheduler(second)
self.assertIs(get_server_args(), second) self.assertIs(get_server_args(), second)
def test_reset_context_clears_owned_store(self): def test_reset_context_clears_owned_store(self):
server_args_module.set_global_server_args_for_scheduler(object()) server_args_module.set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy")
)
reset_context() reset_context()
with self.assertRaises(ValueError): with self.assertRaises(ValueError):
get_server_args() get_server_args()
@@ -74,7 +74,8 @@ class TestConfigBags(CustomTestCase):
def _publish(self): def _publish(self):
sa = ServerArgs(model_path="dummy") sa = ServerArgs(model_path="dummy")
rc.get_context().set_server_args(sa) # Through publish, so the record is resolved the way a process resolves it.
rc.publish(sa, role="test")
return sa return sa
def test_fail_closed_before_publish(self): def test_fail_closed_before_publish(self):
@@ -201,7 +202,14 @@ class TestConfigBags(CustomTestCase):
self.addCleanup(restore_process_state) self.addCleanup(restore_process_state)
def resolve(): def resolve():
return ServerArgs(model_path=config_dir, device="cuda", random_seed=42) server_args = ServerArgs(
model_path=config_dir, device="cuda", random_seed=42
)
# The reference has to be *resolved*, not merely constructed:
# construction is inert, and the point of the sibling is to be an
# independent run of the pipeline over the same raw input.
server_args.resolve_once()
return server_args
sa = resolve() sa = resolve()
rc.publish(sa, role="scheduler") rc.publish(sa, role="scheduler")
@@ -287,7 +295,7 @@ class TestRoleNamespaceEnforcement(CustomTestCase):
rc.get_mm() rc.get_mm()
# A direct set_server_args install is roleless; enforcement only # A direct set_server_args install is roleless; enforcement only
# keys off a recorded publish role. # keys off a recorded publish role.
rc.get_context().set_server_args(ServerArgs(model_path="dummy")) rc.publish(ServerArgs(model_path="dummy"), role="test")
rc.get_exec() rc.get_exec()
def test_off_mode_bag_read_traces_under_torch_compile(self): def test_off_mode_bag_read_traces_under_torch_compile(self):
@@ -25,7 +25,8 @@ class TestContextOverride(CustomTestCase):
def _publish(self): def _publish(self):
sa = ServerArgs(model_path="dummy") sa = ServerArgs(model_path="dummy")
rc.get_context().set_server_args(sa) # Through publish, so the record is resolved the way a process resolves it.
rc.publish(sa, role="test")
return sa return sa
def test_override_writes_bag_not_server_args(self): def test_override_writes_bag_not_server_args(self):
@@ -24,7 +24,11 @@ class TestServerArgsAnnotatedCli(CustomTestCase):
def _parse(self, args_list): def _parse(self, args_list):
args = self.parser.parse_args(["--model", "dummy"] + args_list) args = self.parser.parse_args(["--model", "dummy"] + args_list)
return ServerArgs.from_cli_args(args) server_args = ServerArgs.from_cli_args(args)
# Parsing hands back the raw record; the cases below read values that
# resolution normalises, so resolve here the way a launcher would.
server_args.resolve_once()
return server_args
def test_aliases_and_dest(self): def test_aliases_and_dest(self):
"""Field name drives dest; long forms and short aliases both work.""" """Field name drives dest; long forms and short aliases both work."""
@@ -118,6 +118,12 @@ _LATE_RESOLUTION_DYNAMIC_SITES = {
"parser/template_detection.py": frozenset({"reasoning_parser", "tool_call_parser"}), "parser/template_detection.py": frozenset({"reasoning_parser", "tool_call_parser"}),
} }
# `get_context().override(...)` declares through the same seam, but the fields
# are the caller's -- a test names them one call at a time. There is no static
# set to collect, and nothing resolution decides: whatever a caller overrides
# there is exposure only through that caller's own reads.
_CALLER_SUPPLIED_LATE_SITES = frozenset({"runtime_context.py"})
# Resolution also branches on ambient environment; those shapes are explicit # Resolution also branches on ambient environment; those shapes are explicit
# entries so the written set is the same on every host. `SGLANG_IS_IN_CI` # entries so the written set is the same on every host. `SGLANG_IS_IN_CI`
# makes resolution fill `soft_watchdog_timeout`. # makes resolution fill `soft_watchdog_timeout`.
@@ -439,6 +445,7 @@ class TestSuppliedInstanceExposure(CustomTestCase):
resolved = ServerArgs( resolved = ServerArgs(
model_path=model_path, device="cuda", random_seed=42, **extra model_path=model_path, device="cuda", random_seed=42, **extra
) )
resolved.resolve_once()
except Exception as exc: except Exception as exc:
self.fail( self.fail(
f"the matrix entry {extra} (env={env}) did not resolve in " f"the matrix entry {extra} (env={env}) did not resolve in "
@@ -717,6 +724,8 @@ class TestSuppliedInstanceExposure(CustomTestCase):
if kw.arg and kw.arg != "source": if kw.arg and kw.arg != "source":
written.add(kw.arg) written.add(kw.arg)
elif kw.arg is None: elif kw.arg is None:
if rel in _CALLER_SUPPLIED_LATE_SITES:
continue
dynamic = _LATE_RESOLUTION_DYNAMIC_SITES.get(rel) dynamic = _LATE_RESOLUTION_DYNAMIC_SITES.get(rel)
if dynamic is not None: if dynamic is not None:
constants = { constants = {