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 |
|
| 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 |
|
| 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 |
|
| 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()` |
|
| 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 |
|
| 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
|
the scope. When there is a runner in hand, read its
|
||||||
stamp; that is a different rule from "read the instance".
|
stamp; that is a different rule from "read the instance".
|
||||||
- **Per-instance boundaries** — the tokenizer-manager family, everything under
|
- **Per-instance boundaries** — the tokenizer-manager family, everything under
|
||||||
`entrypoints/`, and the tokenizer-process multimodal processors still read
|
`entrypoints/`, and the tokenizer-process multimodal processors read the bags.
|
||||||
`self.server_args` today. The old justification ("several `Engine`s can share
|
The old justification for keeping them on `self.server_args` ("several
|
||||||
one process, bags are last-publish-wins across them") is **retracted** — owner
|
`Engine`s can share one process, bags are last-publish-wins across them") is
|
||||||
ruling (2026-08-15): a process holds at most one live config at a time
|
**retracted** — owner ruling (2026-08-15): a process holds at most one live
|
||||||
(concurrent multi-Engine is unsupported; sequential rebuild stays legal, unit
|
config at a time (concurrent multi-Engine is unsupported; sequential rebuild
|
||||||
tests rely on it). These reads are scheduled to become bag reads in the
|
stays legal, unit tests rely on it). What still reads the instance in those
|
||||||
bag-read series; treat them as pinned debt, not as a boundary to imitate. What
|
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:
|
genuinely stays per-instance is what differs per *worker* within one engine:
|
||||||
`base_gpu_id` travels as a constructor argument (`MMEncoder(gpu_id=...)`;
|
`base_gpu_id` travels as a constructor argument (`MMEncoder(gpu_id=...)`;
|
||||||
`BaseMultimodalProcessor._fast_image_processor_device` is the shape to copy).
|
`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 —
|
"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
|
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`;
|
copies to read any more). The allow-list is `GrammarManager` and `MMEncoder`;
|
||||||
the tokenizer-manager family, `entrypoints/`, and the tokenizer-process
|
what sits beside it is residue, not a family — and not for one single reason:
|
||||||
multimodal processors sit beside it only as pinned debt — not for one single
|
|
||||||
reason:
|
|
||||||
|
|
||||||
- the tokenizer-manager family and `entrypoints/` are **pinned debt awaiting
|
- the tokenizer-manager family and `entrypoints/` **read the bags**; what is
|
||||||
conversion to bag reads** (the old multi-Engine justification is retracted —
|
left of them in the exposure ratchet is a handful of individually-dispositioned
|
||||||
one process, one live config); the reads still work today because the
|
pairs, not a family awaiting conversion. Read the ratchet for the current set
|
||||||
instance carries resolved values;
|
rather than assuming a directory is off-limits;
|
||||||
- `GrammarManager` is a handed instance — it is constructed with the config its
|
- `GrammarManager` is a handed instance for its residual `self.server_args`
|
||||||
owner hands it and never assumes a published namespace;
|
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,
|
- `MMEncoder` publishes the very instance it is handed (`publish(server_args,
|
||||||
role="encoder")`) and takes its per-worker device as a separate `gpu_id`
|
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
|
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 —
|
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.
|
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
|
Their tests are not one story: a `GrammarManager` built standalone turns the
|
||||||
bag read turns into "config namespace not published".
|
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
|
**Test doubles publish, they do not inject.** A stand-in that carries
|
||||||
`server_args=SimpleNamespace(field=...)` stops working the moment production reads
|
`server_args=SimpleNamespace(field=...)` stops working the moment production reads
|
||||||
|
|||||||
@@ -22,7 +22,12 @@ from typing import Dict, List, NamedTuple, Optional, Tuple
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.parser.reasoning_parser import ReasoningParser
|
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
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -315,7 +320,7 @@ def create_grammar_backend(
|
|||||||
eos_token_ids: Optional[set] = None,
|
eos_token_ids: Optional[set] = None,
|
||||||
think_end_ids: Optional[List[int]] = None,
|
think_end_ids: Optional[List[int]] = None,
|
||||||
) -> Optional[BaseGrammarBackend]:
|
) -> Optional[BaseGrammarBackend]:
|
||||||
name = server_args.grammar_backend
|
name = get_exec().kernel.grammar_backend
|
||||||
|
|
||||||
# Custom grammar backend has the highest priority
|
# Custom grammar backend has the highest priority
|
||||||
if name in GRAMMAR_BACKEND_REGISTRY:
|
if name in GRAMMAR_BACKEND_REGISTRY:
|
||||||
@@ -329,7 +334,7 @@ def create_grammar_backend(
|
|||||||
|
|
||||||
grammar_backend = OutlinesGrammarBackend(
|
grammar_backend = OutlinesGrammarBackend(
|
||||||
tokenizer,
|
tokenizer,
|
||||||
whitespace_pattern=server_args.constrained_json_whitespace_pattern,
|
whitespace_pattern=get_serving().constrained_json_whitespace_pattern,
|
||||||
)
|
)
|
||||||
elif name == "xgrammar":
|
elif name == "xgrammar":
|
||||||
from sglang.srt.constrained.xgrammar_backend import (
|
from sglang.srt.constrained.xgrammar_backend import (
|
||||||
@@ -345,10 +350,10 @@ def create_grammar_backend(
|
|||||||
tokenizer,
|
tokenizer,
|
||||||
vocab_size=vocab_size,
|
vocab_size=vocab_size,
|
||||||
model_eos_token_ids=eos_list,
|
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:
|
except TokenizerNotSupportedError as e:
|
||||||
if server_args.enable_strict_thinking:
|
if get_serving().enable_strict_thinking:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"--enable-strict-thinking requires a grammar backend with "
|
f"--enable-strict-thinking requires a grammar backend with "
|
||||||
f"token filtering support, but XGrammar failed to initialize: "
|
f"token filtering support, but XGrammar failed to initialize: "
|
||||||
@@ -367,13 +372,13 @@ def create_grammar_backend(
|
|||||||
|
|
||||||
grammar_backend = GuidanceBackend(
|
grammar_backend = GuidanceBackend(
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
any_whitespace=not server_args.constrained_json_disable_any_whitespace,
|
any_whitespace=not get_serving().constrained_json_disable_any_whitespace,
|
||||||
whitespace_pattern=server_args.constrained_json_whitespace_pattern,
|
whitespace_pattern=get_serving().constrained_json_whitespace_pattern,
|
||||||
n_vocab=vocab_size,
|
n_vocab=vocab_size,
|
||||||
eos_token_ids=eos_token_ids,
|
eos_token_ids=eos_token_ids,
|
||||||
)
|
)
|
||||||
elif name == "none":
|
elif name == "none":
|
||||||
if server_args.enable_strict_thinking:
|
if get_serving().enable_strict_thinking:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"--enable-strict-thinking requires a grammar backend that supports "
|
"--enable-strict-thinking requires a grammar backend that supports "
|
||||||
"token filtering, but grammar_backend='none' was specified. Use "
|
"token filtering, but grammar_backend='none' was specified. Use "
|
||||||
@@ -384,13 +389,13 @@ def create_grammar_backend(
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"Invalid grammar backend: {name}")
|
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 (
|
from sglang.srt.constrained.reasoner_grammar_backend import (
|
||||||
ReasonerGrammarBackend,
|
ReasonerGrammarBackend,
|
||||||
)
|
)
|
||||||
|
|
||||||
reasoning_parser = ReasoningParser(
|
reasoning_parser = ReasoningParser(
|
||||||
model_type=server_args.reasoning_parser,
|
model_type=get_serving().reasoning_parser,
|
||||||
stream_reasoning=False,
|
stream_reasoning=False,
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
)
|
)
|
||||||
@@ -399,7 +404,7 @@ def create_grammar_backend(
|
|||||||
grammar_backend,
|
grammar_backend,
|
||||||
reasoning_parser,
|
reasoning_parser,
|
||||||
tokenizer,
|
tokenizer,
|
||||||
enable_strict_thinking=server_args.enable_strict_thinking,
|
enable_strict_thinking=get_serving().enable_strict_thinking,
|
||||||
)
|
)
|
||||||
|
|
||||||
return grammar_backend
|
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.multimodal_processor import get_mm_processor, import_processors
|
||||||
from sglang.srt.managers.schedule_batch import Modality, Req
|
from sglang.srt.managers.schedule_batch import Modality, Req
|
||||||
from sglang.srt.multimodal.cache import media_preprocess_kwargs
|
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.server_args import ServerArgs
|
||||||
from sglang.srt.utils import ImageData
|
from sglang.srt.utils import ImageData
|
||||||
from sglang.srt.utils.common import safe_pickle_loads
|
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
|
# context alive for the process instead of creating a temporary context
|
||||||
# whose destruction also closes its per-request socket.
|
# whose destruction also closes its per-request socket.
|
||||||
self.scheduler_context = zmq.Context()
|
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`
|
# When ``encode_urls`` is shared with an :class:`EncoderBootstrapServer`
|
||||||
# (tokenizer manager process), it grows / shrinks in place as encoders
|
# (tokenizer manager process), it grows / shrinks in place as encoders
|
||||||
# register or unregister; the receiver always sees the current set.
|
# register or unregister; the receiver always sees the current set.
|
||||||
@@ -1642,8 +1643,8 @@ class MMReceiverBase(ABC):
|
|||||||
self.embeddings_engine = init_mooncake_transfer_engine(
|
self.embeddings_engine = init_mooncake_transfer_engine(
|
||||||
hostname=self.host,
|
hostname=self.host,
|
||||||
ib_device=(
|
ib_device=(
|
||||||
server_args.disaggregation_ib_device
|
get_disagg().disaggregation_ib_device
|
||||||
or server_args.mooncake_ib_device
|
or get_exec().moe.mooncake_ib_device
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
self.embeddings_buffer = dict()
|
self.embeddings_buffer = dict()
|
||||||
@@ -1689,7 +1690,7 @@ class MMReceiverBase(ABC):
|
|||||||
extra_kwargs["tokenizer_backend"] = server_args.tokenizer_backend
|
extra_kwargs["tokenizer_backend"] = server_args.tokenizer_backend
|
||||||
|
|
||||||
_processor = get_processor(
|
_processor = get_processor(
|
||||||
server_args.tokenizer_path,
|
get_serving().tokenizer_path,
|
||||||
tokenizer_mode=server_args.tokenizer_mode,
|
tokenizer_mode=server_args.tokenizer_mode,
|
||||||
trust_remote_code=server_args.trust_remote_code,
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
revision=server_args.revision,
|
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 (
|
from sglang.srt.model_executor.model_runner_components.load_model_utils import (
|
||||||
maybe_precompile_model_kernels_after_loading,
|
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.cache import parse_content_hash, snapshot_media
|
||||||
from sglang.srt.multimodal.encoder_preprocessing import (
|
from sglang.srt.multimodal.encoder_preprocessing import (
|
||||||
EncoderPreprocessOutput,
|
EncoderPreprocessOutput,
|
||||||
@@ -73,10 +73,15 @@ from sglang.srt.observability.trace import (
|
|||||||
trace_set_thread_info,
|
trace_set_thread_info,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
configured_tp_size,
|
||||||
|
get_device,
|
||||||
get_disagg,
|
get_disagg,
|
||||||
get_exec,
|
get_exec,
|
||||||
get_mm,
|
get_mm,
|
||||||
|
get_model,
|
||||||
|
get_observability,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
|
get_serving,
|
||||||
publish,
|
publish,
|
||||||
)
|
)
|
||||||
from sglang.srt.server_args import (
|
from sglang.srt.server_args import (
|
||||||
@@ -315,7 +320,7 @@ class MMEncoder:
|
|||||||
logger.info(f"init MMEncoder {rank}/{server_args.tp_size}")
|
logger.info(f"init MMEncoder {rank}/{server_args.tp_size}")
|
||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
configure_media_url_security(
|
configure_media_url_security(
|
||||||
server_args.allowed_media_domains,
|
get_mm().allowed_media_domains,
|
||||||
server_args.media_url_max_file_size_mb,
|
server_args.media_url_max_file_size_mb,
|
||||||
)
|
)
|
||||||
self.rank = rank
|
self.rank = rank
|
||||||
@@ -329,7 +334,7 @@ class MMEncoder:
|
|||||||
server_args,
|
server_args,
|
||||||
)
|
)
|
||||||
self.load_config = LoadConfig(
|
self.load_config = LoadConfig(
|
||||||
load_format=server_args.load_format,
|
load_format=get_model().load_format,
|
||||||
download_dir=server_args.download_dir,
|
download_dir=server_args.download_dir,
|
||||||
model_loader_extra_config=server_args.model_loader_extra_config,
|
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,
|
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"
|
self.model_config.hf_config, "model_type", "unknown"
|
||||||
).lower()
|
).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.gpu_id = server_args.base_gpu_id + rank if gpu_id is None else gpu_id
|
||||||
|
|
||||||
self.device_config = DeviceConfig(
|
self.device_config = DeviceConfig(
|
||||||
@@ -354,7 +359,7 @@ class MMEncoder:
|
|||||||
use_image_processor_gpu
|
use_image_processor_gpu
|
||||||
and resolve_image_processor_backend(server_args) != "pil"
|
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()
|
self.model_audio_sr = self._resolve_audio_sr()
|
||||||
logger.info(f"Resolved model audio sample rate: {self.model_audio_sr} Hz")
|
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_model_parallel(tensor_model_parallel_size=server_args.tp_size)
|
||||||
initialize_dp_attention(server_args, self.model_config)
|
initialize_dp_attention(server_args, self.model_config)
|
||||||
|
|
||||||
self.model = get_model(
|
self.model = load_model(
|
||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
load_config=self.load_config,
|
load_config=self.load_config,
|
||||||
device_config=self.device_config,
|
device_config=self.device_config,
|
||||||
@@ -628,7 +633,7 @@ class MMEncoder:
|
|||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
self.image_processor = AutoImageProcessor.from_pretrained(
|
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,
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
revision=server_args.revision,
|
revision=server_args.revision,
|
||||||
**image_processor_kwargs,
|
**image_processor_kwargs,
|
||||||
@@ -639,7 +644,7 @@ class MMEncoder:
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
self.video_processor = AutoVideoProcessor.from_pretrained(
|
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,
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
revision=server_args.revision,
|
revision=server_args.revision,
|
||||||
)
|
)
|
||||||
@@ -650,7 +655,7 @@ class MMEncoder:
|
|||||||
try:
|
try:
|
||||||
# Note: AutoProcessor is used for audio processor
|
# Note: AutoProcessor is used for audio processor
|
||||||
_audio_proc = AutoProcessor.from_pretrained(
|
_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,
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
revision=server_args.revision,
|
revision=server_args.revision,
|
||||||
)
|
)
|
||||||
@@ -2100,7 +2105,7 @@ class MMEncoder:
|
|||||||
|
|
||||||
_zmq_xfer_start = time.perf_counter()
|
_zmq_xfer_start = time.perf_counter()
|
||||||
if (
|
if (
|
||||||
self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
|
get_disagg().encoder_transfer_backend == "zmq_to_scheduler"
|
||||||
and url is not None
|
and url is not None
|
||||||
):
|
):
|
||||||
lock = self.scheduler_send_locks.get(endpoint)
|
lock = self.scheduler_send_locks.get(endpoint)
|
||||||
@@ -2146,7 +2151,7 @@ class MMEncoder:
|
|||||||
if encoder_metrics_collector is not None:
|
if encoder_metrics_collector is not None:
|
||||||
encoder_metrics_collector.observe_transfer(
|
encoder_metrics_collector.observe_transfer(
|
||||||
time.perf_counter() - _zmq_xfer_start,
|
time.perf_counter() - _zmq_xfer_start,
|
||||||
backend=self.server_args.encoder_transfer_backend,
|
backend=get_disagg().encoder_transfer_backend,
|
||||||
)
|
)
|
||||||
return
|
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
|
# No-op for mooncake (its /send is separate). embedding_port=None is
|
||||||
# rejected upfront, so ports is always a concrete list here.
|
# rejected upfront, so ports is always a concrete list here.
|
||||||
req_id = request["req_id"]
|
req_id = request["req_id"]
|
||||||
backend = enc.server_args.encoder_transfer_backend
|
backend = get_disagg().encoder_transfer_backend
|
||||||
|
|
||||||
if backend == "zmq_to_tokenizer":
|
if backend == "zmq_to_tokenizer":
|
||||||
await enc.send(
|
await enc.send(
|
||||||
@@ -3050,7 +3055,7 @@ async def _dp_worker_encode_and_send(
|
|||||||
modality = Modality.from_str(request["modality"])
|
modality = Modality.from_str(request["modality"])
|
||||||
time_stats.modality = modality.name.lower()
|
time_stats.modality = modality.name.lower()
|
||||||
time_stats.set_metrics_collector(encoder_metrics_collector)
|
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.
|
# 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:
|
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
|
global encoder_metrics_collector
|
||||||
if server_args.enable_metrics:
|
if get_observability().enable_metrics:
|
||||||
set_prometheus_multiproc_dir()
|
set_prometheus_multiproc_dir()
|
||||||
labels = {
|
labels = {
|
||||||
"model_name": server_args.served_model_name,
|
"model_name": get_serving().served_model_name,
|
||||||
"dp_rank": str(dp_rank),
|
"dp_rank": str(dp_rank),
|
||||||
}
|
}
|
||||||
if server_args.extra_metric_labels:
|
if get_observability().extra_metric_labels:
|
||||||
labels.update(server_args.extra_metric_labels)
|
labels.update(get_observability().extra_metric_labels)
|
||||||
encoder_metrics_collector = EncoderMetricsCollector(labels)
|
encoder_metrics_collector = EncoderMetricsCollector(labels)
|
||||||
enc.dp_rank = dp_rank
|
enc.dp_rank = dp_rank
|
||||||
|
|
||||||
@@ -3964,14 +3969,14 @@ def launch_server(server_args: ServerArgs):
|
|||||||
global encoder, encoder_metrics_collector
|
global encoder, encoder_metrics_collector
|
||||||
|
|
||||||
# Set up prometheus metrics.
|
# Set up prometheus metrics.
|
||||||
if server_args.enable_metrics:
|
if get_observability().enable_metrics:
|
||||||
set_prometheus_multiproc_dir()
|
set_prometheus_multiproc_dir()
|
||||||
labels = {
|
labels = {
|
||||||
"model_name": server_args.served_model_name,
|
"model_name": get_serving().served_model_name,
|
||||||
"dp_rank": "0",
|
"dp_rank": "0",
|
||||||
}
|
}
|
||||||
if server_args.extra_metric_labels:
|
if get_observability().extra_metric_labels:
|
||||||
labels.update(server_args.extra_metric_labels)
|
labels.update(get_observability().extra_metric_labels)
|
||||||
encoder_metrics_collector = EncoderMetricsCollector(labels)
|
encoder_metrics_collector = EncoderMetricsCollector(labels)
|
||||||
add_prometheus_middleware(app)
|
add_prometheus_middleware(app)
|
||||||
|
|
||||||
@@ -3979,21 +3984,21 @@ def launch_server(server_args: ServerArgs):
|
|||||||
zmq_ctx = zmq.Context(10)
|
zmq_ctx = zmq.Context(10)
|
||||||
ipc_path_prefix = random_uuid()
|
ipc_path_prefix = random_uuid()
|
||||||
port_args = PortArgs.init_new(server_args)
|
port_args = PortArgs.init_new(server_args)
|
||||||
if server_args.dist_init_addr:
|
if get_parallel().dist_init_addr:
|
||||||
na = NetworkAddress.parse(server_args.dist_init_addr)
|
na = NetworkAddress.parse(get_parallel().dist_init_addr)
|
||||||
dist_init_method = na.to_tcp()
|
dist_init_method = na.to_tcp()
|
||||||
else:
|
else:
|
||||||
dist_init_method = NetworkAddress(
|
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()
|
).to_tcp()
|
||||||
if server_args.enable_trace:
|
if get_observability().enable_trace:
|
||||||
process_tracing_init(
|
process_tracing_init(
|
||||||
server_args.otlp_traces_endpoint,
|
get_observability().otlp_traces_endpoint,
|
||||||
"sglang",
|
"sglang",
|
||||||
trace_modules=server_args.trace_modules,
|
trace_modules=get_observability().trace_modules,
|
||||||
)
|
)
|
||||||
trace_set_thread_info("Encoder")
|
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}"
|
schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}"
|
||||||
send_sockets.append(
|
send_sockets.append(
|
||||||
get_zmq_socket(zmq_ctx, zmq.PUSH, schedule_path, bind=False)
|
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)
|
encoder = MMEncoder(server_args, dist_init_method=dist_init_method)
|
||||||
|
|
||||||
# Register this encoder's URL with prefill server(s) if configured.
|
# 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
|
import atexit
|
||||||
|
|
||||||
_register_encoder_url_with_bootstrap(server_args)
|
_register_encoder_url_with_bootstrap(server_args)
|
||||||
atexit.register(_unregister_encoder_url_from_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):
|
def _launch_server_dp(server_args: ServerArgs):
|
||||||
@@ -4083,7 +4088,7 @@ def _launch_server_dp(server_args: ServerArgs):
|
|||||||
proc.start()
|
proc.start()
|
||||||
worker_processes.append(proc)
|
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:
|
if server_args.extra_metric_labels:
|
||||||
labels.update(server_args.extra_metric_labels)
|
labels.update(server_args.extra_metric_labels)
|
||||||
dp_dispatcher = DPDispatcher(
|
dp_dispatcher = DPDispatcher(
|
||||||
@@ -4188,7 +4193,7 @@ async def handle_encode_request(request: dict):
|
|||||||
# when multiple decoder TP ranks POST /encode
|
# when multiple decoder TP ranks POST /encode
|
||||||
# with the same req_id, only the first triggers the VIT forward;
|
# with the same req_id, only the first triggers the VIT forward;
|
||||||
# subsequent callers wait and return the same metadata.
|
# 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:
|
async with encoder._inflight_encode_lock:
|
||||||
if req_id in encoder._inflight_encode_events:
|
if req_id in encoder._inflight_encode_events:
|
||||||
event = encoder._inflight_encode_events[req_id]
|
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()
|
time_stats.set_mm_encode_end_time()
|
||||||
|
|
||||||
if error_msg:
|
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:
|
if request["embedding_port"] is None:
|
||||||
start_background_send(req_id)
|
start_background_send(req_id)
|
||||||
else:
|
else:
|
||||||
@@ -4285,7 +4290,7 @@ async def handle_encode_request(request: dict):
|
|||||||
embedding_port=port,
|
embedding_port=port,
|
||||||
)
|
)
|
||||||
# Signal waiters on failure for mooncake
|
# 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)
|
encoder._inflight_encode_meta.pop(req_id, None)
|
||||||
evt = encoder._inflight_encode_events.pop(req_id, None)
|
evt = encoder._inflight_encode_events.pop(req_id, None)
|
||||||
if evt:
|
if evt:
|
||||||
@@ -4299,7 +4304,7 @@ async def handle_encode_request(request: dict):
|
|||||||
status_code=error_code,
|
status_code=error_code,
|
||||||
content={"status": "error", "message": error_msg, "req_id": req_id},
|
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
|
# Store metadata for duplicate callers and signal them
|
||||||
encoder._inflight_encode_meta[req_id] = (
|
encoder._inflight_encode_meta[req_id] = (
|
||||||
nbytes,
|
nbytes,
|
||||||
@@ -4323,7 +4328,7 @@ async def handle_encode_request(request: dict):
|
|||||||
modality=modality_str, status="success"
|
modality=modality_str, status="success"
|
||||||
)
|
)
|
||||||
return ORJSONResponse(content=request)
|
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'] = }")
|
logger.info(f"{request['embedding_port'] = }")
|
||||||
if request["embedding_port"] is None:
|
if request["embedding_port"] is None:
|
||||||
await encoder.send_with_url(
|
await encoder.send_with_url(
|
||||||
@@ -4347,7 +4352,7 @@ async def handle_encode_request(request: dict):
|
|||||||
modality=modality_str, status="success"
|
modality=modality_str, status="success"
|
||||||
)
|
)
|
||||||
return ORJSONResponse(content=None)
|
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(
|
await encoder.send(
|
||||||
req_id=request["req_id"],
|
req_id=request["req_id"],
|
||||||
prefill_host=request["prefill_host"],
|
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}")
|
logger.error(f"Unexpected error in encoder logic for {req_id}: {error_msg}")
|
||||||
rid_to_err_msg[req_id] = error_msg
|
rid_to_err_msg[req_id] = error_msg
|
||||||
# Ensure inflight waiters are unblocked on unexpected errors
|
# 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)
|
encoder._inflight_encode_meta.pop(req_id, None)
|
||||||
evt = encoder._inflight_encode_events.pop(req_id, None)
|
evt = encoder._inflight_encode_events.pop(req_id, None)
|
||||||
if evt:
|
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_detection import resolve_auto_parsers
|
||||||
from sglang.srt.parser.template_manager import TemplateManager
|
from sglang.srt.parser.template_manager import TemplateManager
|
||||||
from sglang.srt.plugins import load_plugins
|
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.server_args import PortArgs, ServerArgs
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
MultiprocessingSerializer,
|
MultiprocessingSerializer,
|
||||||
@@ -158,9 +165,9 @@ def init_tokenizer_manager(
|
|||||||
template_manager = TemplateManager()
|
template_manager = TemplateManager()
|
||||||
template_manager.initialize_templates(
|
template_manager.initialize_templates(
|
||||||
tokenizer_manager=tokenizer_manager,
|
tokenizer_manager=tokenizer_manager,
|
||||||
model_path=server_args.model_path,
|
model_path=get_model().model_path,
|
||||||
chat_template=server_args.chat_template,
|
chat_template=get_serving().chat_template,
|
||||||
completion_template=server_args.completion_template,
|
completion_template=get_serving().completion_template,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Resolve any remaining auto parsers using template manager's detection results
|
# 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 = (
|
pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = (
|
||||||
_calculate_rank_ranges(
|
_calculate_rank_ranges(
|
||||||
server_args.nnodes,
|
server_args.nnodes,
|
||||||
server_args.pp_size,
|
configured_pp_size(),
|
||||||
tp_size,
|
tp_size,
|
||||||
server_args.node_rank,
|
server_args.node_rank,
|
||||||
)
|
)
|
||||||
@@ -702,7 +709,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
|||||||
daemon_procs = []
|
daemon_procs = []
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Launching {num_daemons} weight cache daemon(s) on node "
|
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"pp_ranks={pp_rank_range.start}..{pp_rank_range.stop - 1}, "
|
||||||
f"tp_ranks={tp_rank_range.start}..{tp_rank_range.stop - 1}, "
|
f"tp_ranks={tp_rank_range.start}..{tp_rank_range.stop - 1}, "
|
||||||
f"dist_init_method={dist_init_method}"
|
f"dist_init_method={dist_init_method}"
|
||||||
@@ -737,7 +744,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
|||||||
"-m",
|
"-m",
|
||||||
"sglang.srt.weight_cache.daemon",
|
"sglang.srt.weight_cache.daemon",
|
||||||
"--model-path",
|
"--model-path",
|
||||||
server_args.model_path,
|
get_model().model_path,
|
||||||
"--gpu-id",
|
"--gpu-id",
|
||||||
str(gpu_id),
|
str(gpu_id),
|
||||||
"--tp-size",
|
"--tp-size",
|
||||||
@@ -745,7 +752,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
|||||||
"--tp-rank",
|
"--tp-rank",
|
||||||
str(tp_rank),
|
str(tp_rank),
|
||||||
"--pp-size",
|
"--pp-size",
|
||||||
str(server_args.pp_size),
|
str(configured_pp_size()),
|
||||||
"--pp-rank",
|
"--pp-rank",
|
||||||
str(pp_rank),
|
str(pp_rank),
|
||||||
"--dp-size",
|
"--dp-size",
|
||||||
@@ -753,14 +760,14 @@ class Engine(EngineScoreMixin, EngineBase):
|
|||||||
"--ep-size",
|
"--ep-size",
|
||||||
str(get_parallel().ep_size),
|
str(get_parallel().ep_size),
|
||||||
"--load-format",
|
"--load-format",
|
||||||
server_args.load_format,
|
get_model().load_format,
|
||||||
"--dtype",
|
"--dtype",
|
||||||
server_args.dtype,
|
get_model().dtype,
|
||||||
"--dist-init-method",
|
"--dist-init-method",
|
||||||
dist_init_method,
|
dist_init_method,
|
||||||
]
|
]
|
||||||
if server_args.quantization:
|
if get_model().quantization:
|
||||||
cmd += ["--quantization", server_args.quantization]
|
cmd += ["--quantization", get_model().quantization]
|
||||||
if (
|
if (
|
||||||
server_args.model_loader_extra_config
|
server_args.model_loader_extra_config
|
||||||
and server_args.model_loader_extra_config != "{}"
|
and server_args.model_loader_extra_config != "{}"
|
||||||
@@ -863,7 +870,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
|||||||
"""
|
"""
|
||||||
scheduler_procs = []
|
scheduler_procs = []
|
||||||
use_dp_controller = (
|
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:
|
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 = (
|
pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = (
|
||||||
_calculate_rank_ranges(
|
_calculate_rank_ranges(
|
||||||
server_args.nnodes,
|
server_args.nnodes,
|
||||||
server_args.pp_size,
|
configured_pp_size(),
|
||||||
server_args.tp_size,
|
server_args.tp_size,
|
||||||
server_args.node_rank,
|
server_args.node_rank,
|
||||||
)
|
)
|
||||||
@@ -983,7 +990,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
|||||||
processes: List[mp.Process] = []
|
processes: List[mp.Process] = []
|
||||||
names: List[str] = []
|
names: List[str] = []
|
||||||
|
|
||||||
if server_args.detokenizer_worker_num <= 1:
|
if get_serving().detokenizer_worker_num <= 1:
|
||||||
proc = mp.Process(
|
proc = mp.Process(
|
||||||
target=run_detokenizer_process_func,
|
target=run_detokenizer_process_func,
|
||||||
args=(server_args, port_args),
|
args=(server_args, port_args),
|
||||||
@@ -996,7 +1003,7 @@ class Engine(EngineScoreMixin, EngineBase):
|
|||||||
router_ipc_name = port_args.detokenizer_ipc_name
|
router_ipc_name = port_args.detokenizer_ipc_name
|
||||||
worker_ipc_names: List[str] = []
|
worker_ipc_names: List[str] = []
|
||||||
try:
|
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}"
|
worker_ipc = f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}"
|
||||||
port_args.detokenizer_ipc_name = worker_ipc
|
port_args.detokenizer_ipc_name = worker_ipc
|
||||||
proc = mp.Process(
|
proc = mp.Process(
|
||||||
|
|||||||
@@ -247,7 +247,7 @@ async def init_multi_tokenizer() -> ServerArgs:
|
|||||||
template_manager = TemplateManager()
|
template_manager = TemplateManager()
|
||||||
template_manager.initialize_templates(
|
template_manager.initialize_templates(
|
||||||
tokenizer_manager=tokenizer_manager,
|
tokenizer_manager=tokenizer_manager,
|
||||||
model_path=server_args.model_path,
|
model_path=get_model().model_path,
|
||||||
chat_template=server_args.chat_template,
|
chat_template=server_args.chat_template,
|
||||||
completion_template=server_args.completion_template,
|
completion_template=server_args.completion_template,
|
||||||
)
|
)
|
||||||
@@ -293,9 +293,9 @@ async def lifespan(fast_api_app: FastAPI):
|
|||||||
"sglang",
|
"sglang",
|
||||||
trace_modules=server_args.trace_modules,
|
trace_modules=server_args.trace_modules,
|
||||||
)
|
)
|
||||||
if server_args.disaggregation_mode == "prefill":
|
if get_disagg().disaggregation_mode == "prefill":
|
||||||
thread_label = "Prefill" + thread_label
|
thread_label = "Prefill" + thread_label
|
||||||
elif server_args.disaggregation_mode == "decode":
|
elif get_disagg().disaggregation_mode == "decode":
|
||||||
thread_label = "Decode" + thread_label
|
thread_label = "Decode" + thread_label
|
||||||
trace_set_thread_info(thread_label)
|
trace_set_thread_info(thread_label)
|
||||||
|
|
||||||
@@ -380,7 +380,7 @@ async def lifespan(fast_api_app: FastAPI):
|
|||||||
# Execute custom warmups
|
# Execute custom warmups
|
||||||
if server_args.warmups is not None:
|
if server_args.warmups is not None:
|
||||||
await execute_warmups(
|
await execute_warmups(
|
||||||
server_args.disaggregation_mode,
|
get_disagg().disaggregation_mode,
|
||||||
server_args.warmups.split(","),
|
server_args.warmups.split(","),
|
||||||
_global_state.tokenizer_manager,
|
_global_state.tokenizer_manager,
|
||||||
)
|
)
|
||||||
@@ -393,7 +393,7 @@ async def lifespan(fast_api_app: FastAPI):
|
|||||||
try:
|
try:
|
||||||
if (
|
if (
|
||||||
getattr(fast_api_app, "is_single_tokenizer_mode", False)
|
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)
|
and not (server_args.smg_grpc_mode or server_args.grpc_mode)
|
||||||
):
|
):
|
||||||
grpc_handle = _start_native_grpc_server_for_runtime(
|
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,
|
tokenizer_manager=_global_state.tokenizer_manager,
|
||||||
template_manager=_global_state.template_manager,
|
template_manager=_global_state.template_manager,
|
||||||
scheduler_info=_global_state.scheduler_info,
|
scheduler_info=_global_state.scheduler_info,
|
||||||
|
grpc_port=get_serving().grpc_port,
|
||||||
)
|
)
|
||||||
if server_args.sidecar is not None:
|
if server_args.sidecar is not None:
|
||||||
from sglang.srt.entrypoints.sidecar import start_sidecar
|
from sglang.srt.entrypoints.sidecar import start_sidecar
|
||||||
@@ -480,7 +481,13 @@ v1_loads_router.route_class = ORJSONRoute
|
|||||||
app.include_router(v1_loads_router)
|
app.include_router(v1_loads_router)
|
||||||
|
|
||||||
from sglang.srt.entrypoints.elastic_ep import router as elastic_ep_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
|
elastic_ep_router.route_class = ORJSONRoute
|
||||||
app.include_router(elastic_ep_router)
|
app.include_router(elastic_ep_router)
|
||||||
@@ -679,10 +686,7 @@ async def health_generate(request: Request) -> Response:
|
|||||||
sampling_params=sampling_params,
|
sampling_params=sampling_params,
|
||||||
log_metrics=False,
|
log_metrics=False,
|
||||||
)
|
)
|
||||||
if (
|
if get_disagg().disaggregation_mode != DisaggregationMode.NULL.value:
|
||||||
_global_state.tokenizer_manager.server_args.disaggregation_mode
|
|
||||||
!= DisaggregationMode.NULL.value
|
|
||||||
):
|
|
||||||
gri.bootstrap_host = FAKE_BOOTSTRAP_HOST
|
gri.bootstrap_host = FAKE_BOOTSTRAP_HOST
|
||||||
gri.bootstrap_room = 0
|
gri.bootstrap_room = 0
|
||||||
else:
|
else:
|
||||||
@@ -2224,7 +2228,7 @@ def _execute_server_warmup(server_args: ServerArgs):
|
|||||||
json_data["input_ids"] = json_data["input_ids"][0]
|
json_data["input_ids"] = json_data["input_ids"][0]
|
||||||
elif (
|
elif (
|
||||||
is_vlm
|
is_vlm
|
||||||
and server_args.disaggregation_mode == "null"
|
and get_disagg().disaggregation_mode == "null"
|
||||||
and model_info["is_generation"]
|
and model_info["is_generation"]
|
||||||
):
|
):
|
||||||
served_model_name = ""
|
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,
|
# _global_state.tokenizer_manager is not initialized in the rust server,
|
||||||
# so we need to get the model name from the model_info
|
# so we need to get the model name from the model_info
|
||||||
served_model_name = model_info.get(
|
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
|
# 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
|
# Only use chat completions format for generation models, not embedding models
|
||||||
json_data = {
|
json_data = {
|
||||||
@@ -2280,7 +2284,7 @@ def _execute_server_warmup(server_args: ServerArgs):
|
|||||||
# Send a warmup request
|
# Send a warmup request
|
||||||
warmup_timeout = envs.SGLANG_WARMUP_TIMEOUT.get()
|
warmup_timeout = envs.SGLANG_WARMUP_TIMEOUT.get()
|
||||||
try:
|
try:
|
||||||
if server_args.disaggregation_mode == "null":
|
if get_disagg().disaggregation_mode == "null":
|
||||||
res = requests.post(
|
res = requests.post(
|
||||||
url + request_name,
|
url + request_name,
|
||||||
json=json_data,
|
json=json_data,
|
||||||
@@ -2314,7 +2318,7 @@ def _execute_server_warmup(server_args: ServerArgs):
|
|||||||
else:
|
else:
|
||||||
logger.info(
|
logger.info(
|
||||||
"Disaggregation warmup failed (mode=%s), status codes: %s",
|
"Disaggregation warmup failed (mode=%s), status codes: %s",
|
||||||
server_args.disaggregation_mode,
|
get_disagg().disaggregation_mode,
|
||||||
failed_status_codes,
|
failed_status_codes,
|
||||||
)
|
)
|
||||||
# In rust-server mode there is no TokenizerManager (readiness is
|
# In rust-server mode there is no TokenizerManager (readiness is
|
||||||
@@ -2368,10 +2372,10 @@ def _wait_and_warmup(
|
|||||||
logger.debug(
|
logger.debug(
|
||||||
"[Elastic EP] Skipping server warmup for elastic joiner "
|
"[Elastic EP] Skipping server warmup for elastic joiner "
|
||||||
"(ep_join_mode=%s)",
|
"(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):
|
if not execute_warmup_func(server_args):
|
||||||
return
|
return
|
||||||
else:
|
else:
|
||||||
@@ -2383,7 +2387,7 @@ def _wait_and_warmup(
|
|||||||
logger.info("The server is fired up and ready to roll!")
|
logger.info("The server is fired up and ready to roll!")
|
||||||
|
|
||||||
if server_args.delete_ckpt_after_loading:
|
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:
|
if server_args.debug_tensor_dump_input_file:
|
||||||
kill_process_tree(os.getpid())
|
kill_process_tree(os.getpid())
|
||||||
@@ -2711,6 +2715,7 @@ def _start_native_grpc_server_for_runtime(
|
|||||||
tokenizer_manager,
|
tokenizer_manager,
|
||||||
template_manager,
|
template_manager,
|
||||||
scheduler_info,
|
scheduler_info,
|
||||||
|
grpc_port,
|
||||||
):
|
):
|
||||||
from sglang.srt.entrypoints.grpc_bridge import RuntimeHandle
|
from sglang.srt.entrypoints.grpc_bridge import RuntimeHandle
|
||||||
from sglang.srt.rust_extensions import load_rust_extension
|
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(
|
grpc_handle = grpc_native.start_server(
|
||||||
host=server_args.host,
|
host=server_args.host,
|
||||||
port=server_args.grpc_port,
|
port=grpc_port,
|
||||||
runtime_handle=runtime_handle,
|
runtime_handle=runtime_handle,
|
||||||
worker_threads=server_args.grpc_worker_threads,
|
worker_threads=server_args.grpc_worker_threads,
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(f"Native gRPC server started on {server_args.host}:{grpc_port}")
|
||||||
f"Native gRPC server started on {server_args.host}:{server_args.grpc_port}"
|
|
||||||
)
|
|
||||||
return grpc_handle
|
return grpc_handle
|
||||||
|
|
||||||
|
|
||||||
@@ -2790,7 +2793,7 @@ def launch_server(
|
|||||||
# and /get_model_info endpoints are static (200 as soon as the 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
|
# binds, before any forward pass), so without this the first real request
|
||||||
# pays the cold-start cost (observed as a >60s first generation).
|
# 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)
|
_execute_server_warmup(server_args)
|
||||||
logger.info("The server is fired up and ready to roll!")
|
logger.info("The server is fired up and ready to roll!")
|
||||||
if launch_callback is not None:
|
if launch_callback is not None:
|
||||||
|
|||||||
@@ -37,6 +37,8 @@ def _loopback_host(host: str) -> str:
|
|||||||
|
|
||||||
|
|
||||||
def build_sidecar_endpoint(server_args) -> 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(
|
return NetworkAddress(
|
||||||
_loopback_host(server_args.host), server_args.grpc_port
|
_loopback_host(server_args.host), server_args.grpc_port
|
||||||
).to_url()
|
).to_url()
|
||||||
|
|||||||
@@ -1,21 +1,17 @@
|
|||||||
from __future__ import annotations
|
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 import HashOracle
|
||||||
from sglang.srt.kv_canary.token_oracle.oracle_manager import TokenOracleManager
|
from sglang.srt.kv_canary.token_oracle.oracle_manager import TokenOracleManager
|
||||||
from sglang.srt.kv_canary.token_oracle.sampler import install_oracle_sampler
|
from sglang.srt.kv_canary.token_oracle.sampler import install_oracle_sampler
|
||||||
|
from sglang.srt.runtime_context import get_exec
|
||||||
if TYPE_CHECKING:
|
|
||||||
from sglang.srt.server_args import ServerArgs
|
|
||||||
|
|
||||||
|
|
||||||
def install_token_oracle_from_env(
|
def install_token_oracle_from_env(*, vocab_size: int) -> Optional[TokenOracleManager]:
|
||||||
*, server_args: ServerArgs, vocab_size: int
|
|
||||||
) -> Optional[TokenOracleManager]:
|
|
||||||
# Must be called before create_sampler() so the factory is present when the
|
# Must be called before create_sampler() so the factory is present when the
|
||||||
# Sampler is first constructed.
|
# Sampler is first constructed.
|
||||||
if server_args.sampling_backend != "token_oracle":
|
if get_exec().kernel.sampling_backend != "token_oracle":
|
||||||
return None
|
return None
|
||||||
|
|
||||||
oracle = HashOracle(vocab_size=vocab_size)
|
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.managers.multi_tokenizer_mixin import MultiHttpWorkerDetokenizerMixin
|
||||||
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
|
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.server_args import PortArgs, ServerArgs
|
||||||
from sglang.srt.utils import configure_logger, freeze_gc, kill_itself_when_parent_died
|
from sglang.srt.utils import configure_logger, freeze_gc, kill_itself_when_parent_died
|
||||||
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
|
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
|
||||||
@@ -128,7 +128,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
|||||||
self.vocab_size = None
|
self.vocab_size = None
|
||||||
else:
|
else:
|
||||||
self.tokenizer = get_tokenizer(
|
self.tokenizer = get_tokenizer(
|
||||||
server_args.tokenizer_path,
|
get_serving().tokenizer_path,
|
||||||
tokenizer_mode=server_args.tokenizer_mode,
|
tokenizer_mode=server_args.tokenizer_mode,
|
||||||
trust_remote_code=server_args.trust_remote_code,
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
revision=server_args.revision,
|
revision=server_args.revision,
|
||||||
@@ -142,11 +142,11 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
|||||||
def init_running_status(self, server_args: ServerArgs):
|
def init_running_status(self, server_args: ServerArgs):
|
||||||
self.decode_status = LimitedCapacityDict(capacity=DETOKENIZER_MAX_STATES)
|
self.decode_status = LimitedCapacityDict(capacity=DETOKENIZER_MAX_STATES)
|
||||||
self.disable_tokenizer_batch_decode = server_args.disable_tokenizer_batch_decode
|
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(
|
self.soft_watchdog = Watchdog.create(
|
||||||
debug_name="DetokenizerManager",
|
debug_name="DetokenizerManager",
|
||||||
watchdog_timeout=server_args.soft_watchdog_timeout,
|
watchdog_timeout=get_device().soft_watchdog_timeout,
|
||||||
soft=True,
|
soft=True,
|
||||||
test_stuck_time=envs.SGLANG_TEST_STUCK_DETOKENIZER.get(),
|
test_stuck_time=envs.SGLANG_TEST_STUCK_DETOKENIZER.get(),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1252,9 +1252,7 @@ class Scheduler(
|
|||||||
and not get_schedule().disable_priority_preemption
|
and not get_schedule().disable_priority_preemption
|
||||||
)
|
)
|
||||||
|
|
||||||
self.new_token_ratio_tracker = NewTokenRatioTracker.from_server_args(
|
self.new_token_ratio_tracker = NewTokenRatioTracker.from_config()
|
||||||
self.server_args
|
|
||||||
)
|
|
||||||
|
|
||||||
def init_soft_watchdog(self, server_args: ServerArgs):
|
def init_soft_watchdog(self, server_args: ServerArgs):
|
||||||
if (x := server_args.soft_watchdog_timeout) is not None:
|
if (x := server_args.soft_watchdog_timeout) is not None:
|
||||||
|
|||||||
@@ -566,7 +566,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
return get_required_capture_hidden_mode(
|
return get_required_capture_hidden_mode(
|
||||||
max(
|
max(
|
||||||
batch.return_hidden_states_mode,
|
batch.return_hidden_states_mode,
|
||||||
get_server_return_hidden_states_mode(server_args),
|
get_server_return_hidden_states_mode(),
|
||||||
),
|
),
|
||||||
batch.spec_info,
|
batch.spec_info,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from dataclasses import dataclass
|
|||||||
from typing import TYPE_CHECKING, Sequence
|
from typing import TYPE_CHECKING, Sequence
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.runtime_context import get_schedule
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
@@ -18,10 +18,10 @@ class NewTokenRatioTracker:
|
|||||||
current: float
|
current: float
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_server_args(cls, server_args: ServerArgs) -> NewTokenRatioTracker:
|
def from_config(cls) -> NewTokenRatioTracker:
|
||||||
init = min(
|
init = min(
|
||||||
envs.SGLANG_INIT_NEW_TOKEN_RATIO.get()
|
envs.SGLANG_INIT_NEW_TOKEN_RATIO.get()
|
||||||
* server_args.schedule_conservativeness,
|
* get_schedule().schedule_conservativeness,
|
||||||
1.0,
|
1.0,
|
||||||
)
|
)
|
||||||
min_ratio = min(
|
min_ratio = min(
|
||||||
|
|||||||
@@ -121,7 +121,18 @@ from sglang.srt.observability.request_metrics_exporter import (
|
|||||||
RequestMetricsExporterManager,
|
RequestMetricsExporterManager,
|
||||||
)
|
)
|
||||||
from sglang.srt.observability.trace import SpanAttributes, extract_trace_headers
|
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.sampling.sampling_params import SamplingParams
|
||||||
from sglang.srt.server_args import (
|
from sglang.srt.server_args import (
|
||||||
PortArgs,
|
PortArgs,
|
||||||
@@ -160,12 +171,13 @@ _REQUEST_STATE_WAIT_TIMEOUT = envs.SGLANG_REQUEST_STATE_WAIT_TIMEOUT.get()
|
|||||||
logger = logging.getLogger(__name__)
|
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."""
|
"""Do not silently turn a failed EPD request into local vision work."""
|
||||||
|
disagg = get_disagg()
|
||||||
if (
|
if (
|
||||||
mm_inputs is None
|
mm_inputs is None
|
||||||
and server_args.language_only
|
and disagg.language_only
|
||||||
and server_args.encoder_transfer_backend == "zmq_to_tokenizer"
|
and disagg.encoder_transfer_backend == "zmq_to_tokenizer"
|
||||||
and request_obj.need_wait_for_mm_inputs
|
and request_obj.need_wait_for_mm_inputs
|
||||||
):
|
):
|
||||||
raise fastapi.HTTPException(
|
raise fastapi.HTTPException(
|
||||||
@@ -405,11 +417,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
self.elastic_last_error = None
|
self.elastic_last_error = None
|
||||||
self.enable_metrics = server_args.enable_metrics
|
self.enable_metrics = server_args.enable_metrics
|
||||||
self.incremental_streaming_output = server_args.incremental_streaming_output
|
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.enable_trace = server_args.enable_trace
|
||||||
self.allow_auto_truncate = server_args.allow_auto_truncate
|
self.allow_auto_truncate = server_args.allow_auto_truncate
|
||||||
self.skip_tokenizer_init = server_args.skip_tokenizer_init
|
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
|
self.crash_dump_folder = server_args.crash_dump_folder
|
||||||
|
|
||||||
# Init model config
|
# Init model config
|
||||||
@@ -501,7 +513,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
self.tokenizer = None
|
self.tokenizer = None
|
||||||
else:
|
else:
|
||||||
self.tokenizer = get_tokenizer(
|
self.tokenizer = get_tokenizer(
|
||||||
server_args.tokenizer_path,
|
get_serving().tokenizer_path,
|
||||||
tokenizer_mode=server_args.tokenizer_mode,
|
tokenizer_mode=server_args.tokenizer_mode,
|
||||||
trust_remote_code=server_args.trust_remote_code,
|
trust_remote_code=server_args.trust_remote_code,
|
||||||
revision=server_args.revision,
|
revision=server_args.revision,
|
||||||
@@ -522,7 +534,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
self.async_dynamic_batch_tokenizer = None
|
self.async_dynamic_batch_tokenizer = None
|
||||||
|
|
||||||
def _validate_cuda_vmm_feature_transport_support(self) -> 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
|
return
|
||||||
|
|
||||||
from sglang.srt.model_loader.utils import get_model_architecture
|
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
|
# 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
|
# serves as the source of truth for available adapters and maps user-friendly LoRA names
|
||||||
# to internally used unique LoRA IDs.
|
# 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.
|
# Lock to serialize LoRA update operations.
|
||||||
# Please note that, unlike `model_update_lock`, this does not block inference, allowing
|
# Please note that, unlike `model_update_lock`, this does not block inference, allowing
|
||||||
# LoRA updates and inference to overlap.
|
# 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
|
# point to their latest LoRARef objects, so that they can be
|
||||||
# dynamically loaded if needed for inference
|
# dynamically loaded if needed for inference
|
||||||
self.lora_ref_cache: Dict[str, LoRARef] = {}
|
self.lora_ref_cache: Dict[str, LoRARef] = {}
|
||||||
if self.server_args.lora_paths is not None:
|
if get_lora().lora_paths is not None:
|
||||||
for lora_ref in self.server_args.lora_paths:
|
for lora_ref in get_lora().lora_paths:
|
||||||
self.lora_ref_cache[lora_ref.lora_name] = lora_ref
|
self.lora_ref_cache[lora_ref.lora_name] = lora_ref
|
||||||
|
|
||||||
def init_disaggregation(self, *, start_pd_bootstrap_service: bool = True):
|
def init_disaggregation(self, *, start_pd_bootstrap_service: bool = True):
|
||||||
# PD Disaggregation
|
# PD Disaggregation
|
||||||
self.disaggregation_mode = DisaggregationMode(
|
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
||||||
self.server_args.disaggregation_mode
|
|
||||||
)
|
|
||||||
# Keep a reference so the bootstrap server is not garbage-collected.
|
# Keep a reference so the bootstrap server is not garbage-collected.
|
||||||
self.bootstrap_server = (
|
self.bootstrap_server = (
|
||||||
start_disagg_service(self.server_args)
|
start_disagg_service(self.server_args)
|
||||||
@@ -688,7 +698,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
# Metrics
|
# Metrics
|
||||||
if self.enable_metrics:
|
if self.enable_metrics:
|
||||||
engine_type = DisaggregationMode.to_engine_type(
|
engine_type = DisaggregationMode.to_engine_type(
|
||||||
self.server_args.disaggregation_mode
|
get_disagg().disaggregation_mode
|
||||||
)
|
)
|
||||||
|
|
||||||
labels = {
|
labels = {
|
||||||
@@ -721,7 +731,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
configure_gc_warning(self.server_args.gc_warning_threshold_secs)
|
configure_gc_warning(self.server_args.gc_warning_threshold_secs)
|
||||||
self.soft_watchdog = Watchdog.create(
|
self.soft_watchdog = Watchdog.create(
|
||||||
debug_name="TokenizerManager",
|
debug_name="TokenizerManager",
|
||||||
watchdog_timeout=self.server_args.soft_watchdog_timeout,
|
watchdog_timeout=get_device().soft_watchdog_timeout,
|
||||||
soft=True,
|
soft=True,
|
||||||
test_stuck_time=envs.SGLANG_TEST_STUCK_TOKENIZER.get(),
|
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
|
isinstance(obj, EmbeddingReqInput) and obj.is_cross_encoder_request
|
||||||
)
|
)
|
||||||
if obj.input_embeds is not None:
|
if obj.input_embeds is not None:
|
||||||
if not self.server_args.disable_radix_cache:
|
if not get_memory().disable_radix_cache:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"input_embeds is provided while disable_radix_cache is False. "
|
"input_embeds is provided while disable_radix_cache is False. "
|
||||||
"Please add `--disable-radix-cache` when you launch the server "
|
"Please add `--disable-radix-cache` when you launch the server "
|
||||||
@@ -1062,7 +1072,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
|
|
||||||
if (
|
if (
|
||||||
not self.server_args.language_only
|
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:
|
if self.server_args.language_only:
|
||||||
mm_inputs = await self.mm_receiver.recv_mm_data(
|
mm_inputs = await self.mm_receiver.recv_mm_data(
|
||||||
@@ -1071,11 +1081,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
prompt=mm_processor_input,
|
prompt=mm_processor_input,
|
||||||
need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs,
|
need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs,
|
||||||
)
|
)
|
||||||
_reject_missing_dispatched_encoder_embedding(
|
_reject_missing_dispatched_encoder_embedding(obj, mm_inputs)
|
||||||
self.server_args, obj, mm_inputs
|
|
||||||
)
|
|
||||||
if mm_inputs is None:
|
if mm_inputs is None:
|
||||||
if self.server_args.language_only:
|
if get_disagg().language_only:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Encoder embedding not available, "
|
"Encoder embedding not available, "
|
||||||
"falling back to local mm processing"
|
"falling back to local mm processing"
|
||||||
@@ -1089,7 +1097,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
)
|
)
|
||||||
elif (
|
elif (
|
||||||
self.server_args.language_only
|
self.server_args.language_only
|
||||||
and self.server_args.encoder_transfer_backend
|
and get_disagg().encoder_transfer_backend
|
||||||
in ["zmq_to_scheduler", "mooncake"]
|
in ["zmq_to_scheduler", "mooncake"]
|
||||||
and not obj.need_wait_for_mm_inputs
|
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(
|
requested_hidden_mode = get_request_return_hidden_states_mode(
|
||||||
obj.return_hidden_states
|
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 requested_hidden_mode > server_hidden_mode:
|
||||||
if server_hidden_mode.need_capture():
|
if server_hidden_mode.need_capture():
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"The requested return_hidden_states mode exceeds the "
|
"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` "
|
"Please launch with `--return-hidden-states-mode full` "
|
||||||
"to allow return_hidden_states=True."
|
"to allow return_hidden_states=True."
|
||||||
)
|
)
|
||||||
@@ -1283,10 +1291,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
def _validate_mm_limits(
|
def _validate_mm_limits(
|
||||||
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
|
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
|
||||||
) -> None:
|
) -> None:
|
||||||
if not self.server_args.limit_mm_data_per_request:
|
if not get_mm().limit_mm_data_per_request:
|
||||||
return
|
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)
|
data = getattr(obj, f"{modality}_data", None)
|
||||||
if data:
|
if data:
|
||||||
count = len(data) if isinstance(data, list) else 1
|
count = len(data) if isinstance(data, list) else 1
|
||||||
@@ -1399,7 +1407,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
bootstrap_room = obj.bootstrap_room
|
bootstrap_room = obj.bootstrap_room
|
||||||
if (
|
if (
|
||||||
bootstrap_room is None
|
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
|
bootstrap_room = self.fake_bootstrap_room_counter
|
||||||
self.fake_bootstrap_room_counter += 1
|
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
|
- Batch tokenization does not support DP attention yet, and it will make everything goes to the first rank currently
|
||||||
"""
|
"""
|
||||||
return batch_size > 0 and (
|
return batch_size > 0 and (
|
||||||
self.server_args.enable_tokenizer_batch_encode
|
get_serving().enable_tokenizer_batch_encode
|
||||||
or (
|
or (
|
||||||
(not get_parallel().enable_dp_attention)
|
(not get_parallel().enable_dp_attention)
|
||||||
and (not self._batch_has_text(batch_size, requests))
|
and (not self._batch_has_text(batch_size, requests))
|
||||||
@@ -2464,7 +2472,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
state.time_stats.set_finished_time()
|
state.time_stats.set_finished_time()
|
||||||
meta_info["e2e_latency"] = state.time_stats.get_e2e_latency()
|
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)
|
self._calculate_spec_decoding_metrics(meta_info, recv_obj, i)
|
||||||
if self.enable_metrics:
|
if self.enable_metrics:
|
||||||
scheduler_time_stats = (
|
scheduler_time_stats = (
|
||||||
@@ -2787,7 +2795,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
):
|
):
|
||||||
# Total number of proposed draft tokens per request.
|
# Total number of proposed draft tokens per request.
|
||||||
num_proposed_drafts = recv_obj.spec_verify_ct[i] * (
|
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]
|
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
|
# This flag will be used in _tokenize_one_request to determine processing path
|
||||||
if should_dispatch:
|
if should_dispatch:
|
||||||
obj.need_wait_for_mm_inputs = True
|
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",
|
"zmq_to_scheduler",
|
||||||
"mooncake",
|
"mooncake",
|
||||||
]:
|
]:
|
||||||
@@ -3603,13 +3611,13 @@ async def print_exception_wrapper(func):
|
|||||||
|
|
||||||
def get_processor_wrapper(server_args):
|
def get_processor_wrapper(server_args):
|
||||||
return get_processor(
|
return get_processor(
|
||||||
server_args.tokenizer_path,
|
get_serving().tokenizer_path,
|
||||||
tokenizer_mode=server_args.tokenizer_mode,
|
tokenizer_mode=get_serving().tokenizer_mode,
|
||||||
trust_remote_code=server_args.trust_remote_code,
|
trust_remote_code=get_model().trust_remote_code,
|
||||||
revision=server_args.revision,
|
revision=get_model().revision,
|
||||||
image_processor_backend=resolve_image_processor_backend(server_args),
|
image_processor_backend=resolve_image_processor_backend(server_args),
|
||||||
tokenizer_backend=server_args.tokenizer_backend,
|
tokenizer_backend=get_serving().tokenizer_backend,
|
||||||
model_name=server_args.model_path,
|
model_name=get_model().model_path,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -567,7 +567,7 @@ class CPUGraphRunner:
|
|||||||
self.return_hidden_states_mode = (
|
self.return_hidden_states_mode = (
|
||||||
CaptureHiddenMode.NULL
|
CaptureHiddenMode.NULL
|
||||||
if model_runner.is_draft_worker
|
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.enable_return_hidden_states = self.return_hidden_states_mode.need_capture()
|
||||||
# bs -> compiled fn (text-only / skip_cross_attention=True)
|
# bs -> compiled fn (text-only / skip_cross_attention=True)
|
||||||
|
|||||||
@@ -32,7 +32,7 @@ import warnings
|
|||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from enum import IntEnum, auto
|
from enum import IntEnum, auto
|
||||||
from functools import total_ordering
|
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
|
import torch
|
||||||
|
|
||||||
@@ -231,11 +231,12 @@ def register_attn_tp_sequence_sharded_predicate(
|
|||||||
_attn_tp_sequence_sharded_predicate = predicate
|
_attn_tp_sequence_sharded_predicate = predicate
|
||||||
|
|
||||||
|
|
||||||
def get_server_return_hidden_states_mode(server_args: Any) -> CaptureHiddenMode:
|
def get_server_return_hidden_states_mode() -> CaptureHiddenMode:
|
||||||
mode = getattr(server_args, "return_hidden_states_mode", None)
|
features = get_exec().features
|
||||||
|
mode = features.return_hidden_states_mode
|
||||||
if mode == "last":
|
if mode == "last":
|
||||||
return CaptureHiddenMode.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.FULL
|
||||||
return CaptureHiddenMode.NULL
|
return CaptureHiddenMode.NULL
|
||||||
|
|
||||||
@@ -719,7 +720,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
if model_runner.is_draft_worker
|
if model_runner.is_draft_worker
|
||||||
else max(
|
else max(
|
||||||
batch.return_hidden_states_mode,
|
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(
|
capture_hidden_mode = get_required_capture_hidden_mode(
|
||||||
|
|||||||
@@ -726,7 +726,6 @@ class ModelRunner:
|
|||||||
self._token_oracle_manager = None
|
self._token_oracle_manager = None
|
||||||
return
|
return
|
||||||
self._token_oracle_manager = install_token_oracle_from_env(
|
self._token_oracle_manager = install_token_oracle_from_env(
|
||||||
server_args=self.server_args,
|
|
||||||
vocab_size=self.model_config.vocab_size,
|
vocab_size=self.model_config.vocab_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -261,8 +261,7 @@ def capture_prefill_graph(
|
|||||||
if (
|
if (
|
||||||
model_runner.spec_algorithm.is_eagle()
|
model_runner.spec_algorithm.is_eagle()
|
||||||
and not model_runner.is_draft_worker
|
and not model_runner.is_draft_worker
|
||||||
and get_server_return_hidden_states_mode(model_runner.server_args)
|
and get_server_return_hidden_states_mode() < CaptureHiddenMode.FULL
|
||||||
< CaptureHiddenMode.FULL
|
|
||||||
and not check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE)
|
and not check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE)
|
||||||
):
|
):
|
||||||
logger.info(
|
logger.info(
|
||||||
|
|||||||
@@ -219,7 +219,7 @@ class BaseRunner(ABC):
|
|||||||
self.return_hidden_states_mode = (
|
self.return_hidden_states_mode = (
|
||||||
CaptureHiddenMode.NULL
|
CaptureHiddenMode.NULL
|
||||||
if model_runner.is_draft_worker
|
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.enable_return_hidden_states = self.return_hidden_states_mode.need_capture()
|
||||||
self.attn_tp_size = get_parallel().attn_tp_size
|
self.attn_tp_size = get_parallel().attn_tp_size
|
||||||
@@ -403,7 +403,7 @@ class BaseRunner(ABC):
|
|||||||
capture_hidden_mode = (
|
capture_hidden_mode = (
|
||||||
CaptureHiddenMode.NULL
|
CaptureHiddenMode.NULL
|
||||||
if mr.is_draft_worker
|
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
|
num_tokens_per_req = 1
|
||||||
# A PD prefill target worker's pool has no SpeculativeState, so a
|
# 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_memory,
|
||||||
get_model,
|
get_model,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_server_args,
|
|
||||||
get_stream,
|
get_stream,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
@@ -159,12 +158,13 @@ for backend in CONCAT_ROPE_BACKENDS:
|
|||||||
AttentionBackendRegistry.register(backend, _handle_concat_rope_backend)
|
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()
|
is_decode = forward_batch.forward_mode.is_decode_or_idle()
|
||||||
if is_decode:
|
if is_decode:
|
||||||
backend = server_args.decode_attention_backend or server_args.attention_backend
|
backend = decode_backend
|
||||||
else:
|
else:
|
||||||
backend = server_args.prefill_attention_backend or server_args.attention_backend
|
backend = prefill_backend
|
||||||
if (
|
if (
|
||||||
forward_batch.forward_mode.is_extend_without_speculative()
|
forward_batch.forward_mode.is_extend_without_speculative()
|
||||||
and backend == "fa3"
|
and backend == "fa3"
|
||||||
@@ -456,7 +456,6 @@ class SarvamMoEMLAAttention(nn.Module):
|
|||||||
self.max_position_embeddings = max_position_embeddings
|
self.max_position_embeddings = max_position_embeddings
|
||||||
self.kv_cache_dtype = get_model().kv_cache_dtype
|
self.kv_cache_dtype = get_model().kv_cache_dtype
|
||||||
|
|
||||||
self._server_args = None
|
|
||||||
self.current_attention_backend = None
|
self.current_attention_backend = None
|
||||||
|
|
||||||
if self.q_lora_rank is 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)
|
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)
|
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)
|
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:
|
if forward_method == AttnForwardMethod.MHA_PREFILL:
|
||||||
return self._run_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)
|
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)
|
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)
|
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:
|
if forward_method == AttnForwardMethod.MHA_PREFILL:
|
||||||
output = self._run_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
|
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)
|
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:
|
if forward_method == AttnForwardMethod.MLA_SEPARATE_ROPE:
|
||||||
attn_output = self.attn_mqa(
|
attn_output = self.attn_mqa(
|
||||||
|
|||||||
+14
-9
@@ -9,7 +9,7 @@ import struct
|
|||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from pathlib import Path
|
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
|
from urllib.parse import unquote, urlparse
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@@ -17,8 +17,7 @@ import torch
|
|||||||
import transformers
|
import transformers
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
from sglang.srt.runtime_context import get_mm, get_model
|
||||||
from sglang.srt.server_args import ServerArgs
|
|
||||||
|
|
||||||
CONTENT_HASH_PREFIX = "sha256:"
|
CONTENT_HASH_PREFIX = "sha256:"
|
||||||
_SHA256_HEX_LENGTH = 64
|
_SHA256_HEX_LENGTH = 64
|
||||||
@@ -379,11 +378,17 @@ def resolve_multimodal_item_hash(
|
|||||||
def build_processor_fingerprint(
|
def build_processor_fingerprint(
|
||||||
processor: Any,
|
processor: Any,
|
||||||
hf_config: Any,
|
hf_config: Any,
|
||||||
server_args: ServerArgs,
|
|
||||||
*,
|
*,
|
||||||
extra: Optional[Mapping[str, Any]] = None,
|
extra: Optional[Mapping[str, Any]] = None,
|
||||||
) -> str:
|
) -> 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_payload = (
|
||||||
processor.preprocess_fingerprint_payload()
|
processor.preprocess_fingerprint_payload()
|
||||||
if isinstance(processor, PreprocessFingerprintProvider)
|
if isinstance(processor, PreprocessFingerprintProvider)
|
||||||
@@ -395,10 +400,10 @@ def build_processor_fingerprint(
|
|||||||
"processor_class": f"{type(processor).__module__}.{type(processor).__qualname__}",
|
"processor_class": f"{type(processor).__module__}.{type(processor).__qualname__}",
|
||||||
"model_type": hf_payload.get("model_type"),
|
"model_type": hf_payload.get("model_type"),
|
||||||
"architectures": hf_payload.get("architectures"),
|
"architectures": hf_payload.get("architectures"),
|
||||||
"model_revision": server_args.revision,
|
"model_revision": get_model().revision,
|
||||||
"processor_revision": server_args.revision,
|
"processor_revision": get_model().revision,
|
||||||
"disable_fast_image_processor": server_args.disable_fast_image_processor,
|
"disable_fast_image_processor": get_mm().disable_fast_image_processor,
|
||||||
"mm_process_config": server_args.mm_process_config or {},
|
"mm_process_config": get_mm().mm_process_config or {},
|
||||||
"processor": processor_payload,
|
"processor": processor_payload,
|
||||||
"extra": extra or {},
|
"extra": extra or {},
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -39,6 +39,7 @@ from sglang.srt.multimodal.transport.cuda_ipc import (
|
|||||||
MmItemMemoryPool,
|
MmItemMemoryPool,
|
||||||
get_mm_feature_pool_size_per_worker,
|
get_mm_feature_pool_size_per_worker,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import get_mm
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
CLIENT_MEDIA_EXCEPTIONS,
|
CLIENT_MEDIA_EXCEPTIONS,
|
||||||
configure_media_url_security,
|
configure_media_url_security,
|
||||||
@@ -212,10 +213,10 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
self.transport_mode = transport_mode
|
self.transport_mode = transport_mode
|
||||||
configure_media_url_security(
|
configure_media_url_security(
|
||||||
server_args.allowed_media_domains,
|
get_mm().allowed_media_domains,
|
||||||
server_args.media_url_max_file_size_mb,
|
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 = (
|
self.mm_feature_transport = (
|
||||||
configured_mm_feature_transport
|
configured_mm_feature_transport
|
||||||
if configured_mm_feature_transport in ("cpu", "cuda_ipc", "cuda_vmm")
|
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.disable_fast_image_processor = self.image_processor_backend == "pil"
|
||||||
self.skip_tokenizer_init = server_args.skip_tokenizer_init
|
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.image_config = mm_process_config.get("image", {})
|
||||||
self.video_config = mm_process_config.get("video", {})
|
self.video_config = mm_process_config.get("video", {})
|
||||||
self.audio_config = mm_process_config.get("audio", {})
|
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
|
# The fingerprint is needed only to build artifact keys. Avoid inspecting
|
||||||
# processor state when this processor will never retain artifacts.
|
# processor state when this processor will never retain artifacts.
|
||||||
self.processor_fingerprint = (
|
self.processor_fingerprint = (
|
||||||
build_processor_fingerprint(self, hf_config, server_args)
|
build_processor_fingerprint(self, hf_config)
|
||||||
if self.mm_preprocess_cache.enabled
|
if self.mm_preprocess_cache.enabled
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -39,6 +39,7 @@ from sglang.srt.multimodal.processors.mimo_audio import (
|
|||||||
MiMoAudioPipeline,
|
MiMoAudioPipeline,
|
||||||
)
|
)
|
||||||
from sglang.srt.multimodal.processors.qwen_vl import smart_nframes
|
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 import ImageData, VideoData
|
||||||
from sglang.srt.utils.common import download_remote_media
|
from sglang.srt.utils.common import download_remote_media
|
||||||
from sglang.utils import logger
|
from sglang.utils import logger
|
||||||
@@ -1588,7 +1589,7 @@ class MiMoV2Processor(BaseMultimodalProcessor):
|
|||||||
processor_config, "video_end_token_id"
|
processor_config, "video_end_token_id"
|
||||||
)
|
)
|
||||||
self.use_image_processor_gpu = envs.SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU.get()
|
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(
|
self.mimo_processor = MiMoProcessor(
|
||||||
tokenizer=self._processor.tokenizer,
|
tokenizer=self._processor.tokenizer,
|
||||||
|
|||||||
@@ -31,7 +31,11 @@ from sglang.srt.ray.engine import (
|
|||||||
_get_bundle_node_ip,
|
_get_bundle_node_ip,
|
||||||
_resolve_bundle_indices,
|
_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.server_args import PortArgs, ServerArgs
|
||||||
from sglang.srt.utils.network import bind_port, get_zmq_socket, get_zmq_socket_on_host
|
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):
|
for node_idx in range(nnodes):
|
||||||
bundle_idx = self.bundle_for_node[node_idx]
|
bundle_idx = self.bundle_for_node[node_idx]
|
||||||
pp_range, tp_range, pp_per_node, tp_per_node = _calculate_rank_ranges(
|
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 pp_rank in pp_range:
|
||||||
for tp_rank in tp_range:
|
for tp_rank in tp_range:
|
||||||
@@ -161,7 +168,7 @@ class RayDataParallelController(DataParallelController):
|
|||||||
tp_rank,
|
tp_rank,
|
||||||
server_args.tp_size,
|
server_args.tp_size,
|
||||||
get_parallel().dp_size,
|
get_parallel().dp_size,
|
||||||
server_args.attn_cp_size,
|
configured_attn_cp_size(),
|
||||||
)
|
)
|
||||||
rank_port_args = PortArgs.init_new(
|
rank_port_args = PortArgs.init_new(
|
||||||
server_args, actual_dp_rank, worker_ports
|
server_args, actual_dp_rank, worker_ports
|
||||||
@@ -202,7 +209,7 @@ class RayDataParallelController(DataParallelController):
|
|||||||
world_size = _compute_world_size(server_args)
|
world_size = _compute_world_size(server_args)
|
||||||
bundle_indices = _resolve_bundle_indices(self.pg, world_size)
|
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:
|
if dp_rank is not None:
|
||||||
start_rank = dp_rank * ranks_per_tp_group
|
start_rank = dp_rank * ranks_per_tp_group
|
||||||
end_rank = start_rank + ranks_per_tp_group
|
end_rank = start_rank + ranks_per_tp_group
|
||||||
@@ -232,7 +239,7 @@ class RayDataParallelController(DataParallelController):
|
|||||||
tp_rank,
|
tp_rank,
|
||||||
server_args.tp_size,
|
server_args.tp_size,
|
||||||
get_parallel().dp_size,
|
get_parallel().dp_size,
|
||||||
server_args.attn_cp_size,
|
configured_attn_cp_size(),
|
||||||
)
|
)
|
||||||
rank_port_args = PortArgs.init_new(
|
rank_port_args = PortArgs.init_new(
|
||||||
server_args, actual_dp_rank, worker_ports
|
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.environ import envs
|
||||||
from sglang.srt.ray.scheduler_actor import SchedulerActor
|
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
|
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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.
|
Normal: dp_size * tp_size * pp_size; DP attention: tp_size * pp_size.
|
||||||
"""
|
"""
|
||||||
if get_parallel().enable_dp_attention:
|
if get_parallel().enable_dp_attention:
|
||||||
return 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 * server_args.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]:
|
def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]:
|
||||||
@@ -269,10 +269,10 @@ class RayEngine(Engine):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if get_parallel().enable_dp_attention:
|
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:
|
else:
|
||||||
total_gpus = (
|
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
|
nnodes = server_args.nnodes
|
||||||
@@ -332,7 +332,7 @@ class RayEngine(Engine):
|
|||||||
pp_range, tp_range, pp_per_node, tp_per_node = (
|
pp_range, tp_range, pp_per_node, tp_per_node = (
|
||||||
_calculate_rank_ranges(
|
_calculate_rank_ranges(
|
||||||
nnodes,
|
nnodes,
|
||||||
server_args.pp_size,
|
configured_pp_size(),
|
||||||
server_args.tp_size,
|
server_args.tp_size,
|
||||||
node_rank=node_idx,
|
node_rank=node_idx,
|
||||||
)
|
)
|
||||||
@@ -449,16 +449,16 @@ class RayEngine(Engine):
|
|||||||
|
|
||||||
if get_parallel().enable_dp_attention:
|
if get_parallel().enable_dp_attention:
|
||||||
# DP attention folds DP into TP — total GPUs = tp_size * pp_size
|
# 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:
|
else:
|
||||||
total_gpus = (
|
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
|
gpus_per_node = total_gpus // server_args.nnodes
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Ray DP cluster: {server_args.nnodes} nodes, "
|
f"Ray DP cluster: {server_args.nnodes} nodes, "
|
||||||
f"{gpus_per_node} GPUs/node, dp_size={get_parallel().dp_size}, "
|
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}"
|
f"enable_dp_attention={get_parallel().enable_dp_attention}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -2,13 +2,13 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import os
|
import os
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
|
||||||
|
|
||||||
os.environ["SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE"] = "1"
|
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.install import install_token_oracle_from_env
|
||||||
from sglang.srt.kv_canary.token_oracle.oracle import HashOracle
|
from sglang.srt.kv_canary.token_oracle.oracle import HashOracle
|
||||||
from sglang.srt.layers.sampler import _CUSTOM_SAMPLER_FACTORIES
|
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.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
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")
|
register_amd_ci(est_time=60, suite="extra-a-test-1-gpu-small-amd")
|
||||||
|
|
||||||
|
|
||||||
def _make_server_args(*, sampling_backend: str) -> SimpleNamespace:
|
def _publish(case, *, sampling_backend: str) -> None:
|
||||||
return SimpleNamespace(sampling_backend=sampling_backend)
|
"""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):
|
class TestInstallTokenOracleFromEnv(CustomTestCase):
|
||||||
def test_install_token_oracle_from_env_disabled_returns_none(self) -> None:
|
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."""
|
"""Verify server-arg-disabled token oracle installation (sampling_backend != 'token_oracle') returns no TokenOracleManager."""
|
||||||
server_args = _make_server_args(sampling_backend="auto")
|
_publish(self, sampling_backend="auto")
|
||||||
hook = install_token_oracle_from_env(server_args=server_args, vocab_size=1000)
|
hook = install_token_oracle_from_env(vocab_size=1000)
|
||||||
self.assertIsNone(hook)
|
self.assertIsNone(hook)
|
||||||
|
|
||||||
def test_install_token_oracle_from_env_enabled_registers_oracle_backend(
|
def test_install_token_oracle_from_env_enabled_registers_oracle_backend(
|
||||||
self,
|
self,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Verify token oracle installation via sampling_backend='token_oracle' registers the oracle backend."""
|
"""Verify token oracle installation via sampling_backend='token_oracle' registers the oracle backend."""
|
||||||
server_args = _make_server_args(sampling_backend="token_oracle")
|
_publish(self, sampling_backend="token_oracle")
|
||||||
hook = install_token_oracle_from_env(server_args=server_args, vocab_size=512)
|
hook = install_token_oracle_from_env(vocab_size=512)
|
||||||
self.assertIsNotNone(hook)
|
self.assertIsNotNone(hook)
|
||||||
self.assertIn("token_oracle", _CUSTOM_SAMPLER_FACTORIES)
|
self.assertIn("token_oracle", _CUSTOM_SAMPLER_FACTORIES)
|
||||||
|
|
||||||
@@ -40,8 +43,8 @@ class TestInstallTokenOracleFromEnv(CustomTestCase):
|
|||||||
self,
|
self,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Verify token oracle installation via sampling_backend='token_oracle' returns a TokenOracleManager wrapping a HashOracle."""
|
"""Verify token oracle installation via sampling_backend='token_oracle' returns a TokenOracleManager wrapping a HashOracle."""
|
||||||
server_args = _make_server_args(sampling_backend="token_oracle")
|
_publish(self, sampling_backend="token_oracle")
|
||||||
hook = install_token_oracle_from_env(server_args=server_args, vocab_size=256)
|
hook = install_token_oracle_from_env(vocab_size=256)
|
||||||
self.assertIsNotNone(hook)
|
self.assertIsNotNone(hook)
|
||||||
self.assertIsInstance(hook.oracle, HashOracle)
|
self.assertIsInstance(hook.oracle, HashOracle)
|
||||||
self.assertEqual(hook.oracle.vocab_size, 256)
|
self.assertEqual(hook.oracle.vocab_size, 256)
|
||||||
|
|||||||
@@ -28,6 +28,7 @@ from sglang.srt.constrained.base_grammar_backend import (
|
|||||||
create_grammar_backend,
|
create_grammar_backend,
|
||||||
register_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
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
register_cpu_ci(2.0, "base-a-test-cpu")
|
register_cpu_ci(2.0, "base-a-test-cpu")
|
||||||
@@ -231,18 +232,39 @@ class TestCreateGrammarBackend(unittest.TestCase):
|
|||||||
GRAMMAR_BACKEND_REGISTRY.clear()
|
GRAMMAR_BACKEND_REGISTRY.clear()
|
||||||
GRAMMAR_BACKEND_REGISTRY.update(self._saved)
|
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(
|
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 = MagicMock()
|
||||||
args.override = lambda source, **updates: [
|
args.override = lambda source, **updates: [
|
||||||
setattr(args, key, value) for key, value in updates.items()
|
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
|
return args
|
||||||
|
|
||||||
def test_none_backend_returns_none(self):
|
def test_none_backend_returns_none(self):
|
||||||
@@ -293,8 +315,9 @@ class TestCreateGrammarBackend(unittest.TestCase):
|
|||||||
def test_outlines_backend(self, mock_outlines_cls):
|
def test_outlines_backend(self, mock_outlines_cls):
|
||||||
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
||||||
mock_outlines_cls.return_value = mock_backend
|
mock_outlines_cls.return_value = mock_backend
|
||||||
args = self._make_server_args("outlines")
|
args = self._make_server_args(
|
||||||
args.constrained_json_whitespace_pattern = r"\s*"
|
"outlines", constrained_json_whitespace_pattern=r"\s*"
|
||||||
|
)
|
||||||
|
|
||||||
result = create_grammar_backend(args, "tok", 32000)
|
result = create_grammar_backend(args, "tok", 32000)
|
||||||
mock_outlines_cls.assert_called_once_with("tok", whitespace_pattern=r"\s*")
|
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):
|
def test_xgrammar_backend(self, mock_xgrammar_cls):
|
||||||
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
||||||
mock_xgrammar_cls.return_value = mock_backend
|
mock_xgrammar_cls.return_value = mock_backend
|
||||||
args = self._make_server_args("xgrammar")
|
args = self._make_server_args(
|
||||||
args.constrained_json_disable_any_whitespace = True
|
"xgrammar", constrained_json_disable_any_whitespace=True
|
||||||
|
)
|
||||||
|
|
||||||
result = create_grammar_backend(args, "tok", 32000, {1, 2})
|
result = create_grammar_backend(args, "tok", 32000, {1, 2})
|
||||||
mock_xgrammar_cls.assert_called_once_with(
|
mock_xgrammar_cls.assert_called_once_with(
|
||||||
@@ -336,9 +360,11 @@ class TestCreateGrammarBackend(unittest.TestCase):
|
|||||||
def test_llguidance_backend(self, mock_guidance_cls):
|
def test_llguidance_backend(self, mock_guidance_cls):
|
||||||
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
mock_backend = MagicMock(spec=BaseGrammarBackend)
|
||||||
mock_guidance_cls.return_value = mock_backend
|
mock_guidance_cls.return_value = mock_backend
|
||||||
args = self._make_server_args("llguidance")
|
args = self._make_server_args(
|
||||||
args.constrained_json_disable_any_whitespace = False
|
"llguidance",
|
||||||
args.constrained_json_whitespace_pattern = r"\s+"
|
constrained_json_disable_any_whitespace=False,
|
||||||
|
constrained_json_whitespace_pattern=r"\s+",
|
||||||
|
)
|
||||||
|
|
||||||
result = create_grammar_backend(args, "tok", 32000, {1, 2})
|
result = create_grammar_backend(args, "tok", 32000, {1, 2})
|
||||||
mock_guidance_cls.assert_called_once_with(
|
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,
|
materialize_kimi_k3_cpu_features,
|
||||||
prepare_kimi_k3_encoder_inputs,
|
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.server_args import resolve_encoder_transfer_backend
|
||||||
from sglang.srt.utils import ImageData
|
from sglang.srt.utils import ImageData
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
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():
|
def test_epd_language_only_rejects_missing_dispatched_embedding():
|
||||||
server_args = SimpleNamespace(
|
override = get_context().override_server_args(
|
||||||
language_only=True,
|
language_only=True,
|
||||||
encoder_transfer_backend="zmq_to_tokenizer",
|
encoder_transfer_backend="zmq_to_tokenizer",
|
||||||
)
|
)
|
||||||
|
override.install()
|
||||||
|
try:
|
||||||
request = SimpleNamespace(need_wait_for_mm_inputs=True)
|
request = SimpleNamespace(need_wait_for_mm_inputs=True)
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
_reject_missing_dispatched_encoder_embedding(server_args, request, None)
|
_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():
|
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
|
The record is produced by actual resolution -- a language-only Kimi-K3
|
||||||
launch at TP2, whose `encoder_transfer_backend` starts at the argument
|
launch at TP2, whose `encoder_transfer_backend` starts at the argument
|
||||||
default `"auto"` (`ENCODER_TRANSFER_BACKEND_CHOICES[0]`) and is filled in
|
default `"auto"` (`ENCODER_TRANSFER_BACKEND_CHOICES[0]`) and is filled in
|
||||||
by `resolve_encoder_transfer_backend` to `"zmq_to_tokenizer"`. Today the
|
by `resolve_encoder_transfer_backend` to `"zmq_to_tokenizer"`. The guard
|
||||||
guard therefore rejects. When step 12 makes the instance raw, this same
|
reads that resolved value out of the published bags, so the rejection
|
||||||
launch hands the guard a record still at `"auto"`, the rejection silently
|
survives the record going raw: what a reader must never do is go back to
|
||||||
stops, and *this test fails* -- which is the signal to give this reader
|
the record for this field.
|
||||||
the resolved value (per-engine overlay or bag) rather than the record.
|
|
||||||
Fixed doubles cannot trip on that change, so the record here must come
|
Fixed doubles cannot trip on that change, so the record here must come
|
||||||
from resolution, not a SimpleNamespace.
|
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)
|
shutil.rmtree(config_dir, ignore_errors=True)
|
||||||
|
|
||||||
assert resolved.encoder_transfer_backend == "zmq_to_tokenizer"
|
assert resolved.encoder_transfer_backend == "zmq_to_tokenizer"
|
||||||
|
# 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)
|
request = SimpleNamespace(need_wait_for_mm_inputs=True)
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
with pytest.raises(HTTPException) as exc_info:
|
||||||
_reject_missing_dispatched_encoder_embedding(resolved, request, None)
|
_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:
|
||||||
|
reset_context()
|
||||||
|
|
||||||
|
|
||||||
def test_epd_allows_local_processing_when_request_was_not_dispatched():
|
def test_epd_allows_local_processing_when_request_was_not_dispatched():
|
||||||
server_args = SimpleNamespace(
|
override = get_context().override_server_args(
|
||||||
language_only=True,
|
language_only=True,
|
||||||
encoder_transfer_backend="zmq_to_tokenizer",
|
encoder_transfer_backend="zmq_to_tokenizer",
|
||||||
)
|
)
|
||||||
|
override.install()
|
||||||
|
try:
|
||||||
request = SimpleNamespace(need_wait_for_mm_inputs=False)
|
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"):
|
def _encoder(model_type="kimi_k3"):
|
||||||
|
|||||||
@@ -8,13 +8,21 @@ from sglang.srt.managers.scheduler_components.batch_result_processor import (
|
|||||||
SchedulerBatchResultProcessor,
|
SchedulerBatchResultProcessor,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
|
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.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
|
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 = Mock()
|
||||||
metrics_reporter.num_generated_tokens = 0
|
metrics_reporter.num_generated_tokens = 0
|
||||||
metrics_reporter.forward_ct_decode = 0
|
metrics_reporter.forward_ct_decode = 0
|
||||||
@@ -26,8 +34,6 @@ def _make_processor(server_mode: str = "full") -> SchedulerBatchResultProcessor:
|
|||||||
server_args=SimpleNamespace(
|
server_args=SimpleNamespace(
|
||||||
enable_metrics=False,
|
enable_metrics=False,
|
||||||
enable_hisparse=False,
|
enable_hisparse=False,
|
||||||
enable_return_hidden_states=True,
|
|
||||||
return_hidden_states_mode=server_mode,
|
|
||||||
),
|
),
|
||||||
model_config=SimpleNamespace(think_end_ids=None),
|
model_config=SimpleNamespace(think_end_ids=None),
|
||||||
token_to_kv_pool_allocator=Mock(),
|
token_to_kv_pool_allocator=Mock(),
|
||||||
@@ -139,7 +145,7 @@ class TestPrefillHiddenStateOffsets(CustomTestCase):
|
|||||||
can_run_cuda_graph=False,
|
can_run_cuda_graph=False,
|
||||||
skipped_output_comm=False,
|
skipped_output_comm=False,
|
||||||
)
|
)
|
||||||
processor = _make_processor(server_mode)
|
processor = _make_processor(self, server_mode)
|
||||||
|
|
||||||
with (
|
with (
|
||||||
patch(
|
patch(
|
||||||
@@ -160,7 +166,7 @@ class TestPrefillHiddenStateOffsets(CustomTestCase):
|
|||||||
|
|
||||||
class TestDecodeHiddenStateRetention(CustomTestCase):
|
class TestDecodeHiddenStateRetention(CustomTestCase):
|
||||||
def test_last_mode_multi_step_storage_stays_bounded(self):
|
def test_last_mode_multi_step_storage_stays_bounded(self):
|
||||||
processor = _make_processor()
|
processor = _make_processor(self)
|
||||||
req = _DecodeReq()
|
req = _DecodeReq()
|
||||||
batch = SimpleNamespace(
|
batch = SimpleNamespace(
|
||||||
reqs=[req],
|
reqs=[req],
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ from unittest.mock import Mock
|
|||||||
|
|
||||||
from sglang.srt.managers.io_struct import GenerateReqInput
|
from sglang.srt.managers.io_struct import GenerateReqInput
|
||||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
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.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
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):
|
class TestHiddenStateServerMode(CustomTestCase):
|
||||||
@staticmethod
|
def _make_tokenizer_manager(self, mode):
|
||||||
def _make_tokenizer_manager(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 = TokenizerManager.__new__(TokenizerManager)
|
||||||
manager.context_len = 128
|
manager.context_len = 128
|
||||||
manager.num_reserved_tokens = 0
|
manager.num_reserved_tokens = 0
|
||||||
manager.allow_auto_truncate = False
|
manager.allow_auto_truncate = False
|
||||||
manager.validate_total_tokens = False
|
manager.validate_total_tokens = False
|
||||||
manager.is_generation = True
|
manager.is_generation = True
|
||||||
manager.server_args = SimpleNamespace(
|
manager.server_args = SimpleNamespace(enable_custom_logit_processor=False)
|
||||||
enable_return_hidden_states=mode is not None,
|
|
||||||
return_hidden_states_mode=mode,
|
|
||||||
enable_custom_logit_processor=False,
|
|
||||||
)
|
|
||||||
manager._validate_token_ids_logprob = Mock()
|
manager._validate_token_ids_logprob = Mock()
|
||||||
return manager
|
return manager
|
||||||
|
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from unittest.mock import MagicMock, patch
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
@@ -80,14 +81,20 @@ class TestBaseProcessorConfigExtraction(CustomTestCase):
|
|||||||
BaseMultimodalProcessor,
|
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 = MagicMock()
|
||||||
server_args.mm_process_config = mm_process_config
|
|
||||||
server_args.mm_processor_worker_num = mm_processor_worker_num
|
server_args.mm_processor_worker_num = mm_processor_worker_num
|
||||||
server_args.mm_io_worker_num = mm_io_worker_num
|
server_args.mm_io_worker_num = mm_io_worker_num
|
||||||
server_args.mm_preprocess_cache_size_mb = None
|
server_args.mm_preprocess_cache_size_mb = None
|
||||||
server_args.tokenizer_worker_num = 1
|
server_args.tokenizer_worker_num = 1
|
||||||
server_args.trust_mm_content_hashes = False
|
server_args.trust_mm_content_hashes = False
|
||||||
server_args.allowed_media_domains = []
|
|
||||||
server_args.media_url_max_file_size_mb = 64
|
server_args.media_url_max_file_size_mb = 64
|
||||||
|
|
||||||
hf_config = MagicMock()
|
hf_config = MagicMock()
|
||||||
@@ -170,8 +177,14 @@ class TestBaseProcessorConfigExtraction(CustomTestCase):
|
|||||||
|
|
||||||
|
|
||||||
class TestMultimodalFeatureTransportRuntime(CustomTestCase):
|
class TestMultimodalFeatureTransportRuntime(CustomTestCase):
|
||||||
@staticmethod
|
def _server_args(self, mm_feature_transport):
|
||||||
def _server_args(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(
|
return SimpleNamespace(
|
||||||
mm_feature_transport=mm_feature_transport,
|
mm_feature_transport=mm_feature_transport,
|
||||||
image_processor_backend="auto",
|
image_processor_backend="auto",
|
||||||
@@ -197,8 +210,7 @@ class TestMultimodalFeatureTransportRuntime(CustomTestCase):
|
|||||||
return processor
|
return processor
|
||||||
|
|
||||||
def test_cuda_ipc_pool_uses_resolved_server_arg(self):
|
def test_cuda_ipc_pool_uses_resolved_server_arg(self):
|
||||||
# The processor module can be imported before this instance is built;
|
# Transport policy resolves from the mm bag, so the test publishes it.
|
||||||
# transport policy must still resolve from the instance's ServerArgs.
|
|
||||||
from sglang.srt.multimodal.processors import base_processor
|
from sglang.srt.multimodal.processors import base_processor
|
||||||
|
|
||||||
with (
|
with (
|
||||||
@@ -792,16 +804,21 @@ class TestDoubleBosGuard(CustomTestCase):
|
|||||||
BaseMultimodalProcessor,
|
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 = MagicMock()
|
||||||
server_args.mm_process_config = {}
|
|
||||||
server_args.mm_processor_worker_num = 0
|
server_args.mm_processor_worker_num = 0
|
||||||
server_args.mm_io_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.disable_fast_image_processor = True
|
||||||
server_args.mm_preprocess_cache_size_mb = None
|
server_args.mm_preprocess_cache_size_mb = None
|
||||||
server_args.tokenizer_worker_num = 1
|
server_args.tokenizer_worker_num = 1
|
||||||
server_args.trust_mm_content_hashes = False
|
server_args.trust_mm_content_hashes = False
|
||||||
server_args.allowed_media_domains = []
|
|
||||||
server_args.media_url_max_file_size_mb = 64
|
server_args.media_url_max_file_size_mb = 64
|
||||||
|
|
||||||
mock_hf_processor = MagicMock()
|
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
|
from sglang.srt.observability.req_time_stats import ( # noqa: E402
|
||||||
APIServerReqTimeStats,
|
APIServerReqTimeStats,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
|
|
||||||
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
|
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:
|
def _make_tokenizer_manager(case) -> TokenizerManager:
|
||||||
"""Create a TokenizerManager with mocked dependencies, bypassing __init__."""
|
"""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 = TokenizerManager.__new__(TokenizerManager)
|
||||||
tm.server_args = MagicMock()
|
tm.server_args = MagicMock()
|
||||||
tm._config_updates = []
|
tm._config_updates = []
|
||||||
@@ -212,7 +220,7 @@ class TestRidToStateCleanupOnAbort(CustomTestCase):
|
|||||||
|
|
||||||
def test_abort_removes_rid_from_state(self):
|
def test_abort_removes_rid_from_state(self):
|
||||||
"""After _handle_abort_req, rid should be removed from rid_to_state."""
|
"""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"
|
rid = "abort_test_rid"
|
||||||
state = _make_req_state(rid)
|
state = _make_req_state(rid)
|
||||||
tm.rid_to_state[rid] = state
|
tm.rid_to_state[rid] = state
|
||||||
@@ -224,7 +232,7 @@ class TestRidToStateCleanupOnAbort(CustomTestCase):
|
|||||||
|
|
||||||
def test_abort_allows_resubmit_same_rid(self):
|
def test_abort_allows_resubmit_same_rid(self):
|
||||||
"""After abort, _init_req_state should accept the same rid again."""
|
"""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"
|
rid = "resubmit_after_abort_rid"
|
||||||
state = _make_req_state(rid)
|
state = _make_req_state(rid)
|
||||||
tm.rid_to_state[rid] = state
|
tm.rid_to_state[rid] = state
|
||||||
@@ -245,7 +253,7 @@ class TestRidToStateCleanupOnAbort(CustomTestCase):
|
|||||||
|
|
||||||
def test_abort_sets_finished_and_notifies(self):
|
def test_abort_sets_finished_and_notifies(self):
|
||||||
"""_handle_abort_req should mark state as finished and set the event."""
|
"""_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"
|
rid = "abort_notify_rid"
|
||||||
state = _make_req_state(rid)
|
state = _make_req_state(rid)
|
||||||
tm.rid_to_state[rid] = state
|
tm.rid_to_state[rid] = state
|
||||||
@@ -266,7 +274,7 @@ class TestRidToStateCleanupOnBatchOutput(CustomTestCase):
|
|||||||
|
|
||||||
def test_batch_output_removes_rid_on_finish(self):
|
def test_batch_output_removes_rid_on_finish(self):
|
||||||
"""When a request finishes in _handle_batch_output, rid should be removed."""
|
"""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"
|
rid = "batch_finish_rid"
|
||||||
state = _make_req_state(rid)
|
state = _make_req_state(rid)
|
||||||
tm.rid_to_state[rid] = state
|
tm.rid_to_state[rid] = state
|
||||||
@@ -278,7 +286,7 @@ class TestRidToStateCleanupOnBatchOutput(CustomTestCase):
|
|||||||
|
|
||||||
def test_batch_output_allows_resubmit_after_finish(self):
|
def test_batch_output_allows_resubmit_after_finish(self):
|
||||||
"""After a request finishes, the same rid can be resubmitted."""
|
"""After a request finishes, the same rid can be resubmitted."""
|
||||||
tm = _make_tokenizer_manager()
|
tm = _make_tokenizer_manager(self)
|
||||||
rid = "batch_resubmit_rid"
|
rid = "batch_resubmit_rid"
|
||||||
state = _make_req_state(rid)
|
state = _make_req_state(rid)
|
||||||
tm.rid_to_state[rid] = state
|
tm.rid_to_state[rid] = state
|
||||||
@@ -299,7 +307,7 @@ class TestRidToStateCleanupOnBatchOutput(CustomTestCase):
|
|||||||
|
|
||||||
def test_batch_output_keeps_rid_when_not_finished(self):
|
def test_batch_output_keeps_rid_when_not_finished(self):
|
||||||
"""When a request is not yet finished, rid should remain in rid_to_state."""
|
"""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"
|
rid = "batch_ongoing_rid"
|
||||||
state = _make_req_state(rid)
|
state = _make_req_state(rid)
|
||||||
tm.rid_to_state[rid] = state
|
tm.rid_to_state[rid] = state
|
||||||
@@ -316,7 +324,7 @@ class TestInitReqStateDuplicateDetection(CustomTestCase):
|
|||||||
|
|
||||||
def test_duplicate_rid_raises_error(self):
|
def test_duplicate_rid_raises_error(self):
|
||||||
"""_init_req_state should raise ValueError if rid already exists."""
|
"""_init_req_state should raise ValueError if rid already exists."""
|
||||||
tm = _make_tokenizer_manager()
|
tm = _make_tokenizer_manager(self)
|
||||||
rid = "duplicate_rid"
|
rid = "duplicate_rid"
|
||||||
state = _make_req_state(rid)
|
state = _make_req_state(rid)
|
||||||
tm.rid_to_state[rid] = state
|
tm.rid_to_state[rid] = state
|
||||||
@@ -334,7 +342,7 @@ class TestInitReqStateDuplicateDetection(CustomTestCase):
|
|||||||
|
|
||||||
def test_unique_rid_succeeds(self):
|
def test_unique_rid_succeeds(self):
|
||||||
"""_init_req_state should succeed with a unique rid."""
|
"""_init_req_state should succeed with a unique rid."""
|
||||||
tm = _make_tokenizer_manager()
|
tm = _make_tokenizer_manager(self)
|
||||||
rid = "unique_rid"
|
rid = "unique_rid"
|
||||||
|
|
||||||
obj = Mock(spec=GenerateReqInput)
|
obj = Mock(spec=GenerateReqInput)
|
||||||
@@ -353,7 +361,7 @@ class TestResubmitAfterCompletion(CustomTestCase):
|
|||||||
|
|
||||||
def test_complete_then_resubmit_same_rid(self):
|
def test_complete_then_resubmit_same_rid(self):
|
||||||
"""A request that completes normally should allow resubmission with the same rid."""
|
"""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"
|
rid = "complete_resubmit_rid"
|
||||||
|
|
||||||
# Phase 1: simulate a request in rid_to_state, then complete it
|
# 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):
|
def test_abort_then_resubmit_same_rid(self):
|
||||||
"""An aborted request should allow resubmission with the same rid."""
|
"""An aborted request should allow resubmission with the same rid."""
|
||||||
tm = _make_tokenizer_manager()
|
tm = _make_tokenizer_manager(self)
|
||||||
rid = "abort_resubmit_rid"
|
rid = "abort_resubmit_rid"
|
||||||
|
|
||||||
# Phase 1: simulate a request, then abort it
|
# Phase 1: simulate a request, then abort it
|
||||||
@@ -413,9 +421,9 @@ class _DummyAsyncCM:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
def _make_tm_for_generate() -> TokenizerManager:
|
def _make_tm_for_generate(case) -> TokenizerManager:
|
||||||
"""Augment the mocked TokenizerManager with what generate_request needs."""
|
"""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.language_only = False
|
||||||
tm.server_args.tokenizer_worker_num = 1
|
tm.server_args.tokenizer_worker_num = 1
|
||||||
tm.server_args.enable_strict_thinking = False
|
tm.server_args.enable_strict_thinking = False
|
||||||
@@ -450,7 +458,7 @@ class TestDiscardPendingReqStates(CustomTestCase):
|
|||||||
"""Direct tests for _discard_pending_req_states."""
|
"""Direct tests for _discard_pending_req_states."""
|
||||||
|
|
||||||
def test_discard_single(self):
|
def test_discard_single(self):
|
||||||
tm = _make_tokenizer_manager()
|
tm = _make_tokenizer_manager(self)
|
||||||
rid = "d_single"
|
rid = "d_single"
|
||||||
tm.rid_to_state[rid] = _make_req_state(rid)
|
tm.rid_to_state[rid] = _make_req_state(rid)
|
||||||
obj = Mock(spec=GenerateReqInput)
|
obj = Mock(spec=GenerateReqInput)
|
||||||
@@ -460,7 +468,7 @@ class TestDiscardPendingReqStates(CustomTestCase):
|
|||||||
self.assertNotIn(rid, tm.rid_to_state)
|
self.assertNotIn(rid, tm.rid_to_state)
|
||||||
|
|
||||||
def test_discard_batch_removes_all(self):
|
def test_discard_batch_removes_all(self):
|
||||||
tm = _make_tokenizer_manager()
|
tm = _make_tokenizer_manager(self)
|
||||||
rids = ["d0", "d1", "d2"]
|
rids = ["d0", "d1", "d2"]
|
||||||
for r in rids:
|
for r in rids:
|
||||||
tm.rid_to_state[r] = _make_req_state(r)
|
tm.rid_to_state[r] = _make_req_state(r)
|
||||||
@@ -473,7 +481,7 @@ class TestDiscardPendingReqStates(CustomTestCase):
|
|||||||
|
|
||||||
def test_discard_ignores_already_removed(self):
|
def test_discard_ignores_already_removed(self):
|
||||||
"""Popping a rid that is no longer present must not raise."""
|
"""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")
|
tm.rid_to_state["p1"] = _make_req_state("p1")
|
||||||
obj = Mock(spec=GenerateReqInput)
|
obj = Mock(spec=GenerateReqInput)
|
||||||
obj.is_single = False
|
obj.is_single = False
|
||||||
@@ -484,7 +492,7 @@ class TestDiscardPendingReqStates(CustomTestCase):
|
|||||||
|
|
||||||
class TestParallelStreamTaskCleanup(CustomTestCase):
|
class TestParallelStreamTaskCleanup(CustomTestCase):
|
||||||
def test_failing_choice_cancels_and_closes_sibling_waiters(self):
|
def test_failing_choice_cancels_and_closes_sibling_waiters(self):
|
||||||
tm = _make_tokenizer_manager()
|
tm = _make_tokenizer_manager(self)
|
||||||
|
|
||||||
async def drive():
|
async def drive():
|
||||||
sibling_closed = asyncio.Event()
|
sibling_closed = asyncio.Event()
|
||||||
@@ -512,7 +520,7 @@ class TestParallelStreamTaskCleanup(CustomTestCase):
|
|||||||
asyncio.run(drive())
|
asyncio.run(drive())
|
||||||
|
|
||||||
def test_failing_non_stream_choice_cancels_and_closes_sibling_waiters(self):
|
def test_failing_non_stream_choice_cancels_and_closes_sibling_waiters(self):
|
||||||
tm = _make_tokenizer_manager()
|
tm = _make_tokenizer_manager(self)
|
||||||
|
|
||||||
async def drive():
|
async def drive():
|
||||||
sibling_closed = asyncio.Event()
|
sibling_closed = asyncio.Event()
|
||||||
@@ -546,7 +554,7 @@ class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def test_single_failure_before_dispatch_cleans_up(self):
|
def test_single_failure_before_dispatch_cleans_up(self):
|
||||||
tm = _make_tm_for_generate()
|
tm = _make_tm_for_generate(self)
|
||||||
rid = "single_overlen"
|
rid = "single_overlen"
|
||||||
obj = _make_generate_obj(rid, is_single=True)
|
obj = _make_generate_obj(rid, is_single=True)
|
||||||
# Simulate over-length rejection during tokenization/validation.
|
# Simulate over-length rejection during tokenization/validation.
|
||||||
@@ -566,7 +574,7 @@ class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase):
|
|||||||
self.assertNotIn(rid, tm.rid_to_state)
|
self.assertNotIn(rid, tm.rid_to_state)
|
||||||
|
|
||||||
def test_batch_failure_before_dispatch_cleans_up_all(self):
|
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"]
|
rids = ["b0", "b1", "b2"]
|
||||||
obj = _make_generate_obj(list(rids), is_single=False)
|
obj = _make_generate_obj(list(rids), is_single=False)
|
||||||
|
|
||||||
@@ -588,7 +596,7 @@ class TestGenerateRequestCleanupOnDispatchFailure(CustomTestCase):
|
|||||||
self.assertNotIn(r, tm.rid_to_state)
|
self.assertNotIn(r, tm.rid_to_state)
|
||||||
|
|
||||||
def test_thinking_budget_rejects_runtime_without_strict_thinking(self):
|
def test_thinking_budget_rejects_runtime_without_strict_thinking(self):
|
||||||
tm = _make_tm_for_generate()
|
tm = _make_tm_for_generate(self)
|
||||||
obj = GenerateReqInput(
|
obj = GenerateReqInput(
|
||||||
text="hello",
|
text="hello",
|
||||||
rid="thinking-budget",
|
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 (
|
from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import (
|
||||||
PrefillCudaGraphRunner,
|
PrefillCudaGraphRunner,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
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):
|
class TestHiddenStateGraphRecapture(CustomTestCase):
|
||||||
def test_server_mode_sets_graph_capture_ceiling(self):
|
def test_server_mode_sets_graph_capture_ceiling(self):
|
||||||
disabled = SimpleNamespace(
|
cases = (
|
||||||
enable_return_hidden_states=False,
|
(dict(enable_return_hidden_states=False), CaptureHiddenMode.NULL),
|
||||||
return_hidden_states_mode=None,
|
(
|
||||||
)
|
dict(
|
||||||
last = SimpleNamespace(
|
enable_return_hidden_states=True, return_hidden_states_mode="last"
|
||||||
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,
|
CaptureHiddenMode.LAST,
|
||||||
)
|
),
|
||||||
self.assertEqual(
|
(
|
||||||
get_server_return_hidden_states_mode(full),
|
dict(
|
||||||
|
enable_return_hidden_states=True, return_hidden_states_mode="full"
|
||||||
|
),
|
||||||
CaptureHiddenMode.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
|
@staticmethod
|
||||||
def _make_runner(runner_cls, capture_hidden_mode):
|
def _make_runner(runner_cls, capture_hidden_mode):
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import (
|
|||||||
PrefillCudaGraphRunner,
|
PrefillCudaGraphRunner,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.runner.shape_key import ShapeKey
|
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.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
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):
|
def test_eagle_target_tc_piecewise_skips_last_mode_capture(self):
|
||||||
eager_runner = object()
|
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(
|
model_runner = SimpleNamespace(
|
||||||
is_draft_worker=False,
|
is_draft_worker=False,
|
||||||
spec_algorithm=SimpleNamespace(is_eagle=lambda: True),
|
spec_algorithm=SimpleNamespace(is_eagle=lambda: True),
|
||||||
server_args=SimpleNamespace(
|
server_args=SimpleNamespace(),
|
||||||
enable_return_hidden_states=True,
|
|
||||||
return_hidden_states_mode="last",
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
with patch.object(
|
with patch.object(
|
||||||
|
|||||||
@@ -66,7 +66,7 @@ from sglang.srt.multimodal.transport.cuda_ipc import (
|
|||||||
DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY,
|
DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY,
|
||||||
CudaIpcTensorTransportProxy,
|
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.server_args import ServerArgs
|
||||||
from sglang.srt.utils import ImageData
|
from sglang.srt.utils import ImageData
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
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):
|
def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls):
|
||||||
server_args = SimpleNamespace(
|
server_args = SimpleNamespace(
|
||||||
mm_feature_transport="cpu",
|
|
||||||
image_processor_backend="auto",
|
image_processor_backend="auto",
|
||||||
disable_fast_image_processor=False,
|
disable_fast_image_processor=False,
|
||||||
skip_tokenizer_init=False,
|
skip_tokenizer_init=False,
|
||||||
mm_process_config={},
|
|
||||||
mm_io_worker_num=0,
|
mm_io_worker_num=0,
|
||||||
mm_processor_worker_num=0,
|
mm_processor_worker_num=0,
|
||||||
tokenizer_worker_num=1,
|
tokenizer_worker_num=1,
|
||||||
@@ -654,9 +652,12 @@ def test_kimi_processor_workers_clone_the_gpu_wrapper(processor_cls, wrapper_cls
|
|||||||
trust_mm_content_hashes=False,
|
trust_mm_content_hashes=False,
|
||||||
base_gpu_id=0,
|
base_gpu_id=0,
|
||||||
rl_on_policy_target=None,
|
rl_on_policy_target=None,
|
||||||
allowed_media_domains=[],
|
|
||||||
media_url_max_file_size_mb=64,
|
media_url_max_file_size_mb=64,
|
||||||
)
|
)
|
||||||
|
# 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(
|
processor = processor_cls(
|
||||||
hf_config=SimpleNamespace(media_placeholder_token_id=42),
|
hf_config=SimpleNamespace(media_placeholder_token_id=42),
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
@@ -844,6 +845,8 @@ def test_kimi_k3_normal_cache_path_connects_real_producer_to_model_consumer():
|
|||||||
tokenizer_worker_num=1,
|
tokenizer_worker_num=1,
|
||||||
mm_preprocess_cache_size_mb=1,
|
mm_preprocess_cache_size_mb=1,
|
||||||
)
|
)
|
||||||
|
publish(server_args, role="tokenizer")
|
||||||
|
try:
|
||||||
processor = KimiK3ImageProcessor(
|
processor = KimiK3ImageProcessor(
|
||||||
hf_config=hf_config,
|
hf_config=hf_config,
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
@@ -881,7 +884,8 @@ def test_kimi_k3_normal_cache_path_connects_real_producer_to_model_consumer():
|
|||||||
try:
|
try:
|
||||||
with (
|
with (
|
||||||
patch(
|
patch(
|
||||||
"sglang.srt.multimodal.processors.kimi_k3.is_cuda", return_value=True
|
"sglang.srt.multimodal.processors.kimi_k3.is_cuda",
|
||||||
|
return_value=True,
|
||||||
),
|
),
|
||||||
patch.object(
|
patch.object(
|
||||||
processor,
|
processor,
|
||||||
@@ -917,9 +921,12 @@ def test_kimi_k3_normal_cache_path_connects_real_producer_to_model_consumer():
|
|||||||
assert cold.mm_items[0].hash == hot.mm_items[0].hash
|
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.mm_items[0].offsets == hot.mm_items[0].offsets == [(3, 3)]
|
||||||
assert (
|
assert (
|
||||||
cold_items[0].model_specific_data[DEFERRED_PREPROCESSING_KEY].backend == "gpu"
|
cold_items[0].model_specific_data[DEFERRED_PREPROCESSING_KEY].backend
|
||||||
|
== "gpu"
|
||||||
)
|
)
|
||||||
torch.testing.assert_close(cold_features, hot_features)
|
torch.testing.assert_close(cold_features, hot_features)
|
||||||
|
finally:
|
||||||
|
reset_context()
|
||||||
|
|
||||||
|
|
||||||
def test_kimi_k3_model_accepts_mixed_cached_eager_and_deferred_artifacts():
|
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
|
from sglang.srt.multimodal.processors.qwen_vl import ( # noqa: E402
|
||||||
QwenVLImageProcessor,
|
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")
|
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.
|
"""A ``QwenVLImageProcessor`` over a tiny hand-built tokenizer.
|
||||||
``image_processor_cls`` picks the HF backend; they resample differently."""
|
``image_processor_cls`` picks the HF backend; they resample differently."""
|
||||||
image_processor_cls = image_processor_cls or HfQwenImageProcessor
|
image_processor_cls = image_processor_cls or HfQwenImageProcessor
|
||||||
@@ -87,6 +89,18 @@ def make_processor(config, image_processor_cls=None):
|
|||||||
allowed_media_domains=[],
|
allowed_media_domains=[],
|
||||||
media_url_max_file_size_mb=64,
|
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(
|
return QwenVLImageProcessor(
|
||||||
hf_config, server_args, processor, None, skip_mm_pool=True
|
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
|
import transformers.models.qwen2_vl as qwen2_vl
|
||||||
|
|
||||||
self.processor = make_processor(
|
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):
|
def tearDown(self):
|
||||||
|
|||||||
@@ -50,7 +50,7 @@ class TestQwenNativeMmHashes(CustomTestCase):
|
|||||||
from sglang.srt.managers.multimodal_processor import import_processors
|
from sglang.srt.managers.multimodal_processor import import_processors
|
||||||
|
|
||||||
import_processors("sglang.srt.multimodal.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):
|
def tearDown(self):
|
||||||
self.processor.io_executor.shutdown()
|
self.processor.io_executor.shutdown()
|
||||||
|
|||||||
@@ -42,6 +42,7 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
|||||||
|
|
||||||
def test_model_class_controls_cuda_vmm_opt_in(self):
|
def test_model_class_controls_cuda_vmm_opt_in(self):
|
||||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
|
|
||||||
class SupportedModel:
|
class SupportedModel:
|
||||||
supports_cuda_vmm_feature_transport = True
|
supports_cuda_vmm_feature_transport = True
|
||||||
@@ -49,8 +50,10 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
|||||||
class UnsupportedModel:
|
class UnsupportedModel:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
override = get_context().override_server_args(mm_feature_transport="cuda_vmm")
|
||||||
|
override.install()
|
||||||
|
self.addCleanup(override.restore)
|
||||||
manager = object.__new__(TokenizerManager)
|
manager = object.__new__(TokenizerManager)
|
||||||
manager.server_args = SimpleNamespace(mm_feature_transport="cuda_vmm")
|
|
||||||
manager.model_config = object()
|
manager.model_config = object()
|
||||||
|
|
||||||
with patch(
|
with patch(
|
||||||
@@ -70,9 +73,12 @@ class TestCudaVmmFeatureTransport(unittest.TestCase):
|
|||||||
|
|
||||||
def test_cpu_transport_skips_model_opt_in_lookup(self):
|
def test_cpu_transport_skips_model_opt_in_lookup(self):
|
||||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
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 = object.__new__(TokenizerManager)
|
||||||
manager.server_args = SimpleNamespace(mm_feature_transport="cpu")
|
|
||||||
manager.model_config = object()
|
manager.model_config = object()
|
||||||
|
|
||||||
with patch(
|
with patch(
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ from sglang.srt.multimodal.cache import (
|
|||||||
resolve_multimodal_item_hash,
|
resolve_multimodal_item_hash,
|
||||||
snapshot_media,
|
snapshot_media,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import publish, reset_context
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
@@ -216,26 +217,32 @@ class TestMediaIdentity(unittest.TestCase):
|
|||||||
return {"model_type": "vlm", "architectures": ["VLM"]}
|
return {"model_type": "vlm", "architectures": ["VLM"]}
|
||||||
|
|
||||||
config = Config()
|
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)
|
def fingerprint(processor, mm_process_config):
|
||||||
changed_args = ServerArgs(
|
# 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",
|
model_path="dummy",
|
||||||
revision="model-revision",
|
revision="model-revision",
|
||||||
disable_fast_image_processor=False,
|
disable_fast_image_processor=False,
|
||||||
mm_process_config={"image": {"max_pixels": 2048}},
|
mm_process_config=mm_process_config,
|
||||||
)
|
),
|
||||||
changed_config = build_processor_fingerprint(
|
role="test",
|
||||||
Processor("gpu"), config, changed_args
|
|
||||||
)
|
)
|
||||||
|
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_backend)
|
||||||
self.assertNotEqual(base, changed_config)
|
self.assertNotEqual(base, changed_config)
|
||||||
|
self.assertEqual(base, same_again)
|
||||||
|
|
||||||
def test_item_hash_namespace_covers_identity_and_processor_output(self):
|
def test_item_hash_namespace_covers_identity_and_processor_output(self):
|
||||||
digest = snapshot_media(b"image").content_digest
|
digest = snapshot_media(b"image").content_digest
|
||||||
|
|||||||
@@ -2161,6 +2161,7 @@ class TestGrpcServerArgs(CustomTestCase):
|
|||||||
tokenizer_manager=MagicMock(),
|
tokenizer_manager=MagicMock(),
|
||||||
template_manager=MagicMock(),
|
template_manager=MagicMock(),
|
||||||
scheduler_info={},
|
scheduler_info={},
|
||||||
|
grpc_port=server_args.grpc_port,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual(handle, "handle")
|
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
|
# 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.
|
# the set of call sites, so the reasons cannot drift away from the code.
|
||||||
_CONFIGURED_SIZE_CALL_SITES = {
|
_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"): (
|
("srt/layers/attention/dsa/dsa_indexer.py", "configured_pp_size"): (
|
||||||
"gates `pp_size > 1 and not get_pp_group()...`; the short circuit is the "
|
"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 "
|
"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 // "
|
"the consumer count is configured fan-out arithmetic (tp_size // "
|
||||||
"dp_size), which is what the record answered before"
|
"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"): (
|
("srt/model_loader/loader.py", "configured_moe_dp_size"): (
|
||||||
"the same dict already carries the live moe_dp_size under 'dp'; this entry "
|
"the same dict already carries the live moe_dp_size under 'dp'; this entry "
|
||||||
"is the configured intent"
|
"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"): (
|
("multimodal_gen/test/unit/test_disagg_trace.py", "_srt_trace_server_args"): (
|
||||||
"a trace fixture publishing its own context"
|
"a trace fixture publishing its own context"
|
||||||
),
|
),
|
||||||
("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
|
# `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)",
|
"batch_overlap/two_batch_overlap.py": "prefill (extend positions)",
|
||||||
"managers/scheduler.py": "prefill (truncation align knobs)",
|
"managers/scheduler.py": "prefill (truncation align knobs)",
|
||||||
"entrypoints/engine.py": "either half (flashinfer version floor)",
|
"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"})
|
_PASSED = frozenset({"model_path", "device", "random_seed"})
|
||||||
|
|
||||||
_EXPOSED = {
|
_EXPOSED = {
|
||||||
|
("entrypoints/sidecar.py", "grpc_port"),
|
||||||
("configs/embedding_model_spec.py", "chunked_prefill_size"),
|
("configs/embedding_model_spec.py", "chunked_prefill_size"),
|
||||||
("configs/embedding_model_spec.py", "cuda_graph_config"),
|
("configs/embedding_model_spec.py", "cuda_graph_config"),
|
||||||
("configs/embedding_model_spec.py", "disable_radix_cache"),
|
("configs/embedding_model_spec.py", "disable_radix_cache"),
|
||||||
@@ -143,24 +144,10 @@ _EXPOSED = {
|
|||||||
("configs/model_config.py", "quantization"),
|
("configs/model_config.py", "quantization"),
|
||||||
("configs/model_config.py", "speculative_algorithm"),
|
("configs/model_config.py", "speculative_algorithm"),
|
||||||
("configs/model_config.py", "speculative_draft_model_quantization"),
|
("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", "disaggregation_bootstrap_port"),
|
||||||
("disaggregation/common/conn.py", "pp_size"),
|
("disaggregation/common/conn.py", "pp_size"),
|
||||||
("disaggregation/decode_kvcache_offload_manager.py", "hicache_io_backend"),
|
("disaggregation/decode_kvcache_offload_manager.py", "hicache_io_backend"),
|
||||||
("disaggregation/decode_kvcache_offload_manager.py", "served_model_name"),
|
("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"),
|
("disaggregation/utils.py", "disaggregation_transfer_backend"),
|
||||||
("distributed/bootstrap.py", "disable_custom_all_reduce"),
|
("distributed/bootstrap.py", "disable_custom_all_reduce"),
|
||||||
("distributed/bootstrap.py", "enable_symm_mem"),
|
("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", "load_format"),
|
||||||
("elastic_ep/expert_backup_manager.py", "mooncake_ib_device"),
|
("elastic_ep/expert_backup_manager.py", "mooncake_ib_device"),
|
||||||
("entrypoints/engine.py", "attn_cp_size"),
|
("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", "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", "moe_dp_size"),
|
||||||
("entrypoints/engine.py", "pp_size"),
|
|
||||||
("entrypoints/engine.py", "quantization"),
|
|
||||||
("entrypoints/engine.py", "reasoning_parser"),
|
("entrypoints/engine.py", "reasoning_parser"),
|
||||||
(
|
(
|
||||||
"entrypoints/engine.py",
|
"entrypoints/engine.py",
|
||||||
"remote_instance_weight_loader_start_seed_via_transfer_engine",
|
"remote_instance_weight_loader_start_seed_via_transfer_engine",
|
||||||
),
|
),
|
||||||
("entrypoints/engine.py", "tool_call_parser"),
|
("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", "ep_dispatch_algorithm"),
|
||||||
("eplb/eplb_manager.py", "expert_distribution_recorder_buffer_size"),
|
("eplb/eplb_manager.py", "expert_distribution_recorder_buffer_size"),
|
||||||
("eplb/expert_distribution.py", "deepep_mode"),
|
("eplb/expert_distribution.py", "deepep_mode"),
|
||||||
@@ -236,7 +209,6 @@ _EXPOSED = {
|
|||||||
("kv_canary/capacities.py", "chunked_prefill_size"),
|
("kv_canary/capacities.py", "chunked_prefill_size"),
|
||||||
("kv_canary/capacities.py", "cuda_graph_config"),
|
("kv_canary/capacities.py", "cuda_graph_config"),
|
||||||
("kv_canary/capacities.py", "speculative_num_draft_tokens"),
|
("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", "attn_cp_size"),
|
||||||
("layers/cp/base.py", "cp_strategy"),
|
("layers/cp/base.py", "cp_strategy"),
|
||||||
("layers/cp/base.py", "enable_prefill_cp"),
|
("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", "moe_dp_size"),
|
||||||
("managers/data_parallel_controller.py", "pp_size"),
|
("managers/data_parallel_controller.py", "pp_size"),
|
||||||
("managers/data_parallel_controller.py", "soft_watchdog_timeout"),
|
("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_bootstrap_port"),
|
||||||
("managers/disagg_service.py", "disaggregation_mode"),
|
("managers/disagg_service.py", "disaggregation_mode"),
|
||||||
("managers/disagg_service.py", "disaggregation_transfer_backend"),
|
("managers/disagg_service.py", "disaggregation_transfer_backend"),
|
||||||
@@ -279,27 +248,7 @@ _EXPOSED = {
|
|||||||
("managers/scheduler.py", "pp_size"),
|
("managers/scheduler.py", "pp_size"),
|
||||||
("managers/scheduler.py", "soft_watchdog_timeout"),
|
("managers/scheduler.py", "soft_watchdog_timeout"),
|
||||||
("managers/scheduler.py", "speculative_algorithm"),
|
("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", "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", "disable_overlap_schedule"),
|
||||||
("managers/tp_worker.py", "model_path"),
|
("managers/tp_worker.py", "model_path"),
|
||||||
("managers/tp_worker.py", "random_seed"),
|
("managers/tp_worker.py", "random_seed"),
|
||||||
@@ -313,8 +262,6 @@ _EXPOSED = {
|
|||||||
("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "served_model_name"),
|
("mem_cache/hybrid_cache/hybrid_pool_assembler.py", "served_model_name"),
|
||||||
("mem_cache/kv_cache_builder.py", "hicache_mem_layout"),
|
("mem_cache/kv_cache_builder.py", "hicache_mem_layout"),
|
||||||
("mem_cache/radix_cache_cpp.py", "enable_hierarchical_cache"),
|
("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", "device"),
|
||||||
("model_executor/model_runner.py", "speculative_algorithm"),
|
("model_executor/model_runner.py", "speculative_algorithm"),
|
||||||
("model_executor/model_runner.py", "speculative_draft_attention_backend"),
|
("model_executor/model_runner.py", "speculative_draft_attention_backend"),
|
||||||
@@ -353,22 +300,11 @@ _EXPOSED = {
|
|||||||
"model_executor/runner_backend/tc_piecewise_cuda_graph_backend.py",
|
"model_executor/runner_backend/tc_piecewise_cuda_graph_backend.py",
|
||||||
"cuda_graph_config",
|
"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", "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", "disaggregation_mode"),
|
||||||
("observability/metrics_collector.py", "prefill_delayer_max_delay_passes"),
|
("observability/metrics_collector.py", "prefill_delayer_max_delay_passes"),
|
||||||
("observability/metrics_collector.py", "served_model_name"),
|
("observability/metrics_collector.py", "served_model_name"),
|
||||||
("parser/template_detection.py", "model_path"),
|
("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_algorithm"),
|
||||||
("speculative/adaptive_spec_params.py", "speculative_eagle_topk"),
|
("speculative/adaptive_spec_params.py", "speculative_eagle_topk"),
|
||||||
("speculative/dflash_worker_v2.py", "speculative_draft_window_size"),
|
("speculative/dflash_worker_v2.py", "speculative_draft_window_size"),
|
||||||
@@ -435,29 +371,20 @@ _EXPOSED_CUDA_ONLY: frozenset = frozenset()
|
|||||||
_OVERRIDDEN_AND_READ = {
|
_OVERRIDDEN_AND_READ = {
|
||||||
("configs/model_config.py", "dtype"),
|
("configs/model_config.py", "dtype"),
|
||||||
("configs/model_config.py", "model_path"),
|
("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"),
|
||||||
(
|
(
|
||||||
"disaggregation/decode_kvcache_offload_manager.py",
|
"disaggregation/decode_kvcache_offload_manager.py",
|
||||||
"hicache_storage_backend_extra_config",
|
"hicache_storage_backend_extra_config",
|
||||||
),
|
),
|
||||||
("disaggregation/encode_server.py", "load_format"),
|
|
||||||
("disaggregation/encode_server.py", "model_path"),
|
|
||||||
(
|
(
|
||||||
"distributed/device_communicators/mooncake_transfer_engine.py",
|
"distributed/device_communicators/mooncake_transfer_engine.py",
|
||||||
"hicache_storage_backend",
|
"hicache_storage_backend",
|
||||||
),
|
),
|
||||||
("dllm/config.py", "model_path"),
|
("dllm/config.py", "model_path"),
|
||||||
("elastic_ep/expert_backup_manager.py", "load_format"),
|
("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/api.py", "speculative_num_steps"),
|
||||||
("kv_canary/capacities.py", "speculative_num_draft_tokens"),
|
("kv_canary/capacities.py", "speculative_num_draft_tokens"),
|
||||||
("managers/scheduler.py", "hicache_storage_backend"),
|
("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"),
|
("managers/tp_worker.py", "model_path"),
|
||||||
("mem_cache/hiradix_cache.py", "hicache_storage_backend"),
|
("mem_cache/hiradix_cache.py", "hicache_storage_backend"),
|
||||||
("mem_cache/hiradix_cache.py", "hicache_storage_backend_extra_config"),
|
("mem_cache/hiradix_cache.py", "hicache_storage_backend_extra_config"),
|
||||||
|
|||||||
Reference in New Issue
Block a user