config: the per-instance families read the bags (#35026)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(),
|
||||
)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"),
|
||||
|
||||
Reference in New Issue
Block a user