config: the per-instance families read the bags (#35026)

This commit is contained in:
Cheng Wan
2026-08-17 16:17:53 -07:00
committed by GitHub
parent a97bc8db32
commit cba3c5d5ac
45 changed files with 909 additions and 530 deletions
+24 -19
View File
@@ -11,7 +11,7 @@ One container owns process-static runtime state: `sglang.srt.runtime_context.Run
| Tier | Accessor | Holds | Lifecycle |
|------|----------|-------|-----------|
| raw config seed | `get_server_args()` | the published `ServerArgs` — the startup record, for debugging, dumps and provenance. **Business code does not read fields off it**: the read ratchet pins that at zero, and "Reading config: the seed is off limits" below says what to read instead, which forms the ratchet sees, and what is outside it by construction (a runtime-computed name; a whole-object hand-off) | published at process entry; re-publish is **last-publish-wins** (the tokenizer publish in the launcher process; sequential engine rebuild in one process, e.g. unit tests) and re-projects the bags; read-only |
| resolved config | `get_exec()` `get_memory()` `get_schedule()` `get_model()` `get_spec()` `get_serving()` `get_observability()` `get_disagg()` `get_lora()` `get_mm()` `get_device()` | namespace **config bags** — the single source of truth for resolved config; leaves are real attributes (dynamo-traceable) | projected from `server_args` at `publish`; mutated only via `get_context().override` |
| resolved config | `get_exec()` `get_memory()` `get_schedule()` `get_model()` `get_spec()` `get_serving()` `get_observability()` `get_disagg()` `get_lora()` `get_mm()` `get_device()` | namespace **config bags** — the single source of truth for resolved config; leaves are real attributes (dynamo-traceable). Each is a **module function of no arguments**, and a module binds the name once: `manager.get_disagg()`, `self.get_disagg = get_disagg`, or a same-named import next to the bag one (`from model_loader import get_model`) all import fine and fail only when that path runs. `ruff --select F811` catches the import collision; `RuntimeContext` has no bag-named member and no `__getattr__`, so the member-call shapes are an `AttributeError` at call time — give it a delegating `__getattr__` and they go silent instead | projected from `server_args` at `publish`; mutated only via `get_context().override` |
| runtime flags | `get_flags()` | state that is *not* a pure function of config: `capture` (cuda-graph lifecycle), `moe` (ACTIVE backends, swappable), `dp` (DP-attention runtime flags) | materialized at subsystem init; groups offer `override()` for tests |
| resources | `get_resources()`, `get_stream(name)`, `get_buffer(name, factory)` | process-level handles: graph pools, EPLB state, EP dispatcher state, named side streams, workspace buffers | lazy; cleared by `reset_context()` |
| per-forward | `get_forward()` | forward-scoped flags (multi-stream switch, MoE output buffer, attn-TP inputs, extend-in-batch) | contextvar-backed; `scoped(**kw)` restores on exit; new threads see defaults |
@@ -108,13 +108,14 @@ bag to override at all.
the scope. When there is a runner in hand, read its
stamp; that is a different rule from "read the instance".
- **Per-instance boundaries** — the tokenizer-manager family, everything under
`entrypoints/`, and the tokenizer-process multimodal processors still read
`self.server_args` today. The old justification ("several `Engine`s can share
one process, bags are last-publish-wins across them") is **retracted** — owner
ruling (2026-08-15): a process holds at most one live config at a time
(concurrent multi-Engine is unsupported; sequential rebuild stays legal, unit
tests rely on it). These reads are scheduled to become bag reads in the
bag-read series; treat them as pinned debt, not as a boundary to imitate. What
`entrypoints/`, and the tokenizer-process multimodal processors read the bags.
The old justification for keeping them on `self.server_args` ("several
`Engine`s can share one process, bags are last-publish-wins across them") is
**retracted** — owner ruling (2026-08-15): a process holds at most one live
config at a time (concurrent multi-Engine is unsupported; sequential rebuild
stays legal, unit tests rely on it). What still reads the instance in those
files is pinned pair by pair in the exposure ratchet, each with its own
disposition; none of it is a boundary to imitate. What
genuinely stays per-instance is what differs per *worker* within one engine:
`base_gpu_id` travels as a constructor argument (`MMEncoder(gpu_id=...)`;
`BaseMultimodalProcessor._fast_image_processor_device` is the shape to copy).
@@ -226,24 +227,28 @@ this).
"Reads that legitimately stay on a ServerArgs instance" above for the full set —
per-instance boundaries and whole-object passes; there are no per-runner config
copies to read any more). The allow-list is `GrammarManager` and `MMEncoder`;
the tokenizer-manager family, `entrypoints/`, and the tokenizer-process
multimodal processors sit beside it only as pinned debt — not for one single
reason:
what sits beside it is residue, not a family — and not for one single reason:
- the tokenizer-manager family and `entrypoints/` are **pinned debt awaiting
conversion to bag reads** (the old multi-Engine justification is retracted —
one process, one live config); the reads still work today because the
instance carries resolved values;
- `GrammarManager` is a handed instance — it is constructed with the config its
owner hands it and never assumes a published namespace;
- the tokenizer-manager family and `entrypoints/` **read the bags**; what is
left of them in the exposure ratchet is a handful of individually-dispositioned
pairs, not a family awaiting conversion. Read the ratchet for the current set
rather than assuming a directory is off-limits;
- `GrammarManager` is a handed instance for its residual `self.server_args`
reads, but backend selection is **not** on the instance any more:
`create_grammar_backend` reads `get_exec().kernel.grammar_backend`, and
`__init__` calls that factory whenever `skip_tokenizer_init` is false. In
production the scheduler process has published; a test that constructs one
without publishing has to keep patching the factory (or publish itself);
- `MMEncoder` publishes the very instance it is handed (`publish(server_args,
role="encoder")`) and takes its per-worker device as a separate `gpu_id`
argument, so its `self.server_args` reads and the bag agree today. They are on
this list as a construction-path convention rather than a semantic exception —
and the residual is real: a post-publish `override` would not reach them.
Their tests tell you the same thing: they construct the object standalone, so a
bag read turns into "config namespace not published".
Their tests are not one story: a `GrammarManager` built standalone turns the
factory's bag read into "config namespace not published" unless the test patches
it or publishes, while `MMEncoder` publishes in its own `__init__` and so needs
no such arrangement.
**Test doubles publish, they do not inject.** A stand-in that carries
`server_args=SimpleNamespace(field=...)` stops working the moment production reads
@@ -22,7 +22,12 @@ from typing import Dict, List, NamedTuple, Optional, Tuple
import torch
from sglang.srt.parser.reasoning_parser import ReasoningParser
from sglang.srt.runtime_context import get_context, get_resources
from sglang.srt.runtime_context import (
get_context,
get_exec,
get_resources,
get_serving,
)
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__)
@@ -315,7 +320,7 @@ def create_grammar_backend(
eos_token_ids: Optional[set] = None,
think_end_ids: Optional[List[int]] = None,
) -> Optional[BaseGrammarBackend]:
name = server_args.grammar_backend
name = get_exec().kernel.grammar_backend
# Custom grammar backend has the highest priority
if name in GRAMMAR_BACKEND_REGISTRY:
@@ -329,7 +334,7 @@ def create_grammar_backend(
grammar_backend = OutlinesGrammarBackend(
tokenizer,
whitespace_pattern=server_args.constrained_json_whitespace_pattern,
whitespace_pattern=get_serving().constrained_json_whitespace_pattern,
)
elif name == "xgrammar":
from sglang.srt.constrained.xgrammar_backend import (
@@ -345,10 +350,10 @@ def create_grammar_backend(
tokenizer,
vocab_size=vocab_size,
model_eos_token_ids=eos_list,
any_whitespace=not server_args.constrained_json_disable_any_whitespace,
any_whitespace=not get_serving().constrained_json_disable_any_whitespace,
)
except TokenizerNotSupportedError as e:
if server_args.enable_strict_thinking:
if get_serving().enable_strict_thinking:
raise ValueError(
f"--enable-strict-thinking requires a grammar backend with "
f"token filtering support, but XGrammar failed to initialize: "
@@ -367,13 +372,13 @@ def create_grammar_backend(
grammar_backend = GuidanceBackend(
tokenizer=tokenizer,
any_whitespace=not server_args.constrained_json_disable_any_whitespace,
whitespace_pattern=server_args.constrained_json_whitespace_pattern,
any_whitespace=not get_serving().constrained_json_disable_any_whitespace,
whitespace_pattern=get_serving().constrained_json_whitespace_pattern,
n_vocab=vocab_size,
eos_token_ids=eos_token_ids,
)
elif name == "none":
if server_args.enable_strict_thinking:
if get_serving().enable_strict_thinking:
raise ValueError(
"--enable-strict-thinking requires a grammar backend that supports "
"token filtering, but grammar_backend='none' was specified. Use "
@@ -384,13 +389,13 @@ def create_grammar_backend(
else:
raise ValueError(f"Invalid grammar backend: {name}")
if server_args.reasoning_parser and think_end_ids:
if get_serving().reasoning_parser and think_end_ids:
from sglang.srt.constrained.reasoner_grammar_backend import (
ReasonerGrammarBackend,
)
reasoning_parser = ReasoningParser(
model_type=server_args.reasoning_parser,
model_type=get_serving().reasoning_parser,
stream_reasoning=False,
tokenizer=tokenizer,
)
@@ -399,7 +404,7 @@ def create_grammar_backend(
grammar_backend,
reasoning_parser,
tokenizer,
enable_strict_thinking=server_args.enable_strict_thinking,
enable_strict_thinking=get_serving().enable_strict_thinking,
)
return grammar_backend
@@ -34,6 +34,7 @@ from sglang.srt.managers.io_struct import GenerateReqInput, TokenizedGenerateReq
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
from sglang.srt.managers.schedule_batch import Modality, Req
from sglang.srt.multimodal.cache import media_preprocess_kwargs
from sglang.srt.runtime_context import get_disagg, get_exec, get_serving
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ImageData
from sglang.srt.utils.common import safe_pickle_loads
@@ -1586,7 +1587,7 @@ class MMReceiverBase(ABC):
# context alive for the process instead of creating a temporary context
# whose destruction also closes its per-request socket.
self.scheduler_context = zmq.Context()
self.encoder_transfer_backend = server_args.encoder_transfer_backend
self.encoder_transfer_backend = get_disagg().encoder_transfer_backend
# When ``encode_urls`` is shared with an :class:`EncoderBootstrapServer`
# (tokenizer manager process), it grows / shrinks in place as encoders
# register or unregister; the receiver always sees the current set.
@@ -1642,8 +1643,8 @@ class MMReceiverBase(ABC):
self.embeddings_engine = init_mooncake_transfer_engine(
hostname=self.host,
ib_device=(
server_args.disaggregation_ib_device
or server_args.mooncake_ib_device
get_disagg().disaggregation_ib_device
or get_exec().moe.mooncake_ib_device
),
)
self.embeddings_buffer = dict()
@@ -1689,7 +1690,7 @@ class MMReceiverBase(ABC):
extra_kwargs["tokenizer_backend"] = server_args.tokenizer_backend
_processor = get_processor(
server_args.tokenizer_path,
get_serving().tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
@@ -57,7 +57,7 @@ from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalSta
from sglang.srt.model_executor.model_runner_components.load_model_utils import (
maybe_precompile_model_kernels_after_loading,
)
from sglang.srt.model_loader import get_model
from sglang.srt.model_loader import get_model as load_model
from sglang.srt.multimodal.cache import parse_content_hash, snapshot_media
from sglang.srt.multimodal.encoder_preprocessing import (
EncoderPreprocessOutput,
@@ -73,10 +73,15 @@ from sglang.srt.observability.trace import (
trace_set_thread_info,
)
from sglang.srt.runtime_context import (
configured_tp_size,
get_device,
get_disagg,
get_exec,
get_mm,
get_model,
get_observability,
get_parallel,
get_serving,
publish,
)
from sglang.srt.server_args import (
@@ -315,7 +320,7 @@ class MMEncoder:
logger.info(f"init MMEncoder {rank}/{server_args.tp_size}")
self.server_args = server_args
configure_media_url_security(
server_args.allowed_media_domains,
get_mm().allowed_media_domains,
server_args.media_url_max_file_size_mb,
)
self.rank = rank
@@ -329,7 +334,7 @@ class MMEncoder:
server_args,
)
self.load_config = LoadConfig(
load_format=server_args.load_format,
load_format=get_model().load_format,
download_dir=server_args.download_dir,
model_loader_extra_config=server_args.model_loader_extra_config,
remote_instance_weight_loader_seed_instance_ip=server_args.remote_instance_weight_loader_seed_instance_ip,
@@ -340,7 +345,7 @@ class MMEncoder:
self.model_config.hf_config, "model_type", "unknown"
).lower()
self.device = server_args.device
self.device = get_device().device
self.gpu_id = server_args.base_gpu_id + rank if gpu_id is None else gpu_id
self.device_config = DeviceConfig(
@@ -354,7 +359,7 @@ class MMEncoder:
use_image_processor_gpu
and resolve_image_processor_backend(server_args) != "pil"
)
self._build_vision_config(server_args.mm_process_config)
self._build_vision_config(get_mm().mm_process_config)
self.model_audio_sr = self._resolve_audio_sr()
logger.info(f"Resolved model audio sample rate: {self.model_audio_sr} Hz")
@@ -368,7 +373,7 @@ class MMEncoder:
initialize_model_parallel(tensor_model_parallel_size=server_args.tp_size)
initialize_dp_attention(server_args, self.model_config)
self.model = get_model(
self.model = load_model(
model_config=self.model_config,
load_config=self.load_config,
device_config=self.device_config,
@@ -628,7 +633,7 @@ class MMEncoder:
)
try:
self.image_processor = AutoImageProcessor.from_pretrained(
server_args.tokenizer_path or server_args.model_path,
get_serving().tokenizer_path or get_model().model_path,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
**image_processor_kwargs,
@@ -639,7 +644,7 @@ class MMEncoder:
try:
self.video_processor = AutoVideoProcessor.from_pretrained(
server_args.tokenizer_path or server_args.model_path,
get_serving().tokenizer_path or get_model().model_path,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
)
@@ -650,7 +655,7 @@ class MMEncoder:
try:
# Note: AutoProcessor is used for audio processor
_audio_proc = AutoProcessor.from_pretrained(
server_args.tokenizer_path or server_args.model_path,
get_serving().tokenizer_path or get_model().model_path,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
)
@@ -2100,7 +2105,7 @@ class MMEncoder:
_zmq_xfer_start = time.perf_counter()
if (
self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
get_disagg().encoder_transfer_backend == "zmq_to_scheduler"
and url is not None
):
lock = self.scheduler_send_locks.get(endpoint)
@@ -2146,7 +2151,7 @@ class MMEncoder:
if encoder_metrics_collector is not None:
encoder_metrics_collector.observe_transfer(
time.perf_counter() - _zmq_xfer_start,
backend=self.server_args.encoder_transfer_backend,
backend=get_disagg().encoder_transfer_backend,
)
return
@@ -3007,7 +3012,7 @@ async def _push_embedding_to_prefill(enc: MMEncoder, request: dict) -> None:
# No-op for mooncake (its /send is separate). embedding_port=None is
# rejected upfront, so ports is always a concrete list here.
req_id = request["req_id"]
backend = enc.server_args.encoder_transfer_backend
backend = get_disagg().encoder_transfer_backend
if backend == "zmq_to_tokenizer":
await enc.send(
@@ -3050,7 +3055,7 @@ async def _dp_worker_encode_and_send(
modality = Modality.from_str(request["modality"])
time_stats.modality = modality.name.lower()
time_stats.set_metrics_collector(encoder_metrics_collector)
backend = enc.server_args.encoder_transfer_backend
backend = get_disagg().encoder_transfer_backend
# URL state lives in main process module globals; workers don't see it.
if backend == "zmq_to_scheduler" and request.get("embedding_port") is None:
@@ -3666,14 +3671,14 @@ async def run_dp_worker(
)
global encoder_metrics_collector
if server_args.enable_metrics:
if get_observability().enable_metrics:
set_prometheus_multiproc_dir()
labels = {
"model_name": server_args.served_model_name,
"model_name": get_serving().served_model_name,
"dp_rank": str(dp_rank),
}
if server_args.extra_metric_labels:
labels.update(server_args.extra_metric_labels)
if get_observability().extra_metric_labels:
labels.update(get_observability().extra_metric_labels)
encoder_metrics_collector = EncoderMetricsCollector(labels)
enc.dp_rank = dp_rank
@@ -3964,14 +3969,14 @@ def launch_server(server_args: ServerArgs):
global encoder, encoder_metrics_collector
# Set up prometheus metrics.
if server_args.enable_metrics:
if get_observability().enable_metrics:
set_prometheus_multiproc_dir()
labels = {
"model_name": server_args.served_model_name,
"model_name": get_serving().served_model_name,
"dp_rank": "0",
}
if server_args.extra_metric_labels:
labels.update(server_args.extra_metric_labels)
if get_observability().extra_metric_labels:
labels.update(get_observability().extra_metric_labels)
encoder_metrics_collector = EncoderMetricsCollector(labels)
add_prometheus_middleware(app)
@@ -3979,21 +3984,21 @@ def launch_server(server_args: ServerArgs):
zmq_ctx = zmq.Context(10)
ipc_path_prefix = random_uuid()
port_args = PortArgs.init_new(server_args)
if server_args.dist_init_addr:
na = NetworkAddress.parse(server_args.dist_init_addr)
if get_parallel().dist_init_addr:
na = NetworkAddress.parse(get_parallel().dist_init_addr)
dist_init_method = na.to_tcp()
else:
dist_init_method = NetworkAddress(
server_args.host or "127.0.0.1", port_args.nccl_port
get_serving().host or "127.0.0.1", port_args.nccl_port
).to_tcp()
if server_args.enable_trace:
if get_observability().enable_trace:
process_tracing_init(
server_args.otlp_traces_endpoint,
get_observability().otlp_traces_endpoint,
"sglang",
trace_modules=server_args.trace_modules,
trace_modules=get_observability().trace_modules,
)
trace_set_thread_info("Encoder")
for rank in range(1, server_args.tp_size):
for rank in range(1, configured_tp_size()):
schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}"
send_sockets.append(
get_zmq_socket(zmq_ctx, zmq.PUSH, schedule_path, bind=False)
@@ -4006,13 +4011,13 @@ def launch_server(server_args: ServerArgs):
encoder = MMEncoder(server_args, dist_init_method=dist_init_method)
# Register this encoder's URL with prefill server(s) if configured.
if server_args.encoder_register_urls:
if get_disagg().encoder_register_urls:
import atexit
_register_encoder_url_with_bootstrap(server_args)
atexit.register(_unregister_encoder_url_from_bootstrap, server_args)
uvicorn.run(app, host=server_args.host, port=server_args.port)
uvicorn.run(app, host=get_serving().host, port=get_serving().port)
def _launch_server_dp(server_args: ServerArgs):
@@ -4083,7 +4088,7 @@ def _launch_server_dp(server_args: ServerArgs):
proc.start()
worker_processes.append(proc)
labels = {"model_name": server_args.served_model_name}
labels = {"model_name": get_serving().served_model_name}
if server_args.extra_metric_labels:
labels.update(server_args.extra_metric_labels)
dp_dispatcher = DPDispatcher(
@@ -4188,7 +4193,7 @@ async def handle_encode_request(request: dict):
# when multiple decoder TP ranks POST /encode
# with the same req_id, only the first triggers the VIT forward;
# subsequent callers wait and return the same metadata.
if encoder.server_args.encoder_transfer_backend == "mooncake":
if get_disagg().encoder_transfer_backend == "mooncake":
async with encoder._inflight_encode_lock:
if req_id in encoder._inflight_encode_events:
event = encoder._inflight_encode_events[req_id]
@@ -4274,7 +4279,7 @@ async def handle_encode_request(request: dict):
time_stats.set_mm_encode_end_time()
if error_msg:
if encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler":
if get_disagg().encoder_transfer_backend == "zmq_to_scheduler":
if request["embedding_port"] is None:
start_background_send(req_id)
else:
@@ -4285,7 +4290,7 @@ async def handle_encode_request(request: dict):
embedding_port=port,
)
# Signal waiters on failure for mooncake
if encoder.server_args.encoder_transfer_backend == "mooncake":
if get_disagg().encoder_transfer_backend == "mooncake":
encoder._inflight_encode_meta.pop(req_id, None)
evt = encoder._inflight_encode_events.pop(req_id, None)
if evt:
@@ -4299,7 +4304,7 @@ async def handle_encode_request(request: dict):
status_code=error_code,
content={"status": "error", "message": error_msg, "req_id": req_id},
)
if encoder.server_args.encoder_transfer_backend == "mooncake":
if get_disagg().encoder_transfer_backend == "mooncake":
# Store metadata for duplicate callers and signal them
encoder._inflight_encode_meta[req_id] = (
nbytes,
@@ -4323,7 +4328,7 @@ async def handle_encode_request(request: dict):
modality=modality_str, status="success"
)
return ORJSONResponse(content=request)
elif encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler":
elif get_disagg().encoder_transfer_backend == "zmq_to_scheduler":
logger.info(f"{request['embedding_port'] = }")
if request["embedding_port"] is None:
await encoder.send_with_url(
@@ -4347,7 +4352,7 @@ async def handle_encode_request(request: dict):
modality=modality_str, status="success"
)
return ORJSONResponse(content=None)
elif encoder.server_args.encoder_transfer_backend == "zmq_to_tokenizer":
elif get_disagg().encoder_transfer_backend == "zmq_to_tokenizer":
await encoder.send(
req_id=request["req_id"],
prefill_host=request["prefill_host"],
@@ -4369,7 +4374,7 @@ async def handle_encode_request(request: dict):
logger.error(f"Unexpected error in encoder logic for {req_id}: {error_msg}")
rid_to_err_msg[req_id] = error_msg
# Ensure inflight waiters are unblocked on unexpected errors
if encoder.server_args.encoder_transfer_backend == "mooncake":
if get_disagg().encoder_transfer_backend == "mooncake":
encoder._inflight_encode_meta.pop(req_id, None)
evt = encoder._inflight_encode_events.pop(req_id, None)
if evt:
+23 -16
View File
@@ -99,7 +99,14 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa
from sglang.srt.parser.template_detection import resolve_auto_parsers
from sglang.srt.parser.template_manager import TemplateManager
from sglang.srt.plugins import load_plugins
from sglang.srt.runtime_context import get_parallel, publish
from sglang.srt.runtime_context import (
configured_pp_size,
get_exec,
get_model,
get_parallel,
get_serving,
publish,
)
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils import (
MultiprocessingSerializer,
@@ -158,9 +165,9 @@ def init_tokenizer_manager(
template_manager = TemplateManager()
template_manager.initialize_templates(
tokenizer_manager=tokenizer_manager,
model_path=server_args.model_path,
chat_template=server_args.chat_template,
completion_template=server_args.completion_template,
model_path=get_model().model_path,
chat_template=get_serving().chat_template,
completion_template=get_serving().completion_template,
)
# Resolve any remaining auto parsers using template manager's detection results
@@ -681,7 +688,7 @@ class Engine(EngineScoreMixin, EngineBase):
pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = (
_calculate_rank_ranges(
server_args.nnodes,
server_args.pp_size,
configured_pp_size(),
tp_size,
server_args.node_rank,
)
@@ -702,7 +709,7 @@ class Engine(EngineScoreMixin, EngineBase):
daemon_procs = []
logger.info(
f"Launching {num_daemons} weight cache daemon(s) on node "
f"{server_args.node_rank} for model={server_args.model_path}, "
f"{server_args.node_rank} for model={get_model().model_path}, "
f"pp_ranks={pp_rank_range.start}..{pp_rank_range.stop - 1}, "
f"tp_ranks={tp_rank_range.start}..{tp_rank_range.stop - 1}, "
f"dist_init_method={dist_init_method}"
@@ -737,7 +744,7 @@ class Engine(EngineScoreMixin, EngineBase):
"-m",
"sglang.srt.weight_cache.daemon",
"--model-path",
server_args.model_path,
get_model().model_path,
"--gpu-id",
str(gpu_id),
"--tp-size",
@@ -745,7 +752,7 @@ class Engine(EngineScoreMixin, EngineBase):
"--tp-rank",
str(tp_rank),
"--pp-size",
str(server_args.pp_size),
str(configured_pp_size()),
"--pp-rank",
str(pp_rank),
"--dp-size",
@@ -753,14 +760,14 @@ class Engine(EngineScoreMixin, EngineBase):
"--ep-size",
str(get_parallel().ep_size),
"--load-format",
server_args.load_format,
get_model().load_format,
"--dtype",
server_args.dtype,
get_model().dtype,
"--dist-init-method",
dist_init_method,
]
if server_args.quantization:
cmd += ["--quantization", server_args.quantization]
if get_model().quantization:
cmd += ["--quantization", get_model().quantization]
if (
server_args.model_loader_extra_config
and server_args.model_loader_extra_config != "{}"
@@ -863,7 +870,7 @@ class Engine(EngineScoreMixin, EngineBase):
"""
scheduler_procs = []
use_dp_controller = (
get_parallel().dp_size > 1 or server_args.ep_join_mode == "scale"
get_parallel().dp_size > 1 or get_exec().moe.ep_join_mode == "scale"
)
if not use_dp_controller:
@@ -876,7 +883,7 @@ class Engine(EngineScoreMixin, EngineBase):
pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = (
_calculate_rank_ranges(
server_args.nnodes,
server_args.pp_size,
configured_pp_size(),
server_args.tp_size,
server_args.node_rank,
)
@@ -983,7 +990,7 @@ class Engine(EngineScoreMixin, EngineBase):
processes: List[mp.Process] = []
names: List[str] = []
if server_args.detokenizer_worker_num <= 1:
if get_serving().detokenizer_worker_num <= 1:
proc = mp.Process(
target=run_detokenizer_process_func,
args=(server_args, port_args),
@@ -996,7 +1003,7 @@ class Engine(EngineScoreMixin, EngineBase):
router_ipc_name = port_args.detokenizer_ipc_name
worker_ipc_names: List[str] = []
try:
for i in range(server_args.detokenizer_worker_num):
for i in range(get_serving().detokenizer_worker_num):
worker_ipc = f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}"
port_args.detokenizer_ipc_name = worker_ipc
proc = mp.Process(
+26 -23
View File
@@ -247,7 +247,7 @@ async def init_multi_tokenizer() -> ServerArgs:
template_manager = TemplateManager()
template_manager.initialize_templates(
tokenizer_manager=tokenizer_manager,
model_path=server_args.model_path,
model_path=get_model().model_path,
chat_template=server_args.chat_template,
completion_template=server_args.completion_template,
)
@@ -293,9 +293,9 @@ async def lifespan(fast_api_app: FastAPI):
"sglang",
trace_modules=server_args.trace_modules,
)
if server_args.disaggregation_mode == "prefill":
if get_disagg().disaggregation_mode == "prefill":
thread_label = "Prefill" + thread_label
elif server_args.disaggregation_mode == "decode":
elif get_disagg().disaggregation_mode == "decode":
thread_label = "Decode" + thread_label
trace_set_thread_info(thread_label)
@@ -380,7 +380,7 @@ async def lifespan(fast_api_app: FastAPI):
# Execute custom warmups
if server_args.warmups is not None:
await execute_warmups(
server_args.disaggregation_mode,
get_disagg().disaggregation_mode,
server_args.warmups.split(","),
_global_state.tokenizer_manager,
)
@@ -393,7 +393,7 @@ async def lifespan(fast_api_app: FastAPI):
try:
if (
getattr(fast_api_app, "is_single_tokenizer_mode", False)
and server_args.grpc_port is not None
and get_serving().grpc_port is not None
and not (server_args.smg_grpc_mode or server_args.grpc_mode)
):
grpc_handle = _start_native_grpc_server_for_runtime(
@@ -401,6 +401,7 @@ async def lifespan(fast_api_app: FastAPI):
tokenizer_manager=_global_state.tokenizer_manager,
template_manager=_global_state.template_manager,
scheduler_info=_global_state.scheduler_info,
grpc_port=get_serving().grpc_port,
)
if server_args.sidecar is not None:
from sglang.srt.entrypoints.sidecar import start_sidecar
@@ -480,7 +481,13 @@ v1_loads_router.route_class = ORJSONRoute
app.include_router(v1_loads_router)
from sglang.srt.entrypoints.elastic_ep import router as elastic_ep_router
from sglang.srt.runtime_context import get_parallel
from sglang.srt.runtime_context import (
get_disagg,
get_exec,
get_model,
get_parallel,
get_serving,
)
elastic_ep_router.route_class = ORJSONRoute
app.include_router(elastic_ep_router)
@@ -679,10 +686,7 @@ async def health_generate(request: Request) -> Response:
sampling_params=sampling_params,
log_metrics=False,
)
if (
_global_state.tokenizer_manager.server_args.disaggregation_mode
!= DisaggregationMode.NULL.value
):
if get_disagg().disaggregation_mode != DisaggregationMode.NULL.value:
gri.bootstrap_host = FAKE_BOOTSTRAP_HOST
gri.bootstrap_room = 0
else:
@@ -2224,7 +2228,7 @@ def _execute_server_warmup(server_args: ServerArgs):
json_data["input_ids"] = json_data["input_ids"][0]
elif (
is_vlm
and server_args.disaggregation_mode == "null"
and get_disagg().disaggregation_mode == "null"
and model_info["is_generation"]
):
served_model_name = ""
@@ -2234,9 +2238,9 @@ def _execute_server_warmup(server_args: ServerArgs):
# _global_state.tokenizer_manager is not initialized in the rust server,
# so we need to get the model name from the model_info
served_model_name = model_info.get(
"model_path", server_args.served_model_name
"model_path", get_serving().served_model_name
)
served_model_name = served_model_name or server_args.model_path
served_model_name = served_model_name or get_model().model_path
# TODO: ChatCompletionRequest does not have bootstrap info required by disaggregation mode, disable image-warmup for now
# Only use chat completions format for generation models, not embedding models
json_data = {
@@ -2280,7 +2284,7 @@ def _execute_server_warmup(server_args: ServerArgs):
# Send a warmup request
warmup_timeout = envs.SGLANG_WARMUP_TIMEOUT.get()
try:
if server_args.disaggregation_mode == "null":
if get_disagg().disaggregation_mode == "null":
res = requests.post(
url + request_name,
json=json_data,
@@ -2314,7 +2318,7 @@ def _execute_server_warmup(server_args: ServerArgs):
else:
logger.info(
"Disaggregation warmup failed (mode=%s), status codes: %s",
server_args.disaggregation_mode,
get_disagg().disaggregation_mode,
failed_status_codes,
)
# In rust-server mode there is no TokenizerManager (readiness is
@@ -2368,10 +2372,10 @@ def _wait_and_warmup(
logger.debug(
"[Elastic EP] Skipping server warmup for elastic joiner "
"(ep_join_mode=%s)",
server_args.ep_join_mode,
get_exec().moe.ep_join_mode,
)
if not server_args.skip_server_warmup and not skip_elastic_joiner_warmup:
if not get_serving().skip_server_warmup and not skip_elastic_joiner_warmup:
if not execute_warmup_func(server_args):
return
else:
@@ -2383,7 +2387,7 @@ def _wait_and_warmup(
logger.info("The server is fired up and ready to roll!")
if server_args.delete_ckpt_after_loading:
delete_directory(server_args.model_path)
delete_directory(get_model().model_path)
if server_args.debug_tensor_dump_input_file:
kill_process_tree(os.getpid())
@@ -2711,6 +2715,7 @@ def _start_native_grpc_server_for_runtime(
tokenizer_manager,
template_manager,
scheduler_info,
grpc_port,
):
from sglang.srt.entrypoints.grpc_bridge import RuntimeHandle
from sglang.srt.rust_extensions import load_rust_extension
@@ -2726,13 +2731,11 @@ def _start_native_grpc_server_for_runtime(
grpc_handle = grpc_native.start_server(
host=server_args.host,
port=server_args.grpc_port,
port=grpc_port,
runtime_handle=runtime_handle,
worker_threads=server_args.grpc_worker_threads,
)
logger.info(
f"Native gRPC server started on {server_args.host}:{server_args.grpc_port}"
)
logger.info(f"Native gRPC server started on {server_args.host}:{grpc_port}")
return grpc_handle
@@ -2790,7 +2793,7 @@ def launch_server(
# and /get_model_info endpoints are static (200 as soon as the server
# binds, before any forward pass), so without this the first real request
# pays the cold-start cost (observed as a >60s first generation).
if not server_args.skip_server_warmup:
if not get_serving().skip_server_warmup:
_execute_server_warmup(server_args)
logger.info("The server is fired up and ready to roll!")
if launch_callback is not None:
+2
View File
@@ -37,6 +37,8 @@ def _loopback_host(host: str) -> str:
def build_sidecar_endpoint(server_args) -> str:
"""Both halves of the endpoint come from the argument: this is a helper
over a config object, callable before anything is published."""
return NetworkAddress(
_loopback_host(server_args.host), server_args.grpc_port
).to_url()
@@ -1,21 +1,17 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Optional
from typing import Optional
from sglang.srt.kv_canary.token_oracle.oracle import HashOracle
from sglang.srt.kv_canary.token_oracle.oracle_manager import TokenOracleManager
from sglang.srt.kv_canary.token_oracle.sampler import install_oracle_sampler
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
from sglang.srt.runtime_context import get_exec
def install_token_oracle_from_env(
*, server_args: ServerArgs, vocab_size: int
) -> Optional[TokenOracleManager]:
def install_token_oracle_from_env(*, vocab_size: int) -> Optional[TokenOracleManager]:
# Must be called before create_sampler() so the factory is present when the
# Sampler is first constructed.
if server_args.sampling_backend != "token_oracle":
if get_exec().kernel.sampling_backend != "token_oracle":
return None
oracle = HashOracle(vocab_size=vocab_size)
@@ -39,7 +39,7 @@ from sglang.srt.managers.io_struct import (
)
from sglang.srt.managers.multi_tokenizer_mixin import MultiHttpWorkerDetokenizerMixin
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
from sglang.srt.runtime_context import publish
from sglang.srt.runtime_context import get_device, get_serving, publish
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils import configure_logger, freeze_gc, kill_itself_when_parent_died
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
@@ -128,7 +128,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
self.vocab_size = None
else:
self.tokenizer = get_tokenizer(
server_args.tokenizer_path,
get_serving().tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
@@ -142,11 +142,11 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
def init_running_status(self, server_args: ServerArgs):
self.decode_status = LimitedCapacityDict(capacity=DETOKENIZER_MAX_STATES)
self.disable_tokenizer_batch_decode = server_args.disable_tokenizer_batch_decode
self.is_tool_call_parser_gpt_oss = server_args.tool_call_parser == "gpt-oss"
self.is_tool_call_parser_gpt_oss = get_serving().tool_call_parser == "gpt-oss"
self.soft_watchdog = Watchdog.create(
debug_name="DetokenizerManager",
watchdog_timeout=server_args.soft_watchdog_timeout,
watchdog_timeout=get_device().soft_watchdog_timeout,
soft=True,
test_stuck_time=envs.SGLANG_TEST_STUCK_DETOKENIZER.get(),
)
+1 -3
View File
@@ -1252,9 +1252,7 @@ class Scheduler(
and not get_schedule().disable_priority_preemption
)
self.new_token_ratio_tracker = NewTokenRatioTracker.from_server_args(
self.server_args
)
self.new_token_ratio_tracker = NewTokenRatioTracker.from_config()
def init_soft_watchdog(self, server_args: ServerArgs):
if (x := server_args.soft_watchdog_timeout) is not None:
@@ -566,7 +566,7 @@ class SchedulerBatchResultProcessor:
return get_required_capture_hidden_mode(
max(
batch.return_hidden_states_mode,
get_server_return_hidden_states_mode(server_args),
get_server_return_hidden_states_mode(),
),
batch.spec_info,
)
@@ -4,7 +4,7 @@ from dataclasses import dataclass
from typing import TYPE_CHECKING, Sequence
from sglang.srt.environ import envs
from sglang.srt.server_args import ServerArgs
from sglang.srt.runtime_context import get_schedule
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
@@ -18,10 +18,10 @@ class NewTokenRatioTracker:
current: float
@classmethod
def from_server_args(cls, server_args: ServerArgs) -> NewTokenRatioTracker:
def from_config(cls) -> NewTokenRatioTracker:
init = min(
envs.SGLANG_INIT_NEW_TOKEN_RATIO.get()
* server_args.schedule_conservativeness,
* get_schedule().schedule_conservativeness,
1.0,
)
min_ratio = min(
+46 -38
View File
@@ -121,7 +121,18 @@ from sglang.srt.observability.request_metrics_exporter import (
RequestMetricsExporterManager,
)
from sglang.srt.observability.trace import SpanAttributes, extract_trace_headers
from sglang.srt.runtime_context import get_parallel
from sglang.srt.runtime_context import (
get_device,
get_disagg,
get_exec,
get_lora,
get_memory,
get_mm,
get_model,
get_parallel,
get_serving,
get_spec,
)
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import (
PortArgs,
@@ -160,12 +171,13 @@ _REQUEST_STATE_WAIT_TIMEOUT = envs.SGLANG_REQUEST_STATE_WAIT_TIMEOUT.get()
logger = logging.getLogger(__name__)
def _reject_missing_dispatched_encoder_embedding(server_args, request_obj, mm_inputs):
def _reject_missing_dispatched_encoder_embedding(request_obj, mm_inputs):
"""Do not silently turn a failed EPD request into local vision work."""
disagg = get_disagg()
if (
mm_inputs is None
and server_args.language_only
and server_args.encoder_transfer_backend == "zmq_to_tokenizer"
and disagg.language_only
and disagg.encoder_transfer_backend == "zmq_to_tokenizer"
and request_obj.need_wait_for_mm_inputs
):
raise fastapi.HTTPException(
@@ -405,11 +417,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.elastic_last_error = None
self.enable_metrics = server_args.enable_metrics
self.incremental_streaming_output = server_args.incremental_streaming_output
self.enable_lora = server_args.enable_lora
self.enable_lora = get_lora().enable_lora
self.enable_trace = server_args.enable_trace
self.allow_auto_truncate = server_args.allow_auto_truncate
self.skip_tokenizer_init = server_args.skip_tokenizer_init
self.preferred_sampling_params = server_args.preferred_sampling_params
self.preferred_sampling_params = get_serving().preferred_sampling_params
self.crash_dump_folder = server_args.crash_dump_folder
# Init model config
@@ -501,7 +513,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.tokenizer = None
else:
self.tokenizer = get_tokenizer(
server_args.tokenizer_path,
get_serving().tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
@@ -522,7 +534,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.async_dynamic_batch_tokenizer = None
def _validate_cuda_vmm_feature_transport_support(self) -> None:
if self.server_args.mm_feature_transport != "cuda_vmm":
if get_mm().mm_feature_transport != "cuda_vmm":
return
from sglang.srt.model_loader.utils import get_model_architecture
@@ -633,7 +645,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# The registry dynamically updates as adapters are loaded / unloaded during runtime. It
# serves as the source of truth for available adapters and maps user-friendly LoRA names
# to internally used unique LoRA IDs.
self.lora_registry = LoRARegistry(self.server_args.lora_paths)
self.lora_registry = LoRARegistry(get_lora().lora_paths)
# Lock to serialize LoRA update operations.
# Please note that, unlike `model_update_lock`, this does not block inference, allowing
# LoRA updates and inference to overlap.
@@ -642,15 +654,13 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# point to their latest LoRARef objects, so that they can be
# dynamically loaded if needed for inference
self.lora_ref_cache: Dict[str, LoRARef] = {}
if self.server_args.lora_paths is not None:
for lora_ref in self.server_args.lora_paths:
if get_lora().lora_paths is not None:
for lora_ref in get_lora().lora_paths:
self.lora_ref_cache[lora_ref.lora_name] = lora_ref
def init_disaggregation(self, *, start_pd_bootstrap_service: bool = True):
# PD Disaggregation
self.disaggregation_mode = DisaggregationMode(
self.server_args.disaggregation_mode
)
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
# Keep a reference so the bootstrap server is not garbage-collected.
self.bootstrap_server = (
start_disagg_service(self.server_args)
@@ -688,7 +698,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Metrics
if self.enable_metrics:
engine_type = DisaggregationMode.to_engine_type(
self.server_args.disaggregation_mode
get_disagg().disaggregation_mode
)
labels = {
@@ -721,7 +731,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
configure_gc_warning(self.server_args.gc_warning_threshold_secs)
self.soft_watchdog = Watchdog.create(
debug_name="TokenizerManager",
watchdog_timeout=self.server_args.soft_watchdog_timeout,
watchdog_timeout=get_device().soft_watchdog_timeout,
soft=True,
test_stuck_time=envs.SGLANG_TEST_STUCK_TOKENIZER.get(),
)
@@ -997,7 +1007,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
isinstance(obj, EmbeddingReqInput) and obj.is_cross_encoder_request
)
if obj.input_embeds is not None:
if not self.server_args.disable_radix_cache:
if not get_memory().disable_radix_cache:
raise ValueError(
"input_embeds is provided while disable_radix_cache is False. "
"Please add `--disable-radix-cache` when you launch the server "
@@ -1062,7 +1072,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if (
not self.server_args.language_only
or self.server_args.encoder_transfer_backend == "zmq_to_tokenizer"
or get_disagg().encoder_transfer_backend == "zmq_to_tokenizer"
):
if self.server_args.language_only:
mm_inputs = await self.mm_receiver.recv_mm_data(
@@ -1071,11 +1081,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
prompt=mm_processor_input,
need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs,
)
_reject_missing_dispatched_encoder_embedding(
self.server_args, obj, mm_inputs
)
_reject_missing_dispatched_encoder_embedding(obj, mm_inputs)
if mm_inputs is None:
if self.server_args.language_only:
if get_disagg().language_only:
logger.warning(
"Encoder embedding not available, "
"falling back to local mm processing"
@@ -1089,7 +1097,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
)
elif (
self.server_args.language_only
and self.server_args.encoder_transfer_backend
and get_disagg().encoder_transfer_backend
in ["zmq_to_scheduler", "mooncake"]
and not obj.need_wait_for_mm_inputs
):
@@ -1256,12 +1264,12 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
requested_hidden_mode = get_request_return_hidden_states_mode(
obj.return_hidden_states
)
server_hidden_mode = get_server_return_hidden_states_mode(self.server_args)
server_hidden_mode = get_server_return_hidden_states_mode()
if requested_hidden_mode > server_hidden_mode:
if server_hidden_mode.need_capture():
raise ValueError(
"The requested return_hidden_states mode exceeds the "
f"server maximum `{self.server_args.return_hidden_states_mode}`. "
f"server maximum `{get_exec().features.return_hidden_states_mode}`. "
"Please launch with `--return-hidden-states-mode full` "
"to allow return_hidden_states=True."
)
@@ -1283,10 +1291,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
def _validate_mm_limits(
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
) -> None:
if not self.server_args.limit_mm_data_per_request:
if not get_mm().limit_mm_data_per_request:
return
for modality, limit in self.server_args.limit_mm_data_per_request.items():
for modality, limit in get_mm().limit_mm_data_per_request.items():
data = getattr(obj, f"{modality}_data", None)
if data:
count = len(data) if isinstance(data, list) else 1
@@ -1399,7 +1407,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
bootstrap_room = obj.bootstrap_room
if (
bootstrap_room is None
and self.server_args.disaggregation_transfer_backend == "fake"
and get_disagg().disaggregation_transfer_backend == "fake"
):
bootstrap_room = self.fake_bootstrap_room_counter
self.fake_bootstrap_room_counter += 1
@@ -1581,7 +1589,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
- Batch tokenization does not support DP attention yet, and it will make everything goes to the first rank currently
"""
return batch_size > 0 and (
self.server_args.enable_tokenizer_batch_encode
get_serving().enable_tokenizer_batch_encode
or (
(not get_parallel().enable_dp_attention)
and (not self._batch_has_text(batch_size, requests))
@@ -2464,7 +2472,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
state.time_stats.set_finished_time()
meta_info["e2e_latency"] = state.time_stats.get_e2e_latency()
if self.server_args.speculative_algorithm:
if get_spec().speculative_algorithm:
self._calculate_spec_decoding_metrics(meta_info, recv_obj, i)
if self.enable_metrics:
scheduler_time_stats = (
@@ -2787,7 +2795,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
):
# Total number of proposed draft tokens per request.
num_proposed_drafts = recv_obj.spec_verify_ct[i] * (
self.server_args.speculative_num_draft_tokens - 1
get_spec().speculative_num_draft_tokens - 1
)
num_correct_drafts = recv_obj.spec_num_correct_drafts[i]
@@ -3484,7 +3492,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# This flag will be used in _tokenize_one_request to determine processing path
if should_dispatch:
obj.need_wait_for_mm_inputs = True
if self.server_args.encoder_transfer_backend in [
if get_disagg().encoder_transfer_backend in [
"zmq_to_scheduler",
"mooncake",
]:
@@ -3603,13 +3611,13 @@ async def print_exception_wrapper(func):
def get_processor_wrapper(server_args):
return get_processor(
server_args.tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
get_serving().tokenizer_path,
tokenizer_mode=get_serving().tokenizer_mode,
trust_remote_code=get_model().trust_remote_code,
revision=get_model().revision,
image_processor_backend=resolve_image_processor_backend(server_args),
tokenizer_backend=server_args.tokenizer_backend,
model_name=server_args.model_path,
tokenizer_backend=get_serving().tokenizer_backend,
model_name=get_model().model_path,
)
@@ -567,7 +567,7 @@ class CPUGraphRunner:
self.return_hidden_states_mode = (
CaptureHiddenMode.NULL
if model_runner.is_draft_worker
else get_server_return_hidden_states_mode(model_runner.server_args)
else get_server_return_hidden_states_mode()
)
self.enable_return_hidden_states = self.return_hidden_states_mode.need_capture()
# bs -> compiled fn (text-only / skip_cross_attention=True)
@@ -32,7 +32,7 @@ import warnings
from dataclasses import dataclass
from enum import IntEnum, auto
from functools import total_ordering
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Set, Tuple, Union
from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Set, Tuple, Union
import torch
@@ -231,11 +231,12 @@ def register_attn_tp_sequence_sharded_predicate(
_attn_tp_sequence_sharded_predicate = predicate
def get_server_return_hidden_states_mode(server_args: Any) -> CaptureHiddenMode:
mode = getattr(server_args, "return_hidden_states_mode", None)
def get_server_return_hidden_states_mode() -> CaptureHiddenMode:
features = get_exec().features
mode = features.return_hidden_states_mode
if mode == "last":
return CaptureHiddenMode.LAST
if mode == "full" or getattr(server_args, "enable_return_hidden_states", False):
if mode == "full" or features.enable_return_hidden_states:
return CaptureHiddenMode.FULL
return CaptureHiddenMode.NULL
@@ -719,7 +720,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
if model_runner.is_draft_worker
else max(
batch.return_hidden_states_mode,
get_server_return_hidden_states_mode(model_runner.server_args),
get_server_return_hidden_states_mode(),
)
)
capture_hidden_mode = get_required_capture_hidden_mode(
@@ -726,7 +726,6 @@ class ModelRunner:
self._token_oracle_manager = None
return
self._token_oracle_manager = install_token_oracle_from_env(
server_args=self.server_args,
vocab_size=self.model_config.vocab_size,
)
@@ -261,8 +261,7 @@ def capture_prefill_graph(
if (
model_runner.spec_algorithm.is_eagle()
and not model_runner.is_draft_worker
and get_server_return_hidden_states_mode(model_runner.server_args)
< CaptureHiddenMode.FULL
and get_server_return_hidden_states_mode() < CaptureHiddenMode.FULL
and not check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE)
):
logger.info(
@@ -219,7 +219,7 @@ class BaseRunner(ABC):
self.return_hidden_states_mode = (
CaptureHiddenMode.NULL
if model_runner.is_draft_worker
else get_server_return_hidden_states_mode(model_runner.server_args)
else get_server_return_hidden_states_mode()
)
self.enable_return_hidden_states = self.return_hidden_states_mode.need_capture()
self.attn_tp_size = get_parallel().attn_tp_size
@@ -403,7 +403,7 @@ class BaseRunner(ABC):
capture_hidden_mode = (
CaptureHiddenMode.NULL
if mr.is_draft_worker
else get_server_return_hidden_states_mode(mr.server_args)
else get_server_return_hidden_states_mode()
)
num_tokens_per_req = 1
# A PD prefill target worker's pool has no SpeculativeState, so a
+7 -14
View File
@@ -67,7 +67,6 @@ from sglang.srt.runtime_context import (
get_memory,
get_model,
get_parallel,
get_server_args,
get_stream,
)
from sglang.srt.utils import (
@@ -159,12 +158,13 @@ for backend in CONCAT_ROPE_BACKENDS:
AttentionBackendRegistry.register(backend, _handle_concat_rope_backend)
def get_attn_forward_method(server_args, forward_batch) -> AttnForwardMethod:
def get_attn_forward_method(forward_batch) -> AttnForwardMethod:
prefill_backend, decode_backend = attention_backends()
is_decode = forward_batch.forward_mode.is_decode_or_idle()
if is_decode:
backend = server_args.decode_attention_backend or server_args.attention_backend
backend = decode_backend
else:
backend = server_args.prefill_attention_backend or server_args.attention_backend
backend = prefill_backend
if (
forward_batch.forward_mode.is_extend_without_speculative()
and backend == "fa3"
@@ -456,7 +456,6 @@ class SarvamMoEMLAAttention(nn.Module):
self.max_position_embeddings = max_position_embeddings
self.kv_cache_dtype = get_model().kv_cache_dtype
self._server_args = None
self.current_attention_backend = None
if self.q_lora_rank is None:
@@ -761,11 +760,9 @@ class SarvamMoEMLAAttention(nn.Module):
q_nope, q_pe = q.split([self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1)
k_pe = latent_cache[..., self.kv_lora_rank :].unsqueeze(1)
if self._server_args is None:
self._server_args = get_server_args()
self._set_current_attention_backend(forward_batch)
forward_method = get_attn_forward_method(self._server_args, forward_batch)
forward_method = get_attn_forward_method(forward_batch)
if forward_method == AttnForwardMethod.MHA_PREFILL:
return self._run_mha_prefill(
@@ -875,10 +872,8 @@ class SarvamMoEMLAAttention(nn.Module):
q_nope, q_pe = q.split([self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1)
k_pe = latent_cache[..., self.kv_lora_rank :].unsqueeze(1)
if self._server_args is None:
self._server_args = get_server_args()
self._set_current_attention_backend(forward_batch)
forward_method = get_attn_forward_method(self._server_args, forward_batch)
forward_method = get_attn_forward_method(forward_batch)
if forward_method == AttnForwardMethod.MHA_PREFILL:
output = self._run_mha_prefill(
@@ -933,11 +928,9 @@ class SarvamMoEMLAAttention(nn.Module):
q_nope_out, k_nope, q_pe, k_pe, forward_batch, zero_allocator = inner_state
if self._server_args is None:
self._server_args = get_server_args()
self._set_current_attention_backend(forward_batch)
forward_method = get_attn_forward_method(self._server_args, forward_batch)
forward_method = get_attn_forward_method(forward_batch)
if forward_method == AttnForwardMethod.MLA_SEPARATE_ROPE:
attn_output = self.attn_mqa(
+14 -9
View File
@@ -9,7 +9,7 @@ import struct
from dataclasses import dataclass
from enum import Enum
from pathlib import Path
from typing import TYPE_CHECKING, Any, Mapping, Optional, Protocol, runtime_checkable
from typing import Any, Mapping, Optional, Protocol, runtime_checkable
from urllib.parse import unquote, urlparse
import numpy as np
@@ -17,8 +17,7 @@ import torch
import transformers
from PIL import Image
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
from sglang.srt.runtime_context import get_mm, get_model
CONTENT_HASH_PREFIX = "sha256:"
_SHA256_HEX_LENGTH = 64
@@ -379,11 +378,17 @@ def resolve_multimodal_item_hash(
def build_processor_fingerprint(
processor: Any,
hf_config: Any,
server_args: ServerArgs,
*,
extra: Optional[Mapping[str, Any]] = None,
) -> str:
"""Fingerprint preprocessing choices that can change processor output."""
"""Fingerprint preprocessing choices that can change processor output.
Every config value comes from the published bags, which is the only source
that answers the *effective* preprocessing config. Taking any of them from
a handed ``ServerArgs`` would let two callers with the same effective
config disagree on the digest -- and an omitted one silently fingerprint
the empty config, which is how incompatible artifacts get reused.
"""
processor_payload = (
processor.preprocess_fingerprint_payload()
if isinstance(processor, PreprocessFingerprintProvider)
@@ -395,10 +400,10 @@ def build_processor_fingerprint(
"processor_class": f"{type(processor).__module__}.{type(processor).__qualname__}",
"model_type": hf_payload.get("model_type"),
"architectures": hf_payload.get("architectures"),
"model_revision": server_args.revision,
"processor_revision": server_args.revision,
"disable_fast_image_processor": server_args.disable_fast_image_processor,
"mm_process_config": server_args.mm_process_config or {},
"model_revision": get_model().revision,
"processor_revision": get_model().revision,
"disable_fast_image_processor": get_mm().disable_fast_image_processor,
"mm_process_config": get_mm().mm_process_config or {},
"processor": processor_payload,
"extra": extra or {},
}
@@ -39,6 +39,7 @@ from sglang.srt.multimodal.transport.cuda_ipc import (
MmItemMemoryPool,
get_mm_feature_pool_size_per_worker,
)
from sglang.srt.runtime_context import get_mm
from sglang.srt.utils import (
CLIENT_MEDIA_EXCEPTIONS,
configure_media_url_security,
@@ -212,10 +213,10 @@ class BaseMultimodalProcessor(ABC):
self.server_args = server_args
self.transport_mode = transport_mode
configure_media_url_security(
server_args.allowed_media_domains,
get_mm().allowed_media_domains,
server_args.media_url_max_file_size_mb,
)
configured_mm_feature_transport = server_args.mm_feature_transport
configured_mm_feature_transport = get_mm().mm_feature_transport
self.mm_feature_transport = (
configured_mm_feature_transport
if configured_mm_feature_transport in ("cpu", "cuda_ipc", "cuda_vmm")
@@ -231,7 +232,7 @@ class BaseMultimodalProcessor(ABC):
self.disable_fast_image_processor = self.image_processor_backend == "pil"
self.skip_tokenizer_init = server_args.skip_tokenizer_init
mm_process_config = self.server_args.mm_process_config
mm_process_config = get_mm().mm_process_config
self.image_config = mm_process_config.get("image", {})
self.video_config = mm_process_config.get("video", {})
self.audio_config = mm_process_config.get("audio", {})
@@ -255,7 +256,7 @@ class BaseMultimodalProcessor(ABC):
# The fingerprint is needed only to build artifact keys. Avoid inspecting
# processor state when this processor will never retain artifacts.
self.processor_fingerprint = (
build_processor_fingerprint(self, hf_config, server_args)
build_processor_fingerprint(self, hf_config)
if self.mm_preprocess_cache.enabled
else None
)
@@ -39,6 +39,7 @@ from sglang.srt.multimodal.processors.mimo_audio import (
MiMoAudioPipeline,
)
from sglang.srt.multimodal.processors.qwen_vl import smart_nframes
from sglang.srt.runtime_context import get_device
from sglang.srt.utils import ImageData, VideoData
from sglang.srt.utils.common import download_remote_media
from sglang.utils import logger
@@ -1588,7 +1589,7 @@ class MiMoV2Processor(BaseMultimodalProcessor):
processor_config, "video_end_token_id"
)
self.use_image_processor_gpu = envs.SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU.get()
device = server_args.device if self.use_image_processor_gpu else None
device = get_device().device if self.use_image_processor_gpu else None
self.mimo_processor = MiMoProcessor(
tokenizer=self._processor.tokenizer,
@@ -31,7 +31,11 @@ from sglang.srt.ray.engine import (
_get_bundle_node_ip,
_resolve_bundle_indices,
)
from sglang.srt.runtime_context import get_parallel
from sglang.srt.runtime_context import (
configured_attn_cp_size,
configured_pp_size,
get_parallel,
)
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils.network import bind_port, get_zmq_socket, get_zmq_socket_on_host
@@ -144,7 +148,10 @@ class RayDataParallelController(DataParallelController):
for node_idx in range(nnodes):
bundle_idx = self.bundle_for_node[node_idx]
pp_range, tp_range, pp_per_node, tp_per_node = _calculate_rank_ranges(
nnodes, server_args.pp_size, server_args.tp_size, node_rank=node_idx
nnodes,
configured_pp_size(),
server_args.tp_size,
node_rank=node_idx,
)
for pp_rank in pp_range:
for tp_rank in tp_range:
@@ -161,7 +168,7 @@ class RayDataParallelController(DataParallelController):
tp_rank,
server_args.tp_size,
get_parallel().dp_size,
server_args.attn_cp_size,
configured_attn_cp_size(),
)
rank_port_args = PortArgs.init_new(
server_args, actual_dp_rank, worker_ports
@@ -202,7 +209,7 @@ class RayDataParallelController(DataParallelController):
world_size = _compute_world_size(server_args)
bundle_indices = _resolve_bundle_indices(self.pg, world_size)
ranks_per_tp_group = server_args.tp_size * server_args.pp_size
ranks_per_tp_group = server_args.tp_size * configured_pp_size()
if dp_rank is not None:
start_rank = dp_rank * ranks_per_tp_group
end_rank = start_rank + ranks_per_tp_group
@@ -232,7 +239,7 @@ class RayDataParallelController(DataParallelController):
tp_rank,
server_args.tp_size,
get_parallel().dp_size,
server_args.attn_cp_size,
configured_attn_cp_size(),
)
rank_port_args = PortArgs.init_new(
server_args, actual_dp_rank, worker_ports
+9 -9
View File
@@ -32,7 +32,7 @@ from sglang.srt.entrypoints.engine import (
)
from sglang.srt.environ import envs
from sglang.srt.ray.scheduler_actor import SchedulerActor
from sglang.srt.runtime_context import get_parallel
from sglang.srt.runtime_context import configured_pp_size, get_parallel
from sglang.srt.server_args import PortArgs, ServerArgs
logger = logging.getLogger(__name__)
@@ -109,8 +109,8 @@ def _compute_world_size(server_args: ServerArgs) -> int:
Normal: dp_size * tp_size * pp_size; DP attention: tp_size * pp_size.
"""
if get_parallel().enable_dp_attention:
return server_args.tp_size * server_args.pp_size
return get_parallel().dp_size * server_args.tp_size * server_args.pp_size
return server_args.tp_size * configured_pp_size()
return get_parallel().dp_size * server_args.tp_size * configured_pp_size()
def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]:
@@ -269,10 +269,10 @@ class RayEngine(Engine):
)
if get_parallel().enable_dp_attention:
total_gpus = server_args.tp_size * server_args.pp_size
total_gpus = server_args.tp_size * configured_pp_size()
else:
total_gpus = (
get_parallel().dp_size * server_args.tp_size * server_args.pp_size
get_parallel().dp_size * server_args.tp_size * configured_pp_size()
)
nnodes = server_args.nnodes
@@ -332,7 +332,7 @@ class RayEngine(Engine):
pp_range, tp_range, pp_per_node, tp_per_node = (
_calculate_rank_ranges(
nnodes,
server_args.pp_size,
configured_pp_size(),
server_args.tp_size,
node_rank=node_idx,
)
@@ -449,16 +449,16 @@ class RayEngine(Engine):
if get_parallel().enable_dp_attention:
# DP attention folds DP into TP — total GPUs = tp_size * pp_size
total_gpus = server_args.tp_size * server_args.pp_size
total_gpus = server_args.tp_size * configured_pp_size()
else:
total_gpus = (
get_parallel().dp_size * server_args.tp_size * server_args.pp_size
get_parallel().dp_size * server_args.tp_size * configured_pp_size()
)
gpus_per_node = total_gpus // server_args.nnodes
logger.info(
f"Ray DP cluster: {server_args.nnodes} nodes, "
f"{gpus_per_node} GPUs/node, dp_size={get_parallel().dp_size}, "
f"tp_size={server_args.tp_size}, pp_size={server_args.pp_size}, "
f"tp_size={server_args.tp_size}, pp_size={configured_pp_size()}, "
f"enable_dp_attention={get_parallel().enable_dp_attention}"
)
@@ -2,13 +2,13 @@ from __future__ import annotations
import os
import unittest
from types import SimpleNamespace
os.environ["SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE"] = "1"
from sglang.srt.kv_canary.token_oracle.install import install_token_oracle_from_env
from sglang.srt.kv_canary.token_oracle.oracle import HashOracle
from sglang.srt.layers.sampler import _CUSTOM_SAMPLER_FACTORIES
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.test_utils import CustomTestCase
@@ -16,23 +16,26 @@ register_cuda_ci(est_time=60, stage="extra-a", runner_config="1-gpu-small")
register_amd_ci(est_time=60, suite="extra-a-test-1-gpu-small-amd")
def _make_server_args(*, sampling_backend: str) -> SimpleNamespace:
return SimpleNamespace(sampling_backend=sampling_backend)
def _publish(case, *, sampling_backend: str) -> None:
"""The gate reads the published config, so the test publishes one."""
override = get_context().override_server_args(sampling_backend=sampling_backend)
override.install()
case.addCleanup(override.restore)
class TestInstallTokenOracleFromEnv(CustomTestCase):
def test_install_token_oracle_from_env_disabled_returns_none(self) -> None:
"""Verify server-arg-disabled token oracle installation (sampling_backend != 'token_oracle') returns no TokenOracleManager."""
server_args = _make_server_args(sampling_backend="auto")
hook = install_token_oracle_from_env(server_args=server_args, vocab_size=1000)
_publish(self, sampling_backend="auto")
hook = install_token_oracle_from_env(vocab_size=1000)
self.assertIsNone(hook)
def test_install_token_oracle_from_env_enabled_registers_oracle_backend(
self,
) -> None:
"""Verify token oracle installation via sampling_backend='token_oracle' registers the oracle backend."""
server_args = _make_server_args(sampling_backend="token_oracle")
hook = install_token_oracle_from_env(server_args=server_args, vocab_size=512)
_publish(self, sampling_backend="token_oracle")
hook = install_token_oracle_from_env(vocab_size=512)
self.assertIsNotNone(hook)
self.assertIn("token_oracle", _CUSTOM_SAMPLER_FACTORIES)
@@ -40,8 +43,8 @@ class TestInstallTokenOracleFromEnv(CustomTestCase):
self,
) -> None:
"""Verify token oracle installation via sampling_backend='token_oracle' returns a TokenOracleManager wrapping a HashOracle."""
server_args = _make_server_args(sampling_backend="token_oracle")
hook = install_token_oracle_from_env(server_args=server_args, vocab_size=256)
_publish(self, sampling_backend="token_oracle")
hook = install_token_oracle_from_env(vocab_size=256)
self.assertIsNotNone(hook)
self.assertIsInstance(hook.oracle, HashOracle)
self.assertEqual(hook.oracle.vocab_size, 256)
@@ -28,6 +28,7 @@ from sglang.srt.constrained.base_grammar_backend import (
create_grammar_backend,
register_grammar_backend,
)
from sglang.srt.runtime_context import get_context # noqa: E402
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(2.0, "base-a-test-cpu")
@@ -231,18 +232,39 @@ class TestCreateGrammarBackend(unittest.TestCase):
GRAMMAR_BACKEND_REGISTRY.clear()
GRAMMAR_BACKEND_REGISTRY.update(self._saved)
def _publish(self, **fields):
"""Set the config the factory reads.
The factory takes every config value off the published bags, so a test
that sets one on the handed object would be setting something the
factory does not read -- which is how a mismatch between the two used
to stay invisible here.
"""
override = get_context().override_server_args(**fields)
override.install()
self.addCleanup(override.restore)
def _make_server_args(
self, backend="none", reasoning_parser=None, enable_strict_thinking=False
self,
backend="none",
reasoning_parser=None,
enable_strict_thinking=False,
**fields
):
published = {
"grammar_backend": backend,
"reasoning_parser": reasoning_parser,
"enable_strict_thinking": enable_strict_thinking,
"constrained_json_whitespace_pattern": None,
"constrained_json_disable_any_whitespace": False,
}
published.update(fields)
self._publish(**published)
# Handed on to plugin-registered backends; not a config source.
args = MagicMock()
args.override = lambda source, **updates: [
setattr(args, key, value) for key, value in updates.items()
]
args.grammar_backend = backend
args.reasoning_parser = reasoning_parser
args.enable_strict_thinking = enable_strict_thinking
args.constrained_json_whitespace_pattern = None
args.constrained_json_disable_any_whitespace = False
return args
def test_none_backend_returns_none(self):
@@ -293,8 +315,9 @@ class TestCreateGrammarBackend(unittest.TestCase):
def test_outlines_backend(self, mock_outlines_cls):
mock_backend = MagicMock(spec=BaseGrammarBackend)
mock_outlines_cls.return_value = mock_backend
args = self._make_server_args("outlines")
args.constrained_json_whitespace_pattern = r"\s*"
args = self._make_server_args(
"outlines", constrained_json_whitespace_pattern=r"\s*"
)
result = create_grammar_backend(args, "tok", 32000)
mock_outlines_cls.assert_called_once_with("tok", whitespace_pattern=r"\s*")
@@ -304,8 +327,9 @@ class TestCreateGrammarBackend(unittest.TestCase):
def test_xgrammar_backend(self, mock_xgrammar_cls):
mock_backend = MagicMock(spec=BaseGrammarBackend)
mock_xgrammar_cls.return_value = mock_backend
args = self._make_server_args("xgrammar")
args.constrained_json_disable_any_whitespace = True
args = self._make_server_args(
"xgrammar", constrained_json_disable_any_whitespace=True
)
result = create_grammar_backend(args, "tok", 32000, {1, 2})
mock_xgrammar_cls.assert_called_once_with(
@@ -336,9 +360,11 @@ class TestCreateGrammarBackend(unittest.TestCase):
def test_llguidance_backend(self, mock_guidance_cls):
mock_backend = MagicMock(spec=BaseGrammarBackend)
mock_guidance_cls.return_value = mock_backend
args = self._make_server_args("llguidance")
args.constrained_json_disable_any_whitespace = False
args.constrained_json_whitespace_pattern = r"\s+"
args = self._make_server_args(
"llguidance",
constrained_json_disable_any_whitespace=False,
constrained_json_whitespace_pattern=r"\s+",
)
result = create_grammar_backend(args, "tok", 32000, {1, 2})
mock_guidance_cls.assert_called_once_with(
@@ -42,7 +42,7 @@ from sglang.srt.multimodal.kimi_k3_image_processing import (
materialize_kimi_k3_cpu_features,
prepare_kimi_k3_encoder_inputs,
)
from sglang.srt.runtime_context import get_context
from sglang.srt.runtime_context import get_context, publish, reset_context
from sglang.srt.server_args import resolve_encoder_transfer_backend
from sglang.srt.utils import ImageData
from sglang.test.ci.ci_register import register_cpu_ci
@@ -78,29 +78,32 @@ def test_kimi_k3_encoder_transfer_backend_auto_avoids_tp_fanout():
def test_epd_language_only_rejects_missing_dispatched_embedding():
server_args = SimpleNamespace(
override = get_context().override_server_args(
language_only=True,
encoder_transfer_backend="zmq_to_tokenizer",
)
request = SimpleNamespace(need_wait_for_mm_inputs=True)
override.install()
try:
request = SimpleNamespace(need_wait_for_mm_inputs=True)
with pytest.raises(HTTPException) as exc_info:
_reject_missing_dispatched_encoder_embedding(server_args, request, None)
with pytest.raises(HTTPException) as exc_info:
_reject_missing_dispatched_encoder_embedding(request, None)
assert getattr(exc_info.value, "status_code", None) == 503
assert getattr(exc_info.value, "status_code", None) == 503
finally:
override.restore()
def test_epd_rejection_reads_the_resolved_transfer_backend():
"""Tripwire for step 12: this guard fires on the *resolved* backend.
"""This guard fires on the *resolved* backend.
The record is produced by actual resolution -- a language-only Kimi-K3
launch at TP2, whose `encoder_transfer_backend` starts at the argument
default `"auto"` (`ENCODER_TRANSFER_BACKEND_CHOICES[0]`) and is filled in
by `resolve_encoder_transfer_backend` to `"zmq_to_tokenizer"`. Today the
guard therefore rejects. When step 12 makes the instance raw, this same
launch hands the guard a record still at `"auto"`, the rejection silently
stops, and *this test fails* -- which is the signal to give this reader
the resolved value (per-engine overlay or bag) rather than the record.
by `resolve_encoder_transfer_backend` to `"zmq_to_tokenizer"`. The guard
reads that resolved value out of the published bags, so the rejection
survives the record going raw: what a reader must never do is go back to
the record for this field.
Fixed doubles cannot trip on that change, so the record here must come
from resolution, not a SimpleNamespace.
"""
@@ -188,20 +191,30 @@ def test_epd_rejection_reads_the_resolved_transfer_backend():
shutil.rmtree(config_dir, ignore_errors=True)
assert resolved.encoder_transfer_backend == "zmq_to_tokenizer"
request = SimpleNamespace(need_wait_for_mm_inputs=True)
with pytest.raises(HTTPException) as exc_info:
_reject_missing_dispatched_encoder_embedding(resolved, request, None)
assert getattr(exc_info.value, "status_code", None) == 503
# Publish that record: the guard reads the resolved value out of the bags,
# so a raw record does not silently disable the rejection.
publish(resolved, role="tokenizer")
try:
request = SimpleNamespace(need_wait_for_mm_inputs=True)
with pytest.raises(HTTPException) as exc_info:
_reject_missing_dispatched_encoder_embedding(request, None)
assert getattr(exc_info.value, "status_code", None) == 503
finally:
reset_context()
def test_epd_allows_local_processing_when_request_was_not_dispatched():
server_args = SimpleNamespace(
override = get_context().override_server_args(
language_only=True,
encoder_transfer_backend="zmq_to_tokenizer",
)
request = SimpleNamespace(need_wait_for_mm_inputs=False)
override.install()
try:
request = SimpleNamespace(need_wait_for_mm_inputs=False)
_reject_missing_dispatched_encoder_embedding(server_args, request, None)
_reject_missing_dispatched_encoder_embedding(request, None)
finally:
override.restore()
def _encoder(model_type="kimi_k3"):
@@ -8,13 +8,21 @@ from sglang.srt.managers.scheduler_components.batch_result_processor import (
SchedulerBatchResultProcessor,
)
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
def _make_processor(server_mode: str = "full") -> SchedulerBatchResultProcessor:
def _make_processor(case, server_mode: str = "full") -> SchedulerBatchResultProcessor:
# The server-side hidden-state ceiling is a bag leaf.
override = get_context().override_server_args(
enable_return_hidden_states=True,
return_hidden_states_mode=server_mode,
)
override.install()
case.addCleanup(override.restore)
metrics_reporter = Mock()
metrics_reporter.num_generated_tokens = 0
metrics_reporter.forward_ct_decode = 0
@@ -26,8 +34,6 @@ def _make_processor(server_mode: str = "full") -> SchedulerBatchResultProcessor:
server_args=SimpleNamespace(
enable_metrics=False,
enable_hisparse=False,
enable_return_hidden_states=True,
return_hidden_states_mode=server_mode,
),
model_config=SimpleNamespace(think_end_ids=None),
token_to_kv_pool_allocator=Mock(),
@@ -139,7 +145,7 @@ class TestPrefillHiddenStateOffsets(CustomTestCase):
can_run_cuda_graph=False,
skipped_output_comm=False,
)
processor = _make_processor(server_mode)
processor = _make_processor(self, server_mode)
with (
patch(
@@ -160,7 +166,7 @@ class TestPrefillHiddenStateOffsets(CustomTestCase):
class TestDecodeHiddenStateRetention(CustomTestCase):
def test_last_mode_multi_step_storage_stays_bounded(self):
processor = _make_processor()
processor = _make_processor(self)
req = _DecodeReq()
batch = SimpleNamespace(
reqs=[req],
@@ -4,6 +4,7 @@ from unittest.mock import Mock
from sglang.srt.managers.io_struct import GenerateReqInput
from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -11,19 +12,21 @@ register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestHiddenStateServerMode(CustomTestCase):
@staticmethod
def _make_tokenizer_manager(mode):
def _make_tokenizer_manager(self, mode):
# The server-side hidden-state mode is a bag leaf.
override = get_context().override_server_args(
enable_return_hidden_states=mode is not None,
return_hidden_states_mode=mode,
)
override.install()
self.addCleanup(override.restore)
manager = TokenizerManager.__new__(TokenizerManager)
manager.context_len = 128
manager.num_reserved_tokens = 0
manager.allow_auto_truncate = False
manager.validate_total_tokens = False
manager.is_generation = True
manager.server_args = SimpleNamespace(
enable_return_hidden_states=mode is not None,
return_hidden_states_mode=mode,
enable_custom_logit_processor=False,
)
manager.server_args = SimpleNamespace(enable_custom_logit_processor=False)
manager._validate_token_ids_logprob = Mock()
return manager
@@ -7,6 +7,7 @@ from unittest.mock import MagicMock, patch
import torch
from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_context
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -80,14 +81,20 @@ class TestBaseProcessorConfigExtraction(CustomTestCase):
BaseMultimodalProcessor,
)
# The multimodal config comes from the bags.
override = get_context().override_server_args(
mm_process_config=mm_process_config,
allowed_media_domains=[],
)
override.install()
self.addCleanup(override.restore)
server_args = MagicMock()
server_args.mm_process_config = mm_process_config
server_args.mm_processor_worker_num = mm_processor_worker_num
server_args.mm_io_worker_num = mm_io_worker_num
server_args.mm_preprocess_cache_size_mb = None
server_args.tokenizer_worker_num = 1
server_args.trust_mm_content_hashes = False
server_args.allowed_media_domains = []
server_args.media_url_max_file_size_mb = 64
hf_config = MagicMock()
@@ -170,8 +177,14 @@ class TestBaseProcessorConfigExtraction(CustomTestCase):
class TestMultimodalFeatureTransportRuntime(CustomTestCase):
@staticmethod
def _server_args(mm_feature_transport):
def _server_args(self, mm_feature_transport):
override = get_context().override_server_args(
mm_feature_transport=mm_feature_transport,
mm_process_config={},
allowed_media_domains=[],
)
override.install()
self.addCleanup(override.restore)
return SimpleNamespace(
mm_feature_transport=mm_feature_transport,
image_processor_backend="auto",
@@ -197,8 +210,7 @@ class TestMultimodalFeatureTransportRuntime(CustomTestCase):
return processor
def test_cuda_ipc_pool_uses_resolved_server_arg(self):
# The processor module can be imported before this instance is built;
# transport policy must still resolve from the instance's ServerArgs.
# Transport policy resolves from the mm bag, so the test publishes it.
from sglang.srt.multimodal.processors import base_processor
with (
@@ -792,16 +804,21 @@ class TestDoubleBosGuard(CustomTestCase):
BaseMultimodalProcessor,
)
override = get_context().override_server_args(
mm_process_config={},
mm_feature_transport="cpu",
allowed_media_domains=[],
)
override.install()
self.addCleanup(override.restore)
server_args = MagicMock()
server_args.mm_process_config = {}
server_args.mm_processor_worker_num = 0
server_args.mm_io_worker_num = 0
server_args.mm_feature_transport = "cpu"
server_args.disable_fast_image_processor = True
server_args.mm_preprocess_cache_size_mb = None
server_args.tokenizer_worker_num = 1
server_args.trust_mm_content_hashes = False
server_args.allowed_media_domains = []
server_args.media_url_max_file_size_mb = 64
mock_hf_processor = MagicMock()
@@ -35,6 +35,7 @@ from sglang.srt.managers.tokenizer_manager import ( # noqa: E402
from sglang.srt.observability.req_time_stats import ( # noqa: E402
APIServerReqTimeStats,
)
from sglang.srt.runtime_context import get_context
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
@@ -103,8 +104,15 @@ _PER_REQUEST_OPTIONAL_FIELDS = frozenset(
)
def _make_tokenizer_manager() -> TokenizerManager:
"""Create a TokenizerManager with mocked dependencies, bypassing __init__."""
def _make_tokenizer_manager(case) -> TokenizerManager:
"""Create a TokenizerManager with mocked dependencies, bypassing __init__.
The config it reads comes from the bags, so the stand-in needs a published
config rather than attributes on a mock.
"""
override = get_context().override_server_args(speculative_algorithm=None)
override.install()
case.addCleanup(override.restore)
tm = TokenizerManager.__new__(TokenizerManager)
tm.server_args = MagicMock()
tm._config_updates = []
@@ -212,7 +220,7 @@ class TestRidToStateCleanupOnAbort(CustomTestCase):
def test_abort_removes_rid_from_state(self):
"""After _handle_abort_req, rid should be removed from rid_to_state."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "abort_test_rid"
state = _make_req_state(rid)
tm.rid_to_state[rid] = state
@@ -224,7 +232,7 @@ class TestRidToStateCleanupOnAbort(CustomTestCase):
def test_abort_allows_resubmit_same_rid(self):
"""After abort, _init_req_state should accept the same rid again."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "resubmit_after_abort_rid"
state = _make_req_state(rid)
tm.rid_to_state[rid] = state
@@ -245,7 +253,7 @@ class TestRidToStateCleanupOnAbort(CustomTestCase):
def test_abort_sets_finished_and_notifies(self):
"""_handle_abort_req should mark state as finished and set the event."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "abort_notify_rid"
state = _make_req_state(rid)
tm.rid_to_state[rid] = state
@@ -266,7 +274,7 @@ class TestRidToStateCleanupOnBatchOutput(CustomTestCase):
def test_batch_output_removes_rid_on_finish(self):
"""When a request finishes in _handle_batch_output, rid should be removed."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "batch_finish_rid"
state = _make_req_state(rid)
tm.rid_to_state[rid] = state
@@ -278,7 +286,7 @@ class TestRidToStateCleanupOnBatchOutput(CustomTestCase):
def test_batch_output_allows_resubmit_after_finish(self):
"""After a request finishes, the same rid can be resubmitted."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "batch_resubmit_rid"
state = _make_req_state(rid)
tm.rid_to_state[rid] = state
@@ -299,7 +307,7 @@ class TestRidToStateCleanupOnBatchOutput(CustomTestCase):
def test_batch_output_keeps_rid_when_not_finished(self):
"""When a request is not yet finished, rid should remain in rid_to_state."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "batch_ongoing_rid"
state = _make_req_state(rid)
tm.rid_to_state[rid] = state
@@ -316,7 +324,7 @@ class TestInitReqStateDuplicateDetection(CustomTestCase):
def test_duplicate_rid_raises_error(self):
"""_init_req_state should raise ValueError if rid already exists."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "duplicate_rid"
state = _make_req_state(rid)
tm.rid_to_state[rid] = state
@@ -334,7 +342,7 @@ class TestInitReqStateDuplicateDetection(CustomTestCase):
def test_unique_rid_succeeds(self):
"""_init_req_state should succeed with a unique rid."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "unique_rid"
obj = Mock(spec=GenerateReqInput)
@@ -353,7 +361,7 @@ class TestResubmitAfterCompletion(CustomTestCase):
def test_complete_then_resubmit_same_rid(self):
"""A request that completes normally should allow resubmission with the same rid."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "complete_resubmit_rid"
# Phase 1: simulate a request in rid_to_state, then complete it
@@ -379,7 +387,7 @@ class TestResubmitAfterCompletion(CustomTestCase):
def test_abort_then_resubmit_same_rid(self):
"""An aborted request should allow resubmission with the same rid."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "abort_resubmit_rid"
# Phase 1: simulate a request, then abort it
@@ -413,9 +421,9 @@ class _DummyAsyncCM:
return False
def _make_tm_for_generate() -> TokenizerManager:
def _make_tm_for_generate(case) -> TokenizerManager:
"""Augment the mocked TokenizerManager with what generate_request needs."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(case)
tm.server_args.language_only = False
tm.server_args.tokenizer_worker_num = 1
tm.server_args.enable_strict_thinking = False
@@ -450,7 +458,7 @@ class TestDiscardPendingReqStates(CustomTestCase):
"""Direct tests for _discard_pending_req_states."""
def test_discard_single(self):
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rid = "d_single"
tm.rid_to_state[rid] = _make_req_state(rid)
obj = Mock(spec=GenerateReqInput)
@@ -460,7 +468,7 @@ class TestDiscardPendingReqStates(CustomTestCase):
self.assertNotIn(rid, tm.rid_to_state)
def test_discard_batch_removes_all(self):
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
rids = ["d0", "d1", "d2"]
for r in rids:
tm.rid_to_state[r] = _make_req_state(r)
@@ -473,7 +481,7 @@ class TestDiscardPendingReqStates(CustomTestCase):
def test_discard_ignores_already_removed(self):
"""Popping a rid that is no longer present must not raise."""
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
tm.rid_to_state["p1"] = _make_req_state("p1")
obj = Mock(spec=GenerateReqInput)
obj.is_single = False
@@ -484,7 +492,7 @@ class TestDiscardPendingReqStates(CustomTestCase):
class TestParallelStreamTaskCleanup(CustomTestCase):
def test_failing_choice_cancels_and_closes_sibling_waiters(self):
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
async def drive():
sibling_closed = asyncio.Event()
@@ -512,7 +520,7 @@ class TestParallelStreamTaskCleanup(CustomTestCase):
asyncio.run(drive())
def test_failing_non_stream_choice_cancels_and_closes_sibling_waiters(self):
tm = _make_tokenizer_manager()
tm = _make_tokenizer_manager(self)
async def drive():
sibling_closed = asyncio.Event()
@@ -546,7 +554,7 @@ class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase):
"""
def test_single_failure_before_dispatch_cleans_up(self):
tm = _make_tm_for_generate()
tm = _make_tm_for_generate(self)
rid = "single_overlen"
obj = _make_generate_obj(rid, is_single=True)
# Simulate over-length rejection during tokenization/validation.
@@ -566,7 +574,7 @@ class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase):
self.assertNotIn(rid, tm.rid_to_state)
def test_batch_failure_before_dispatch_cleans_up_all(self):
tm = _make_tm_for_generate()
tm = _make_tm_for_generate(self)
rids = ["b0", "b1", "b2"]
obj = _make_generate_obj(list(rids), is_single=False)
@@ -588,7 +596,7 @@ class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase):
self.assertNotIn(r, tm.rid_to_state)
def test_thinking_budget_rejects_runtime_without_strict_thinking(self):
tm = _make_tm_for_generate()
tm = _make_tm_for_generate(self)
obj = GenerateReqInput(
text="hello",
rid="thinking-budget",
@@ -15,6 +15,7 @@ from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import (
PrefillCudaGraphRunner,
)
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -23,31 +24,29 @@ register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestHiddenStateGraphRecapture(CustomTestCase):
def test_server_mode_sets_graph_capture_ceiling(self):
disabled = SimpleNamespace(
enable_return_hidden_states=False,
return_hidden_states_mode=None,
)
last = SimpleNamespace(
enable_return_hidden_states=True,
return_hidden_states_mode="last",
)
full = SimpleNamespace(
enable_return_hidden_states=True,
return_hidden_states_mode="full",
)
self.assertEqual(
get_server_return_hidden_states_mode(disabled),
CaptureHiddenMode.NULL,
)
self.assertEqual(
get_server_return_hidden_states_mode(last),
CaptureHiddenMode.LAST,
)
self.assertEqual(
get_server_return_hidden_states_mode(full),
CaptureHiddenMode.FULL,
cases = (
(dict(enable_return_hidden_states=False), CaptureHiddenMode.NULL),
(
dict(
enable_return_hidden_states=True, return_hidden_states_mode="last"
),
CaptureHiddenMode.LAST,
),
(
dict(
enable_return_hidden_states=True, return_hidden_states_mode="full"
),
CaptureHiddenMode.FULL,
),
)
for fields, expected in cases:
with self.subTest(**fields):
override = get_context().override_server_args(**fields)
override.install()
try:
self.assertEqual(get_server_return_hidden_states_mode(), expected)
finally:
override.restore()
@staticmethod
def _make_runner(runner_cls, capture_hidden_mode):
@@ -17,6 +17,7 @@ from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import (
PrefillCudaGraphRunner,
)
from sglang.srt.model_executor.runner.shape_key import ShapeKey
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -111,13 +112,17 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
def test_eagle_target_tc_piecewise_skips_last_mode_capture(self):
eager_runner = object()
# The server-side hidden-state ceiling is a bag leaf.
override = get_context().override_server_args(
enable_return_hidden_states=True,
return_hidden_states_mode="last",
)
override.install()
self.addCleanup(override.restore)
model_runner = SimpleNamespace(
is_draft_worker=False,
spec_algorithm=SimpleNamespace(is_eagle=lambda: True),
server_args=SimpleNamespace(
enable_return_hidden_states=True,
return_hidden_states_mode="last",
),
server_args=SimpleNamespace(),
)
with patch.object(
+105 -98
View File
@@ -66,7 +66,7 @@ from sglang.srt.multimodal.transport.cuda_ipc import (
DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY,
CudaIpcTensorTransportProxy,
)
from sglang.srt.runtime_context import get_context, get_parallel
from sglang.srt.runtime_context import get_context, get_parallel, publish, reset_context
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ImageData
from sglang.test.ci.ci_register import register_cpu_ci
@@ -642,11 +642,9 @@ def _k3_preprocess_config(
)
def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls):
server_args = SimpleNamespace(
mm_feature_transport="cpu",
image_processor_backend="auto",
disable_fast_image_processor=False,
skip_tokenizer_init=False,
mm_process_config={},
mm_io_worker_num=0,
mm_processor_worker_num=0,
tokenizer_worker_num=1,
@@ -654,32 +652,35 @@ def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls
trust_mm_content_hashes=False,
base_gpu_id=0,
rl_on_policy_target=None,
allowed_media_domains=[],
media_url_max_file_size_mb=64,
)
processor = processor_cls(
hf_config=SimpleNamespace(media_placeholder_token_id=42),
server_args=server_args,
_processor=_HFProcessor(),
transport_mode=None,
)
try:
worker_processor = asyncio.run(
processor.mm_processor_executor.run(lambda *, processor: processor)
# The multimodal config comes from the bags.
with get_context().override_server_args(
mm_feature_transport="cpu", mm_process_config={}, allowed_media_domains=[]
):
processor = processor_cls(
hf_config=SimpleNamespace(media_placeholder_token_id=42),
server_args=server_args,
_processor=_HFProcessor(),
transport_mode=None,
)
assert isinstance(processor._processor, wrapper_cls)
assert isinstance(worker_processor, wrapper_cls)
assert worker_processor is not processor._processor
if processor_cls is KimiK3ImageProcessor:
fingerprint_config = processor.preprocess_fingerprint_payload()[
"wrapped_processor"
]
assert isinstance(fingerprint_config, KimiK3PreprocessConfig)
assert fingerprint_config.patch_size == 14
finally:
processor.mm_processor_executor.shutdown()
processor.io_executor.shutdown()
processor.cpu_executor.shutdown()
try:
worker_processor = asyncio.run(
processor.mm_processor_executor.run(lambda *, processor: processor)
)
assert isinstance(processor._processor, wrapper_cls)
assert isinstance(worker_processor, wrapper_cls)
assert worker_processor is not processor._processor
if processor_cls is KimiK3ImageProcessor:
fingerprint_config = processor.preprocess_fingerprint_payload()[
"wrapped_processor"
]
assert isinstance(fingerprint_config, KimiK3PreprocessConfig)
assert fingerprint_config.patch_size == 14
finally:
processor.mm_processor_executor.shutdown()
processor.io_executor.shutdown()
processor.cpu_executor.shutdown()
def test_kimi_k3_expands_image_placeholders_with_original_dimensions():
@@ -844,82 +845,88 @@ def test_kimi_k3_normal_cache_path_connects_real_producer_to_model_consumer():
tokenizer_worker_num=1,
mm_preprocess_cache_size_mb=1,
)
processor = KimiK3ImageProcessor(
hf_config=hf_config,
server_args=server_args,
_processor=hf_processor,
transport_mode=None,
)
image = Image.new("RGB", (28, 28), color=(1, 2, 3))
encoded_image = io.BytesIO()
image.save(encoded_image, format="PNG")
image_data = ImageData(
url="data:image/png;base64,"
+ base64.b64encode(encoded_image.getvalue()).decode()
)
request = SimpleNamespace(video_data=None, mm_content_hashes=None)
class _Tower(nn.Module):
device = torch.device("cpu")
patch_size = 14
def __init__(self):
super().__init__()
self.patch_embed = SimpleNamespace(
proj=SimpleNamespace(weight=torch.empty(1, dtype=torch.float32))
)
def forward(self, pixel_values, _grid_thws):
return pixel_values
model = KimiK3ForConditionalGeneration.__new__(KimiK3ForConditionalGeneration)
nn.Module.__init__(model)
model.use_data_parallel = False
model.vision_tower = _Tower()
model.mm_projector = _Projector()
publish(server_args, role="tokenizer")
try:
with (
patch(
"sglang.srt.multimodal.processors.kimi_k3.is_cuda", return_value=True
),
patch.object(
processor,
"prepare_artifact_batch",
wraps=processor.prepare_artifact_batch,
) as prepare_artifacts,
):
cold = asyncio.run(
processor.process_mm_data_async([image_data], [1, 42, 2], request)
)
hot = asyncio.run(
processor.process_mm_data_async([image_data], [3, 42, 4], request)
)
cold_items = pickle.loads(pickle.dumps(cold.mm_items))
hot_items = pickle.loads(pickle.dumps(hot.mm_items))
processor = KimiK3ImageProcessor(
hf_config=hf_config,
server_args=server_args,
_processor=hf_processor,
transport_mode=None,
)
image = Image.new("RGB", (28, 28), color=(1, 2, 3))
encoded_image = io.BytesIO()
image.save(encoded_image, format="PNG")
image_data = ImageData(
url="data:image/png;base64,"
+ base64.b64encode(encoded_image.getvalue()).decode()
)
request = SimpleNamespace(video_data=None, mm_content_hashes=None)
with (
patch("sglang.srt.models.kimi_k3.configured_tp_size", return_value=1),
patch(
"sglang.srt.multimodal.processors.kimi_k25._gpu_preprocess_images",
return_value=(
torch.ones((4, 3), dtype=torch.float32),
torch.tensor([[1, 2, 2]], dtype=torch.int64),
class _Tower(nn.Module):
device = torch.device("cpu")
patch_size = 14
def __init__(self):
super().__init__()
self.patch_embed = SimpleNamespace(
proj=SimpleNamespace(weight=torch.empty(1, dtype=torch.float32))
)
def forward(self, pixel_values, _grid_thws):
return pixel_values
model = KimiK3ForConditionalGeneration.__new__(KimiK3ForConditionalGeneration)
nn.Module.__init__(model)
model.use_data_parallel = False
model.vision_tower = _Tower()
model.mm_projector = _Projector()
try:
with (
patch(
"sglang.srt.multimodal.processors.kimi_k3.is_cuda",
return_value=True,
),
),
):
cold_features = model.get_image_feature(cold_items)
hot_features = model.get_image_feature(hot_items)
finally:
processor.shutdown()
patch.object(
processor,
"prepare_artifact_batch",
wraps=processor.prepare_artifact_batch,
) as prepare_artifacts,
):
cold = asyncio.run(
processor.process_mm_data_async([image_data], [1, 42, 2], request)
)
hot = asyncio.run(
processor.process_mm_data_async([image_data], [3, 42, 4], request)
)
cold_items = pickle.loads(pickle.dumps(cold.mm_items))
hot_items = pickle.loads(pickle.dumps(hot.mm_items))
assert prepare_artifacts.call_count == 1
assert cold.mm_items[0].hash == hot.mm_items[0].hash
assert cold.mm_items[0].offsets == hot.mm_items[0].offsets == [(3, 3)]
assert (
cold_items[0].model_specific_data[DEFERRED_PREPROCESSING_KEY].backend == "gpu"
)
torch.testing.assert_close(cold_features, hot_features)
with (
patch("sglang.srt.models.kimi_k3.configured_tp_size", return_value=1),
patch(
"sglang.srt.multimodal.processors.kimi_k25._gpu_preprocess_images",
return_value=(
torch.ones((4, 3), dtype=torch.float32),
torch.tensor([[1, 2, 2]], dtype=torch.int64),
),
),
):
cold_features = model.get_image_feature(cold_items)
hot_features = model.get_image_feature(hot_items)
finally:
processor.shutdown()
assert prepare_artifacts.call_count == 1
assert cold.mm_items[0].hash == hot.mm_items[0].hash
assert cold.mm_items[0].offsets == hot.mm_items[0].offsets == [(3, 3)]
assert (
cold_items[0].model_specific_data[DEFERRED_PREPROCESSING_KEY].backend
== "gpu"
)
torch.testing.assert_close(cold_features, hot_features)
finally:
reset_context()
def test_kimi_k3_model_accepts_mixed_cached_eager_and_deferred_artifacts():
@@ -21,11 +21,13 @@ maybe_stub_sgl_kernel()
from sglang.srt.multimodal.processors.qwen_vl import ( # noqa: E402
QwenVLImageProcessor,
)
from sglang.srt.runtime_context import publish, reset_context # noqa: E402
from sglang.srt.server_args import ServerArgs # noqa: E402
register_cpu_ci(est_time=0, suite="base-a-test-cpu", disabled="Qwen test fixtures")
def make_processor(config, image_processor_cls=None):
def make_processor(case, config, image_processor_cls=None):
"""A ``QwenVLImageProcessor`` over a tiny hand-built tokenizer.
``image_processor_cls`` picks the HF backend; they resample differently."""
image_processor_cls = image_processor_cls or HfQwenImageProcessor
@@ -87,6 +89,18 @@ def make_processor(config, image_processor_cls=None):
allowed_media_domains=[],
media_url_max_file_size_mb=64,
)
# The processor reads its media policy, transport and per-modality limits
# from the mm bag, so the fixture publishes before building it.
publish(
ServerArgs(
model_path="dummy",
mm_feature_transport=server_args.mm_feature_transport,
mm_process_config=server_args.mm_process_config,
allowed_media_domains=server_args.allowed_media_domains,
),
role="tokenizer",
)
case.addCleanup(reset_context)
return QwenVLImageProcessor(
hf_config, server_args, processor, None, skip_mm_pool=True
)
@@ -50,7 +50,9 @@ class TestQwenE2eParity(CustomTestCase):
import transformers.models.qwen2_vl as qwen2_vl
self.processor = make_processor(
PROCESSOR_CONFIGS["qwen2_5_vl"], getattr(qwen2_vl, self.image_processor)
self,
PROCESSOR_CONFIGS["qwen2_5_vl"],
getattr(qwen2_vl, self.image_processor),
)
def tearDown(self):
@@ -50,7 +50,7 @@ class TestQwenNativeMmHashes(CustomTestCase):
from sglang.srt.managers.multimodal_processor import import_processors
import_processors("sglang.srt.multimodal.processors")
self.processor = make_processor(PROCESSOR_CONFIGS["qwen2_5_vl"])
self.processor = make_processor(self, PROCESSOR_CONFIGS["qwen2_5_vl"])
def tearDown(self):
self.processor.io_executor.shutdown()
@@ -42,6 +42,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
def test_model_class_controls_cuda_vmm_opt_in(self):
from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.runtime_context import get_context
class SupportedModel:
supports_cuda_vmm_feature_transport = True
@@ -49,8 +50,10 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
class UnsupportedModel:
pass
override = get_context().override_server_args(mm_feature_transport="cuda_vmm")
override.install()
self.addCleanup(override.restore)
manager = object.__new__(TokenizerManager)
manager.server_args = SimpleNamespace(mm_feature_transport="cuda_vmm")
manager.model_config = object()
with patch(
@@ -70,9 +73,12 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
def test_cpu_transport_skips_model_opt_in_lookup(self):
from sglang.srt.managers.tokenizer_manager import TokenizerManager
from sglang.srt.runtime_context import get_context
override = get_context().override_server_args(mm_feature_transport="cpu")
override.install()
self.addCleanup(override.restore)
manager = object.__new__(TokenizerManager)
manager.server_args = SimpleNamespace(mm_feature_transport="cpu")
manager.model_config = object()
with patch(
@@ -21,6 +21,7 @@ from sglang.srt.multimodal.cache import (
resolve_multimodal_item_hash,
snapshot_media,
)
from sglang.srt.runtime_context import publish, reset_context
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
@@ -216,26 +217,32 @@ class TestMediaIdentity(unittest.TestCase):
return {"model_type": "vlm", "architectures": ["VLM"]}
config = Config()
args = ServerArgs(
model_path="dummy",
revision="model-revision",
disable_fast_image_processor=False,
mm_process_config={"image": {"max_pixels": 1024}},
)
base = build_processor_fingerprint(Processor("gpu"), config, args)
changed_backend = build_processor_fingerprint(Processor("cpu"), config, args)
changed_args = ServerArgs(
model_path="dummy",
revision="model-revision",
disable_fast_image_processor=False,
mm_process_config={"image": {"max_pixels": 2048}},
)
changed_config = build_processor_fingerprint(
Processor("gpu"), config, changed_args
)
def fingerprint(processor, mm_process_config):
# The digest reads the effective config, so the test publishes it
# rather than handing one in: that is the only source the function
# has, and two callers with the same effective config must agree.
publish(
ServerArgs(
model_path="dummy",
revision="model-revision",
disable_fast_image_processor=False,
mm_process_config=mm_process_config,
),
role="test",
)
return build_processor_fingerprint(processor, config)
self.addCleanup(reset_context)
small = {"image": {"max_pixels": 1024}}
base = fingerprint(Processor("gpu"), small)
changed_backend = fingerprint(Processor("cpu"), small)
changed_config = fingerprint(Processor("gpu"), {"image": {"max_pixels": 2048}})
same_again = fingerprint(Processor("gpu"), small)
self.assertNotEqual(base, changed_backend)
self.assertNotEqual(base, changed_config)
self.assertEqual(base, same_again)
def test_item_hash_namespace_covers_identity_and_processor_output(self):
digest = snapshot_media(b"image").content_digest
@@ -2161,6 +2161,7 @@ class TestGrpcServerArgs(CustomTestCase):
tokenizer_manager=MagicMock(),
template_manager=MagicMock(),
scheduler_info={},
grpc_port=server_args.grpc_port,
)
self.assertEqual(handle, "handle")
@@ -52,6 +52,23 @@ _SLOT_OWNERS = ("srt/runtime_context.py", "srt/server_args.py", "srt/arg_groups/
# live topology cannot answer there. The test below asserts this map is exactly
# the set of call sites, so the reasons cannot drift away from the code.
_CONFIGURED_SIZE_CALL_SITES = {
("srt/entrypoints/engine.py", "configured_pp_size"): (
"the launch path decides how many scheduler processes to spawn; it runs "
"before any of them exists, so there is no group to ask"
),
("srt/ray/engine.py", "configured_pp_size"): (
"the Ray driver sizes the actor placement group; the actors it is about "
"to create are the ones that will hold the process groups"
),
("srt/ray/data_parallel_controller.py", "configured_pp_size"): (
"same placement arithmetic on the DP path -- ranks per TP group, "
"computed in the driver before the actors start"
),
("srt/ray/data_parallel_controller.py", "configured_attn_cp_size"): (
"the attention-CP factor of that same placement arithmetic, and the one "
"size whose live value cannot express the configured intent when "
"attn_cp_size > moe_dp_size aliases the groups"
),
("srt/layers/attention/dsa/dsa_indexer.py", "configured_pp_size"): (
"gates `pp_size > 1 and not get_pp_group()...`; the short circuit is the "
"point, since with PP off the group is never touched, which is what lets "
@@ -91,6 +108,10 @@ _CONFIGURED_SIZE_CALL_SITES = {
"the consumer count is configured fan-out arithmetic (tp_size // "
"dp_size), which is what the record answered before"
),
("srt/disaggregation/encode_server.py", "configured_tp_size"): (
"the encode server's launch entry sizes its workers before it has "
"spawned any of them"
),
("srt/model_loader/loader.py", "configured_moe_dp_size"): (
"the same dict already carries the live moe_dp_size under 'dp'; this entry "
"is the configured intent"
@@ -0,0 +1,280 @@
"""Launch paths read the configured parallel sizes, not the live ones.
`get_parallel().pp_size` and its four siblings are read-through properties over
the process groups, so they answer only after distributed init. The launcher
decides how many processes to spawn *before* that, and a live read there raises
`Distributed environment is not initialized` -- a startup crash no unit test
reaches, because nothing short of booting a server runs the launcher.
"""
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=9, suite="base-a-test-cpu")
_PACKAGE_ROOT = pathlib.Path(sglang.__file__).resolve().parent
# Live-shadowed sizes a launch path is known to have read. ParallelContext
# shadows more properties than these (every `_v(name, ...)` one raises the same
# "Distributed environment is not initialized"); this dict carries the ones a
# `configured_*` accessor answers, so it is a remedy map, not a census.
_LIVE_SHADOWED = {
"tp_size": "configured_tp_size()",
"pp_size": "configured_pp_size()",
"moe_dp_size": "configured_moe_dp_size()",
"attn_cp_size": "configured_attn_cp_size()",
"dcp_size": "a configured accessor (none exists yet; add one beside configured_pp_size)",
}
# Launch paths that decide how many children to spawn are derived below
# from the spawn itself. These launch without a size-driven spawn, so no
# derivation reaches them and they are carried by hand.
_HAND_CARRIED = (
"srt/entrypoints/http_server.py",
"srt/entrypoints/sidecar.py",
"srt/ray/data_parallel_controller.py",
"srt/ray/engine.py",
"srt/ray/http_server.py",
)
def _multiprocessing_names(tree):
"""Names bound to multiprocessing, to one of its start contexts, or to the
process constructors themselves."""
modules, constructors = set(), set()
for node in ast.walk(tree):
if isinstance(node, ast.Import):
for a in node.names:
if a.name == "multiprocessing" or a.name.startswith("multiprocessing."):
modules.add(a.asname or a.name.split(".")[0])
elif a.name == "torch.multiprocessing":
modules.add(a.asname or "torch")
elif isinstance(node, ast.ImportFrom):
if node.module in (
"multiprocessing",
"multiprocessing.context",
"torch.multiprocessing",
):
constructors |= {
a.asname or a.name for a in node.names if a.name == "Process"
}
elif node.module == "concurrent.futures":
constructors |= {
a.asname or a.name
for a in node.names
if a.name == "ProcessPoolExecutor"
}
for node in ast.walk(tree):
if isinstance(node, ast.Assign) and isinstance(node.value, ast.Call):
func = node.value.func
if (
isinstance(func, ast.Attribute)
and func.attr == "get_context"
and isinstance(func.value, ast.Name)
and func.value.id in modules
):
modules |= {t.id for t in node.targets if isinstance(t, ast.Name)}
return modules, constructors
def _configured_accessors() -> frozenset:
"""The `configured_*_size()` names `runtime_context` exports.
Derived from that module, so a new accessor keeps its launcher watched
without a second list here.
"""
tree = ast.parse(
(_PACKAGE_ROOT / "srt/runtime_context.py").read_text(encoding="utf-8-sig")
)
names = frozenset(
node.name
for node in tree.body
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef))
and node.name.startswith("configured_")
and node.name.endswith("_size")
)
assert names, (
"no configured_*_size accessors found in runtime_context; the "
"derivation is broken, not the tree"
)
return names
def _spawns_from_a_size(tree) -> bool:
"""Does any function here construct a child process *and* read one of the
five sizes -- live off the parallel bag, or through its `configured_*_size()`
answer? That is a spawn count decided from the topology.
Counting the configured read too is what keeps a launcher watched after it
is converted. Deriving on the live read alone means the file drops out of
the scan the moment it stops offending, so the guard would only ever watch
the launchers that already fail it.
"""
configured = _configured_accessors()
modules, constructors = _multiprocessing_names(tree)
names, qualified = _parallel_bag_names(tree)
aliases = _bag_aliases(tree, names, qualified)
for fn in ast.walk(tree):
if not isinstance(fn, (ast.FunctionDef, ast.AsyncFunctionDef)):
continue
spawns = reads = False
for node in ast.walk(fn):
if isinstance(node, ast.Call):
func = node.func
if isinstance(func, ast.Attribute) and func.attr in (
"Process",
"ProcessPoolExecutor",
"Popen",
"spawn",
):
# `mp.Process(`, `mp.get_context("spawn").Process(` and
# `subprocess.Popen(` all reach a child process; the
# receiver of a chained call is itself a call, so this
# cannot require a bare Name.
spawns = True
elif isinstance(func, ast.Name) and func.id in constructors:
spawns = True
if (isinstance(func, ast.Name) and func.id in configured) or (
isinstance(func, ast.Attribute) and func.attr in configured
):
reads = True
elif (
isinstance(node, ast.Attribute)
and node.attr in _LIVE_SHADOWED
and (
_is_parallel_bag_call(node.value, names, qualified)
or (isinstance(node.value, ast.Name) and node.value.id in aliases)
)
):
# A record read (`server_args.tp_size`) sizes a spawn too, but
# it cannot raise pre-dist; only the bag read is this guard's
# subject, so only it forces a module into _PRE_DIST.
reads = True
if spawns and reads:
return True
return False
def _parallel_bag_names(tree):
"""What this module calls `get_parallel`, plus any runtime_context alias.
A literal-name match reads only one spelling; an aliased import or a
module-qualified call is the same read with a different surface.
"""
names, modules = set(), set()
for node in ast.walk(tree):
if (
isinstance(node, ast.ImportFrom)
and node.module
and node.module.endswith("runtime_context")
):
names |= {
a.asname or a.name for a in node.names if a.name == "get_parallel"
}
elif isinstance(node, ast.Import):
for a in node.names:
if a.name.endswith("runtime_context"):
modules.add(a.asname or a.name.split(".")[0])
return names, modules
def _is_parallel_bag_call(node, names, modules) -> bool:
if not isinstance(node, ast.Call):
return False
if isinstance(node.func, ast.Name):
return node.func.id in names
return (
isinstance(node.func, ast.Attribute)
and node.func.attr == "get_parallel"
and isinstance(node.func.value, ast.Name)
and node.func.value.id in modules
)
def _bag_aliases(tree, names, qualified):
"""Locals bound to the parallel bag: `p = get_parallel()` then `p.pp_size`
is the same read one line later."""
return {
target.id
for node in ast.walk(tree)
if isinstance(node, ast.Assign)
and _is_parallel_bag_call(node.value, names, qualified)
for target in node.targets
if isinstance(target, ast.Name)
}
def _launch_paths():
"""(relative path, tree) per module that runs before its process groups.
A module that sizes a spawn loop from a parallel-bag size is derived from
the spawn itself; `_HAND_CARRIED` holds the launch entries that spawn
nothing, which no derivation can reach.
"""
seen = {}
for path in sorted(_PACKAGE_ROOT.rglob("*.py")):
source = path.read_text()
# Every spawn shape below names Process, ProcessPoolExecutor or Popen.
if not any(name in source for name in ("Process", "Popen", "spawn")):
continue
try:
tree = ast.parse(source)
except SyntaxError:
continue
if _spawns_from_a_size(tree):
seen[str(path.relative_to(_PACKAGE_ROOT))] = tree
for rel in _HAND_CARRIED:
seen.setdefault(rel, ast.parse((_PACKAGE_ROOT / rel).read_text()))
return sorted(seen.items())
class TestLaunchPathsReadConfiguredSizes(CustomTestCase):
def test_no_live_topology_read_before_distributed_init(self):
offenders = []
for rel, tree in _launch_paths():
names, modules = _parallel_bag_names(tree)
aliases = _bag_aliases(tree, names, modules)
for node in ast.walk(tree):
if isinstance(node, ast.Attribute) and node.attr in _LIVE_SHADOWED:
base = node.value
if _is_parallel_bag_call(base, names, modules) or (
isinstance(base, ast.Name) and base.id in aliases
):
offenders.append(
f"{rel}:{node.lineno} reads the live {node.attr}; "
f"use {_LIVE_SHADOWED[node.attr]}"
)
elif (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Name)
and node.func.id == "getattr"
and len(node.args) >= 2
and isinstance(node.args[1], ast.Constant)
and node.args[1].value in _LIVE_SHADOWED
and (
_is_parallel_bag_call(node.args[0], names, modules)
or (
isinstance(node.args[0], ast.Name)
and node.args[0].id in aliases
)
)
):
offenders.append(
f"{rel}:{node.lineno} reads the live "
f"{node.args[1].value} through getattr; "
f"use {_LIVE_SHADOWED[node.args[1].value]}"
)
self.assertEqual(
offenders,
[],
"launch paths run before distributed init:\n " + "\n ".join(offenders),
)
if __name__ == "__main__":
unittest.main()
@@ -104,9 +104,6 @@ _UNREAD_ENTRIES: dict = {
("multimodal_gen/test/unit/test_disagg_trace.py", "_srt_trace_server_args"): (
"a trace fixture publishing its own context"
),
("srt/managers/detokenizer_manager.py", "run_detokenizer_process"): (
"DetokenizerManager reads the handed instance at this revision"
),
}
# `publish` itself and its named wrappers live here; a call inside them is the
@@ -49,6 +49,7 @@ _PAIR_READERS = {
"batch_overlap/two_batch_overlap.py": "prefill (extend positions)",
"managers/scheduler.py": "prefill (truncation align knobs)",
"entrypoints/engine.py": "either half (flashinfer version floor)",
"models/sarvam_moe.py": "the half serving the forward (attn dispatch)",
}
@@ -129,6 +129,7 @@ _ENV_MATRIX = (({}, {"SGLANG_IS_IN_CI": "true"}),)
_PASSED = frozenset({"model_path", "device", "random_seed"})
_EXPOSED = {
("entrypoints/sidecar.py", "grpc_port"),
("configs/embedding_model_spec.py", "chunked_prefill_size"),
("configs/embedding_model_spec.py", "cuda_graph_config"),
("configs/embedding_model_spec.py", "disable_radix_cache"),
@@ -143,24 +144,10 @@ _EXPOSED = {
("configs/model_config.py", "quantization"),
("configs/model_config.py", "speculative_algorithm"),
("configs/model_config.py", "speculative_draft_model_quantization"),
("constrained/base_grammar_backend.py", "grammar_backend"),
("constrained/base_grammar_backend.py", "reasoning_parser"),
("disaggregation/common/conn.py", "disaggregation_bootstrap_port"),
("disaggregation/common/conn.py", "pp_size"),
("disaggregation/decode_kvcache_offload_manager.py", "hicache_io_backend"),
("disaggregation/decode_kvcache_offload_manager.py", "served_model_name"),
("disaggregation/encode_receiver.py", "disaggregation_ib_device"),
("disaggregation/encode_receiver.py", "encoder_transfer_backend"),
("disaggregation/encode_receiver.py", "mooncake_ib_device"),
("disaggregation/encode_receiver.py", "tokenizer_path"),
("disaggregation/encode_server.py", "allowed_media_domains"),
("disaggregation/encode_server.py", "device"),
("disaggregation/encode_server.py", "encoder_transfer_backend"),
("disaggregation/encode_server.py", "load_format"),
("disaggregation/encode_server.py", "mm_process_config"),
("disaggregation/encode_server.py", "model_path"),
("disaggregation/encode_server.py", "served_model_name"),
("disaggregation/encode_server.py", "tokenizer_path"),
("disaggregation/utils.py", "disaggregation_transfer_backend"),
("distributed/bootstrap.py", "disable_custom_all_reduce"),
("distributed/bootstrap.py", "enable_symm_mem"),
@@ -201,28 +188,14 @@ _EXPOSED = {
("elastic_ep/expert_backup_manager.py", "load_format"),
("elastic_ep/expert_backup_manager.py", "mooncake_ib_device"),
("entrypoints/engine.py", "attn_cp_size"),
("entrypoints/engine.py", "detokenizer_worker_num"),
("entrypoints/engine.py", "dtype"),
("entrypoints/engine.py", "enable_symm_mem"),
("entrypoints/engine.py", "ep_join_mode"),
("entrypoints/engine.py", "load_format"),
("entrypoints/engine.py", "model_path"),
("entrypoints/engine.py", "moe_dp_size"),
("entrypoints/engine.py", "pp_size"),
("entrypoints/engine.py", "quantization"),
("entrypoints/engine.py", "reasoning_parser"),
(
"entrypoints/engine.py",
"remote_instance_weight_loader_start_seed_via_transfer_engine",
),
("entrypoints/engine.py", "tool_call_parser"),
("entrypoints/http_server.py", "disaggregation_mode"),
("entrypoints/http_server.py", "ep_join_mode"),
("entrypoints/http_server.py", "grpc_port"),
("entrypoints/http_server.py", "model_path"),
("entrypoints/http_server.py", "served_model_name"),
("entrypoints/http_server.py", "skip_server_warmup"),
("entrypoints/sidecar.py", "grpc_port"),
("eplb/eplb_manager.py", "ep_dispatch_algorithm"),
("eplb/eplb_manager.py", "expert_distribution_recorder_buffer_size"),
("eplb/expert_distribution.py", "deepep_mode"),
@@ -236,7 +209,6 @@ _EXPOSED = {
("kv_canary/capacities.py", "chunked_prefill_size"),
("kv_canary/capacities.py", "cuda_graph_config"),
("kv_canary/capacities.py", "speculative_num_draft_tokens"),
("kv_canary/token_oracle/install.py", "sampling_backend"),
("layers/cp/base.py", "attn_cp_size"),
("layers/cp/base.py", "cp_strategy"),
("layers/cp/base.py", "enable_prefill_cp"),
@@ -259,9 +231,6 @@ _EXPOSED = {
("managers/data_parallel_controller.py", "moe_dp_size"),
("managers/data_parallel_controller.py", "pp_size"),
("managers/data_parallel_controller.py", "soft_watchdog_timeout"),
("managers/detokenizer_manager.py", "soft_watchdog_timeout"),
("managers/detokenizer_manager.py", "tokenizer_path"),
("managers/detokenizer_manager.py", "tool_call_parser"),
("managers/disagg_service.py", "disaggregation_bootstrap_port"),
("managers/disagg_service.py", "disaggregation_mode"),
("managers/disagg_service.py", "disaggregation_transfer_backend"),
@@ -279,27 +248,7 @@ _EXPOSED = {
("managers/scheduler.py", "pp_size"),
("managers/scheduler.py", "soft_watchdog_timeout"),
("managers/scheduler.py", "speculative_algorithm"),
(
"managers/scheduler_components/new_token_ratio_tracker.py",
"schedule_conservativeness",
),
("managers/tokenizer_manager.py", "disable_radix_cache"),
("managers/tokenizer_manager.py", "disaggregation_mode"),
("managers/tokenizer_manager.py", "disaggregation_transfer_backend"),
("managers/tokenizer_manager.py", "enable_lora"),
("managers/tokenizer_manager.py", "enable_tokenizer_batch_encode"),
("managers/tokenizer_manager.py", "encoder_transfer_backend"),
("managers/tokenizer_manager.py", "limit_mm_data_per_request"),
("managers/tokenizer_manager.py", "lora_paths"),
("managers/tokenizer_manager.py", "mm_feature_transport"),
("managers/tokenizer_manager.py", "model_path"),
("managers/tokenizer_manager.py", "preferred_sampling_params"),
("managers/tokenizer_manager.py", "return_hidden_states_mode"),
("managers/tokenizer_manager.py", "served_model_name"),
("managers/tokenizer_manager.py", "soft_watchdog_timeout"),
("managers/tokenizer_manager.py", "speculative_algorithm"),
("managers/tokenizer_manager.py", "speculative_num_draft_tokens"),
("managers/tokenizer_manager.py", "tokenizer_path"),
("managers/tp_worker.py", "disable_overlap_schedule"),
("managers/tp_worker.py", "model_path"),
("managers/tp_worker.py", "random_seed"),
@@ -313,8 +262,6 @@ _EXPOSED = {
("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "served_model_name"),
("mem_cache/kv_cache_builder.py", "hicache_mem_layout"),
("mem_cache/radix_cache_cpp.py", "enable_hierarchical_cache"),
("model_executor/forward_batch_info.py", "enable_return_hidden_states"),
("model_executor/forward_batch_info.py", "return_hidden_states_mode"),
("model_executor/model_runner.py", "device"),
("model_executor/model_runner.py", "speculative_algorithm"),
("model_executor/model_runner.py", "speculative_draft_attention_backend"),
@@ -353,22 +300,11 @@ _EXPOSED = {
"model_executor/runner_backend/tc_piecewise_cuda_graph_backend.py",
"cuda_graph_config",
),
("models/sarvam_moe.py", "attention_backend"),
("models/sarvam_moe.py", "decode_attention_backend"),
("models/sarvam_moe.py", "prefill_attention_backend"),
("multimodal/cache/identity.py", "mm_process_config"),
("multimodal/processors/base_processor.py", "allowed_media_domains"),
("multimodal/processors/base_processor.py", "image_processor_backend"),
("multimodal/processors/base_processor.py", "mm_feature_transport"),
("multimodal/processors/base_processor.py", "mm_process_config"),
("multimodal/processors/mimo_v2.py", "device"),
("observability/metrics_collector.py", "disaggregation_mode"),
("observability/metrics_collector.py", "prefill_delayer_max_delay_passes"),
("observability/metrics_collector.py", "served_model_name"),
("parser/template_detection.py", "model_path"),
("ray/data_parallel_controller.py", "attn_cp_size"),
("ray/data_parallel_controller.py", "pp_size"),
("ray/engine.py", "pp_size"),
("speculative/adaptive_spec_params.py", "speculative_algorithm"),
("speculative/adaptive_spec_params.py", "speculative_eagle_topk"),
("speculative/dflash_worker_v2.py", "speculative_draft_window_size"),
@@ -435,29 +371,20 @@ _EXPOSED_CUDA_ONLY: frozenset = frozenset()
_OVERRIDDEN_AND_READ = {
("configs/model_config.py", "dtype"),
("configs/model_config.py", "model_path"),
("constrained/base_grammar_backend.py", "grammar_backend"),
("disaggregation/decode_kvcache_offload_manager.py", "hicache_storage_backend"),
(
"disaggregation/decode_kvcache_offload_manager.py",
"hicache_storage_backend_extra_config",
),
("disaggregation/encode_server.py", "load_format"),
("disaggregation/encode_server.py", "model_path"),
(
"distributed/device_communicators/mooncake_transfer_engine.py",
"hicache_storage_backend",
),
("dllm/config.py", "model_path"),
("elastic_ep/expert_backup_manager.py", "load_format"),
("entrypoints/engine.py", "dtype"),
("entrypoints/engine.py", "load_format"),
("entrypoints/engine.py", "model_path"),
("entrypoints/http_server.py", "model_path"),
("kv_canary/api.py", "speculative_num_steps"),
("kv_canary/capacities.py", "speculative_num_draft_tokens"),
("managers/scheduler.py", "hicache_storage_backend"),
("managers/tokenizer_manager.py", "model_path"),
("managers/tokenizer_manager.py", "speculative_num_draft_tokens"),
("managers/tp_worker.py", "model_path"),
("mem_cache/hiradix_cache.py", "hicache_storage_backend"),
("mem_cache/hiradix_cache.py", "hicache_storage_backend_extra_config"),