config: retire the multi-engine accommodation in the runtime context (#35022)

This commit is contained in:
Cheng Wan
2026-08-17 16:15:33 -07:00
committed by GitHub
parent bc312d185d
commit 2b278b4ac4
8 changed files with 172 additions and 283 deletions
+25 -25
View File
@@ -10,7 +10,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** (in-process tokenizer build, multi-Engine) 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) | 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()` |
@@ -79,11 +79,6 @@ re-projects its own bags, so a parent-side override is lost. Values that feed
construction before any bag exists (group init reads `server_args.tp_size`) have no construction before any bag exists (group init reads `server_args.tp_size`) have no
bag to override at all. bag to override at all.
- **Nested publishes**: `get_context().preserve_config()` snapshots the enclosing
lifecycle (including its post-publish overrides) and reinstates it on exit. No
production caller is left — the draft build was the last one, and per-runner
values are constructor arguments now — so it survives for tests and for a future
construction step that genuinely has to publish a private copy.
### Reads that legitimately stay on a `ServerArgs` instance ### Reads that legitimately stay on a `ServerArgs` instance
@@ -110,28 +105,31 @@ 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 read `entrypoints/`, and the tokenizer-process multimodal processors still read
`self.server_args`: several `Engine`s can share one process, and the process-global `self.server_args` today. The old justification ("several `Engine`s can share
bags are last-publish-wins across engines. `base_gpu_id` also differs per engine, one process, bags are last-publish-wins across them") is **retracted** — owner
so no process-global value can stand in for it — ruling (2026-08-15): a process holds at most one live config at a time
`BaseMultimodalProcessor._fast_image_processor_device` is the shape to copy. (The (concurrent multi-Engine is unsupported; sequential rebuild stays legal, unit
encode-server DP workers used to specialize a config copy for the same reason; tests rely on it). These reads are scheduled to become bag reads in the
their device now travels as `MMEncoder(gpu_id=...)`.) bag-read series; treat them as pinned debt, not as a boundary to imitate. What
genuinely stays per-instance is what differs per *worker* within one engine:
`base_gpu_id` travels as a constructor argument (`MMEncoder(gpu_id=...)`;
`BaseMultimodalProcessor._fast_image_processor_device` is the shape to copy).
- **Whole-object passes** (`f(server_args)` handing the instance along) keep the - **Whole-object passes** (`f(server_args)` handing the instance along) keep the
supplied-instance contract; don't rewrite the parameter reads unless the supplied-instance contract; don't rewrite the parameter reads unless the
field is runtime-mutated (see the elastic-EP `ep_size` case in field is runtime-mutated (see the elastic-EP `ep_size` case in
`eplb/expert_location.py`) — **or the field is one that resolution fills in `eplb/expert_location.py`) — **or the field is one that resolution fills in
and the callee runs in a process that has published.** That second case is and the callee runs in a process that has published.** That second case is
step-12 debt, not a style question: the record is destined to carry the pinned debt, not a style question: the record is destined to carry the
user's raw input, so `server_args.page_size` inside a runner-owned user's raw input, so `server_args.page_size` inside a runner-owned
constructor will read the raw pre-resolution value instead of the effective constructor will read the raw pre-resolution value instead of the effective
one. Debt means a decision, not automatically a bag read: pick where the one. Debt means a decision, not automatically a bag read: pick where the
value should come from — usually the `get_*()` bag, sometimes a runner stamp value should come from — usually the `get_*()` bag, sometimes a runner stamp
or a constructor argument (the per-mode attention pair and the encode-server or a constructor argument (the per-mode attention pair and the encode-server
`gpu_id` above are dispositions of exactly this debt). And the per-instance `gpu_id` above are dispositions of exactly this debt). The per-instance
boundaries above stay exempt from this unless-clause: a multi-Engine site boundaries above are **not** exempt from this unless-clause (the multi-Engine
must not become a process-global bag read even for a resolution-filled exemption is retracted); each one gets its own disposition.
field. `test_supplied_instance_exposure_ratchet.py` `test_supplied_instance_exposure_ratchet.py`
pins the remaining set — three spellings of the read: `server_args.field`, pins the remaining set — three spellings of the read: `server_args.field`,
literal-name `getattr(server_args, "field", default)`, and the parked form literal-name `getattr(server_args, "field", default)`, and the parked form
(`self.x = server_args` in a method that takes the parameter, read as (`self.x = server_args` in a method that takes the parameter, read as
@@ -224,13 +222,15 @@ this).
`self.server_args.field` is still right for handed per-instance config (see `self.server_args.field` is still right for handed per-instance config (see
"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 the copies to read any more). The allow-list is `GrammarManager` and `MMEncoder`;
tokenizer-manager family, `entrypoints/`, the tokenizer-process multimodal the tokenizer-manager family, `entrypoints/`, and the tokenizer-process
processors, `GrammarManager`, `MMEncoder` — but 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 the multi-Engine case - the tokenizer-manager family and `entrypoints/` are **pinned debt awaiting
proper: several of them can live in one process, so a bag read would answer conversion to bag reads** (the old multi-Engine justification is retracted —
from whichever Engine published last; one process, one live config); the reads still work today because the
instance carries resolved values;
- `GrammarManager` is a handed instance — it is constructed with the config its - `GrammarManager` is a handed instance — it is constructed with the config its
owner hands it and never assumes a published namespace; owner hands it and never assumes a published namespace;
- `MMEncoder` publishes the very instance it is handed (`publish(server_args, - `MMEncoder` publishes the very instance it is handed (`publish(server_args,
@@ -474,7 +474,7 @@ Never module-skip a test "until the migration settles" — seed the context inst
## Where to read the code ## Where to read the code
Key source files: `python/sglang/srt/runtime_context.py` (the container, every tier, Key source files: `python/sglang/srt/runtime_context.py` (the container, every tier,
`publish`, `_ConfigBag`, `preserve_config`, `override_server_args`), `publish`, `_ConfigBag`, `override_server_args`),
`python/sglang/srt/arg_groups/overrides.py` (override registry, passes, `python/sglang/srt/arg_groups/overrides.py` (override registry, passes,
`declare_late_resolution`), `python/sglang/srt/server_args.py` (`NS` metadata, `declare_late_resolution`), `python/sglang/srt/server_args.py` (`NS` metadata,
`Arg(..., resolvable=True)`, `__setattr__` strict guard), and the guardrail tests under `Arg(..., resolvable=True)`, `__setattr__` strict guard), and the guardrail tests under
+4 -5
View File
@@ -278,8 +278,8 @@ class Engine(EngineScoreMixin, EngineBase):
self.template_manager = template_manager self.template_manager = template_manager
self._scheduler_init_result = scheduler_init_result self._scheduler_init_result = scheduler_init_result
# Engine-spawned weight cache daemons owned by *this* instance (empty # Engine-spawned weight cache daemons owned by *this* instance (empty
# unless --weight-cache-mode daemon). Kept per-instance so two Engines # unless --weight-cache-mode daemon), so shutdown() reaps exactly what
# in one process each reap only their own daemons in shutdown(). # this Engine spawned.
self._weight_cache_daemon_procs = weight_cache_daemon_procs self._weight_cache_daemon_procs = weight_cache_daemon_procs
if tokenizer_manager is not None: if tokenizer_manager is not None:
tokenizer_manager._subprocess_watchdog = subprocess_watchdog tokenizer_manager._subprocess_watchdog = subprocess_watchdog
@@ -1109,9 +1109,8 @@ class Engine(EngineScoreMixin, EngineBase):
): ):
resolve_auto_parsers(server_args) resolve_auto_parsers(server_args)
# Launch daemons (daemon mode only). Handles are threaded back to the # Launch daemons (daemon mode only). The handles travel back to the
# owning Engine instance (not a class attr) so two Engines in one process # Engine that spawned them; shutdown() reaps from there.
# don't clobber each other's daemon list.
weight_cache_daemon_procs: List = [] weight_cache_daemon_procs: List = []
if server_args.weight_cache_mode == "daemon": if server_args.weight_cache_mode == "daemon":
weight_cache_daemon_procs = cls._launch_weight_cache_daemons(server_args) weight_cache_daemon_procs = cls._launch_weight_cache_daemons(server_args)
@@ -2057,11 +2057,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
def record_config_updates(self, source: str, **fields) -> None: def record_config_updates(self, source: str, **fields) -> None:
"""Record a control-plane config change for this engine. """Record a control-plane config change for this engine.
Per-engine state: several ``Engine``s can share one tokenizer process. These are post-startup facts the config bags do not model (weight
The readback endpoints overlay these onto the startup config. The version, model path, the tokenizer's HiCache mirror); the readback
process-global sibling is ``RuntimeContext.override`` / endpoints overlay them onto the startup config. The process-global
``resolved_server_args_dict``, which writes the config bags every sibling is ``RuntimeContext.override`` / ``resolved_server_args_dict``,
process shares. which writes the config bags.
""" """
unknown = sorted(f for f in fields if f not in _SERVER_ARGS_FIELDS) unknown = sorted(f for f in fields if f not in _SERVER_ARGS_FIELDS)
if unknown: if unknown:
+11 -71
View File
@@ -722,34 +722,6 @@ def _build_config_bags(server_args: Any) -> dict:
return tops return tops
def _snapshot_bag_values(bags: dict | None) -> dict | None:
"""Per-leaf value snapshot of a config-bag tree (bags are mutated in
place by ``override``, so reference snapshots alias live state)."""
if bags is None:
return None
snap: dict = {}
def walk(prefix: str, bag) -> None:
snap[prefix] = dict(object.__getattribute__(bag, "_fields"))
for name, sub in object.__getattribute__(bag, "_subs").items():
walk(f"{prefix}.{name}", sub)
for name, bag in bags.items():
walk(name, bag)
return snap
def _restore_bag_values(bags: dict, snap: dict) -> None:
def walk(prefix: str, bag) -> None:
for key, value in snap[prefix].items():
bag._set(key, value)
for name, sub in object.__getattribute__(bag, "_subs").items():
walk(f"{prefix}.{name}", sub)
for name, bag in bags.items():
walk(name, bag)
class RuntimeContext: class RuntimeContext:
"""Container for the structured runtime accessors; exposes ``parallel``, """Container for the structured runtime accessors; exposes ``parallel``,
``server_args``, the resolved config namespace bags, ``flags``, ``server_args``, the resolved config namespace bags, ``flags``,
@@ -953,12 +925,12 @@ class RuntimeContext:
``ServerArgs`` field names, so overlaying them onto the top level of ``ServerArgs`` field names, so overlaying them onto the top level of
either base is exact. either base is exact.
This covers the process-global bags only. Per-engine control-plane This covers the process-global bags only. Control-plane facts the bags
changes (weight version, model path, the tokenizer's HiCache mirror) do not model (weight version, model path, the tokenizer's HiCache
live on the tokenizer manager — several ``Engine``s can share one mirror) live on the tokenizer manager, and
process — and ``TokenizerManager.resolved_config_dict`` overlays those ``TokenizerManager.resolved_config_dict`` overlays those for the
for the top-level ``/server_info`` body. The two are separate logs, not top-level ``/server_info`` body. The two are separate logs, not one
one merged dict. merged dict.
""" """
d = dict(vars(self.server_args)) if base is None else dict(base) d = dict(vars(self.server_args)) if base is None else dict(base)
for _source, fields in self._overrides_log: for _source, fields in self._overrides_log:
@@ -987,33 +959,6 @@ class RuntimeContext:
""" """
return _ServerArgsOverride(self, fields) return _ServerArgsOverride(self, fields)
@contextmanager
def preserve_config(self):
"""Snapshot the full config lifecycle and reinstate it verbatim on exit.
For nested construction steps that publish a private ``ServerArgs``
copy (e.g. a draft-worker build) and must leave the enclosing
lifecycle — including its post-publish overrides — untouched.
"""
prev_server_args = self._server_args
prev_bags = self._config_bags
prev_bag_values = _snapshot_bag_values(prev_bags)
prev_overrides_log = list(self._overrides_log)
prev_publish_role = self._publish_role
prev_parallel_config = self.parallel._config
prev_capture = self.flags.capture.enable_torch_compile
try:
yield
finally:
self._server_args = prev_server_args
self._config_bags = prev_bags
if prev_bags is not None:
_restore_bag_values(prev_bags, prev_bag_values)
self._overrides_log = prev_overrides_log
self._publish_role = prev_publish_role
self.parallel._config = prev_parallel_config
self.flags.capture.enable_torch_compile = prev_capture
class _ServerArgsOverride: class _ServerArgsOverride:
"""Scoped config override (see ``RuntimeContext.override_server_args``). """Scoped config override (see ``RuntimeContext.override_server_args``).
@@ -1198,7 +1143,6 @@ def get_observability() -> _ConfigBag:
ROLE_NAMESPACE_SETS: dict[str, frozenset[str] | None] = { ROLE_NAMESPACE_SETS: dict[str, frozenset[str] | None] = {
# Reads (almost) everything by design — the model-executing process. # Reads (almost) everything by design — the model-executing process.
"scheduler": None, "scheduler": None,
"launcher": None,
"test": None, "test": None,
# Audited (record-mode smokes, plain + DP-attention): the DP controller # Audited (record-mode smokes, plain + DP-attention): the DP controller
# reads only the elastic-EP gate; its module's static read set agrees. # reads only the elastic-EP gate; its module's static read set agrees.
@@ -1317,21 +1261,17 @@ def publish(server_args, *, role: str, hf_config: Any = None) -> RuntimeContext:
Records the process ``role`` (``tokenizer`` / ``scheduler`` / Records the process ``role`` (``tokenizer`` / ``scheduler`` /
``dp_controller`` / ``encoder`` / ``expert_backup`` / ``dp_controller`` / ``encoder`` / ``expert_backup`` /
``weight_cache_daemon`` / ``launcher`` / ``test``) and ``weight_cache_daemon`` / ``test``) and
projects the config bags. Draft workers skip publish (they must not clobber projects the config bags. Draft workers skip publish (they must not clobber
the target). ``role`` is provenance, and — when ``SGLANG_ROLE_NAMESPACES`` the target). ``role`` is provenance, and — when ``SGLANG_ROLE_NAMESPACES``
is ``enforce`` — the key into ``ROLE_NAMESPACE_SETS`` for fail-closed is ``enforce`` — the key into ``ROLE_NAMESPACE_SETS`` for fail-closed
namespace-read enforcement (``record`` audits the reads instead). namespace-read enforcement (``record`` audits the reads instead).
``hf_config`` is accepted for forward-compat and currently unused. ``hf_config`` is accepted for forward-compat and currently unused.
Normally one call per process, but re-publish is allowed and is A process holds at most one live config: the bags always describe the
**last-publish-wins** (bags re-projected, provenance reset, role engine running now. Re-publish is allowed and is **last-publish-wins**
overwritten). Two sanctioned multi-publish shapes exist: the in-process (bags re-projected, provenance reset, role overwritten), which is what
Engine builds its ``TokenizerManager`` inside the launcher process (the lets one process rebuild an engine after shutting the previous one down.
process ends up with the tokenizer publish), and multiple Engines in one
process publish in sequence — which is exactly why per-instance managers
must read ``self.server_args`` for anything engine-specific rather than
the process-global bags.
""" """
if _ROLE_NS_MODE == "enforce" and role not in ROLE_NAMESPACE_SETS: if _ROLE_NS_MODE == "enforce" and role not in ROLE_NAMESPACE_SETS:
# Fail closed at publish time, not at the first stray read. # Fail closed at publish time, not at the first stray read.
@@ -37,6 +37,20 @@ TEST_CLASSIFICATION_BASE_MODEL = os.environ.get(
_CLS_NUM_LABELS = AutoConfig.from_pretrained(TEST_CLASSIFICATION_BASE_MODEL).num_labels _CLS_NUM_LABELS = AutoConfig.from_pretrained(TEST_CLASSIFICATION_BASE_MODEL).num_labels
def _collect_scores(engine_kwargs, calls):
"""Boot one engine, run ``calls`` through score(), shut it down.
A process holds one live config, so the reference engine must be gone
before the engine under test boots.
"""
engine = Engine(**engine_kwargs)
try:
return [engine.score(**call).scores for call in calls]
finally:
engine.shutdown()
torch.cuda.empty_cache()
class TestMISServerArgsValidation(unittest.TestCase): class TestMISServerArgsValidation(unittest.TestCase):
"""Test ServerArgs defaults for MIS mode.""" """Test ServerArgs defaults for MIS mode."""
@@ -50,8 +64,24 @@ class TestMISServerArgsValidation(unittest.TestCase):
class TestMultiItemScoringOptimization(CustomTestCase): class TestMultiItemScoringOptimization(CustomTestCase):
"""Test the Multi-Item Scoring (MIS) optimization with generation models.""" """Test the Multi-Item Scoring (MIS) optimization with generation models."""
CONSISTENCY_CALL = dict(
query="Is this a fact?\n",
items=[" The sun rises in the east"],
label_token_ids=[9454, 2753],
apply_softmax=True,
)
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
(cls.non_mis_consistency_scores,) = _collect_scores(
dict(
model_path=TEST_MODEL_NAME,
disable_radix_cache=True,
chunked_prefill_size=-1,
mem_fraction_static=0.15,
),
[cls.CONSISTENCY_CALL],
)
cls.engine = Engine( cls.engine = Engine(
model_path=TEST_MODEL_NAME, model_path=TEST_MODEL_NAME,
disable_radix_cache=True, disable_radix_cache=True,
@@ -60,19 +90,11 @@ class TestMultiItemScoringOptimization(CustomTestCase):
attention_backend="flashinfer", attention_backend="flashinfer",
mem_fraction_static=0.15, mem_fraction_static=0.15,
) )
cls.non_mis_engine = Engine(
model_path=TEST_MODEL_NAME,
disable_radix_cache=True,
chunked_prefill_size=-1,
mem_fraction_static=0.15,
)
@classmethod @classmethod
def tearDownClass(cls): def tearDownClass(cls):
if cls.engine is not None: if cls.engine is not None:
cls.engine.shutdown() cls.engine.shutdown()
if cls.non_mis_engine is not None:
cls.non_mis_engine.shutdown()
torch.cuda.empty_cache() torch.cuda.empty_cache()
def test_mis_basic(self): def test_mis_basic(self):
@@ -98,23 +120,8 @@ class TestMultiItemScoringOptimization(CustomTestCase):
def test_mis_consistency_with_single_item(self): def test_mis_consistency_with_single_item(self):
"""MIS with one item should match non-MIS scoring closely.""" """MIS with one item should match non-MIS scoring closely."""
query = "Is this a fact?\n" mis_scores = self.engine.score(**self.CONSISTENCY_CALL).scores
items = [" The sun rises in the east"] non_mis_scores = self.non_mis_consistency_scores
label_token_ids = [9454, 2753]
mis_scores = self.engine.score(
query=query,
items=items,
label_token_ids=label_token_ids,
apply_softmax=True,
).scores
non_mis_scores = self.non_mis_engine.score(
query=query,
items=items,
label_token_ids=label_token_ids,
apply_softmax=True,
).scores
self.assertEqual(len(mis_scores), 1) self.assertEqual(len(mis_scores), 1)
self.assertEqual(len(non_mis_scores), 1) self.assertEqual(len(non_mis_scores), 1)
@@ -151,14 +158,29 @@ class TestMultiItemScoringClassification(CustomTestCase):
Pre-trained Qwen3ForSequenceClassification, so the head weights are Pre-trained Qwen3ForSequenceClassification, so the head weights are
deterministic. One class rather than four because the CI harness demands an deterministic. One class rather than four because the CI harness demands an
idle GPU at every setUpClass -- splitting these means re-booting the same idle GPU at every setUpClass -- splitting these means re-booting the same
two engines instead of sharing them. score() is stateless and the radix engines instead of sharing them. score() is stateless and the radix cache
cache is off, so sharing is safe. is off, so sharing is safe.
""" """
NUM_LABELS = _CLS_NUM_LABELS NUM_LABELS = _CLS_NUM_LABELS
FALLBACK_CALL = dict(query="Test:", items=["A", "B"], apply_softmax=True)
SINGLE_VS_MIS_CALL = dict(
query="Rate this option:",
items=[" Option A", " Option B", " Option C"],
apply_softmax=True,
)
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
cls.non_mis_fallback_scores, cls.non_mis_single_scores = _collect_scores(
dict(
model_path=TEST_CLASSIFICATION_BASE_MODEL,
disable_radix_cache=True,
mem_fraction_static=0.15,
),
[cls.FALLBACK_CALL, cls.SINGLE_VS_MIS_CALL],
)
cls.engine = Engine( cls.engine = Engine(
model_path=TEST_CLASSIFICATION_BASE_MODEL, model_path=TEST_CLASSIFICATION_BASE_MODEL,
disable_radix_cache=True, disable_radix_cache=True,
@@ -167,17 +189,11 @@ class TestMultiItemScoringClassification(CustomTestCase):
attention_backend="flashinfer", attention_backend="flashinfer",
mem_fraction_static=0.15, mem_fraction_static=0.15,
) )
cls.non_mis_engine = Engine(
model_path=TEST_CLASSIFICATION_BASE_MODEL,
disable_radix_cache=True,
mem_fraction_static=0.15,
)
@classmethod @classmethod
def tearDownClass(cls): def tearDownClass(cls):
for engine in (cls.engine, cls.non_mis_engine): if cls.engine is not None:
if engine is not None: cls.engine.shutdown()
engine.shutdown()
torch.cuda.empty_cache() torch.cuda.empty_cache()
def test_classification_mis_basic(self): def test_classification_mis_basic(self):
@@ -215,9 +231,7 @@ class TestMultiItemScoringClassification(CustomTestCase):
def test_classification_non_mis_fallback(self): def test_classification_non_mis_fallback(self):
"""Classification model works correctly without --enable-mis.""" """Classification model works correctly without --enable-mis."""
scores = self.non_mis_engine.score( scores = self.non_mis_fallback_scores
query="Test:", items=["A", "B"], apply_softmax=True
).scores
self.assertEqual(len(scores), 2) self.assertEqual(len(scores), 2)
for score_list in scores: for score_list in scores:
@@ -395,14 +409,8 @@ class TestMultiItemScoringClassification(CustomTestCase):
perturbs hidden states; after softmax the scores should still land perturbs hidden states; after softmax the scores should still land
within places=1 (+-0.05). within places=1 (+-0.05).
""" """
query = "Rate this option:" non_mis_scores = self.non_mis_single_scores
items = [" Option A", " Option B", " Option C"] mis_scores = self.engine.score(**self.SINGLE_VS_MIS_CALL).scores
non_mis_scores = self.non_mis_engine.score(
query=query, items=items, apply_softmax=True
).scores
mis_scores = self.engine.score(
query=query, items=items, apply_softmax=True
).scores
self.assertEqual(len(mis_scores), len(non_mis_scores)) self.assertEqual(len(mis_scores), len(non_mis_scores))
for i, (ms, ns) in enumerate(zip(mis_scores, non_mis_scores)): for i, (ms, ns) in enumerate(zip(mis_scores, non_mis_scores)):
@@ -420,48 +428,73 @@ class TestMultiItemScoringParity(CustomTestCase):
"""Test that MIS produces the same results as single-item scoring.""" """Test that MIS produces the same results as single-item scoring."""
@classmethod @classmethod
def setUpClass(cls): def _cases(cls):
cls.engine_single = Engine( """The scoring calls both engines run, keyed by the test that reads them."""
model_path=TEST_MODEL_NAME, tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME)
disable_radix_cache=True,
log_level="error", def label_ids(labels):
mem_fraction_static=0.15, return [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
)
cls.engine_mis = Engine( return {
model_path=TEST_MODEL_NAME, "basic": dict(
disable_radix_cache=True, query="Rate this option:",
chunked_prefill_size=-1, items=[" Option A", " Option B", " Option C"],
log_level="error", label_token_ids=label_ids([" good", " bad"]),
enable_mis=True, apply_softmax=True,
attention_backend="flashinfer", ),
mem_fraction_static=0.15, "tokenized": dict(
) query=tokenizer.encode("Rate this option:", add_special_tokens=False),
items=[
tokenizer.encode(item, add_special_tokens=False)
for item in [" Option X", " Option Y"]
],
label_token_ids=label_ids([" good", " bad"]),
apply_softmax=True,
),
"no_softmax": dict(
query="The weather today is",
items=[" sunny", " cloudy", " rainy"],
label_token_ids=label_ids([" nice", " bad"]),
apply_softmax=False,
),
"many_items": dict(
query="Rate this option from 1 to 5:",
items=[f" Option {i}" for i in range(10)],
label_token_ids=label_ids([" 1", " 2", " 3", " 4", " 5"]),
apply_softmax=True,
),
}
@classmethod @classmethod
def tearDownClass(cls): def setUpClass(cls):
if cls.engine_single is not None: cases = cls._cases()
cls.engine_single.shutdown() names, calls = list(cases), list(cases.values())
if cls.engine_mis is not None: base = dict(
cls.engine_mis.shutdown() model_path=TEST_MODEL_NAME,
torch.cuda.empty_cache() disable_radix_cache=True,
log_level="error",
mem_fraction_static=0.15,
)
cls.single_scores = dict(zip(names, _collect_scores(base, calls)))
cls.mis_scores = dict(
zip(
names,
_collect_scores(
dict(
base,
chunked_prefill_size=-1,
enable_mis=True,
attention_backend="flashinfer",
),
calls,
),
)
)
def _compare_scores( def _compare_scores(self, test_name):
self, query, items, label_token_ids=None, apply_softmax=True, test_name=""
):
"""Compare MIS vs single-item scoring results.""" """Compare MIS vs single-item scoring results."""
single_scores = self.engine_single.score( single_scores = self.single_scores[test_name]
query=query, mis_scores = self.mis_scores[test_name]
items=items,
label_token_ids=label_token_ids,
apply_softmax=apply_softmax,
).scores
mis_scores = self.engine_mis.score(
query=query,
items=items,
label_token_ids=label_token_ids,
apply_softmax=apply_softmax,
).scores
self.assertEqual( self.assertEqual(
len(mis_scores), len(single_scores), f"{test_name}: count mismatch" len(mis_scores), len(single_scores), f"{test_name}: count mismatch"
@@ -477,40 +510,16 @@ class TestMultiItemScoringParity(CustomTestCase):
) )
def test_parity_basic(self): def test_parity_basic(self):
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME) self._compare_scores("basic")
query = "Rate this option:"
items = [" Option A", " Option B", " Option C"]
labels = [" good", " bad"]
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
self._compare_scores(query, items, label_ids, test_name="basic")
def test_parity_tokenized_inputs(self): def test_parity_tokenized_inputs(self):
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME) self._compare_scores("tokenized")
query = "Rate this option:"
items = [" Option X", " Option Y"]
labels = [" good", " bad"]
query_ids = tokenizer.encode(query, add_special_tokens=False)
items_ids = [tokenizer.encode(i, add_special_tokens=False) for i in items]
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
self._compare_scores(query_ids, items_ids, label_ids, test_name="tokenized")
def test_parity_without_softmax(self): def test_parity_without_softmax(self):
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME) self._compare_scores("no_softmax")
query = "The weather today is"
items = [" sunny", " cloudy", " rainy"]
labels = [" nice", " bad"]
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
self._compare_scores(
query, items, label_ids, apply_softmax=False, test_name="no_softmax"
)
def test_parity_many_items(self): def test_parity_many_items(self):
tokenizer = AutoTokenizer.from_pretrained(TEST_MODEL_NAME) self._compare_scores("many_items")
query = "Rate this option from 1 to 5:"
items = [f" Option {i}" for i in range(10)]
labels = [" 1", " 2", " 3", " 4", " 5"]
label_ids = [tokenizer.encode(lb, add_special_tokens=False)[0] for lb in labels]
self._compare_scores(query, items, label_ids, test_name="many_items")
if __name__ == "__main__": if __name__ == "__main__":
@@ -2,8 +2,8 @@
Regression: runtime updates (weight version, model path, HiCache attach) were Regression: runtime updates (weight version, model path, HiCache attach) were
written onto the manager's ServerArgs instance so that the readback endpoints written onto the manager's ServerArgs instance so that the readback endpoints
would show them. They are per-engine — several Engines can share a tokenizer would show them. The record stays pristine; the updates live in a separate log
process — so they live on the manager and the endpoints overlay them. that the endpoints overlay on top of it.
""" """
import re import re
@@ -39,11 +39,6 @@ class TestTokenizerConfigUpdates(CustomTestCase):
manager.record_config_updates("test", weight_version="v2") manager.record_config_updates("test", weight_version="v2")
self.assertEqual(manager.server_args.weight_version, "v1") self.assertEqual(manager.server_args.weight_version, "v1")
def test_two_engines_keep_their_own_updates(self):
first, second = _manager(weight_version="v1"), _manager(weight_version="v1")
first.record_config_updates("test", weight_version="v2")
self.assertEqual(second.config_value("weight_version"), "v1")
def test_the_readback_dict_carries_the_updates(self): def test_the_readback_dict_carries_the_updates(self):
manager = _manager(hicache_storage_backend=None) manager = _manager(hicache_storage_backend=None)
manager.record_config_updates( manager.record_config_updates(
@@ -1,9 +1,9 @@
"""The fast-image-processor device comes from the processor's own ServerArgs. """The fast-image-processor device comes from the processor's own ServerArgs.
Regression: the device decision read the published global ServerArgs, which is Regression: the device decision read the published global ServerArgs, so every
last-publish-wins. Two engines in one tokenizer process then shared whichever processor answered with one process-wide device. The encode-server DP workers
config published last, so one engine's images were preprocessed on the other each drive their own GPU, which no process-global value can express — the
engine's GPU. device has to come from what the worker was handed.
""" """
import unittest import unittest
@@ -109,60 +109,6 @@ class TestContextOverride(CustomTestCase):
with self.assertRaises(AttributeError): with self.assertRaises(AttributeError):
sa.page_size = 999 sa.page_size = 999
def test_preserve_config_keeps_post_publish_overrides(self):
# A nested build (e.g. a draft worker) publishes its own private copy;
# on exit the target's resolved bags — including post-publish
# overrides — must be reinstated verbatim, not re-projected from the
# pristine record (which would silently drop the overrides).
target = self._publish()
rc.get_context().override(
"ModelRunner.configure_kv_cache_dtype", kv_cache_dtype="fp8_e4m3"
)
draft = ServerArgs(model_path="dummy", kv_cache_dtype="bf16")
with rc.get_context().preserve_config():
rc.get_context().set_server_args(draft)
# Inside the scope the draft's bags are live...
self.assertEqual(rc.get_model().kv_cache_dtype, "bf16")
# ...and its own post-publish overrides work as usual.
rc.get_context().override("draft-load", kv_cache_dtype="fp8_e5m2")
self.assertEqual(rc.get_model().kv_cache_dtype, "fp8_e5m2")
# Target slot, bags, and provenance restored verbatim.
self.assertIs(rc.get_context().server_args, target)
self.assertEqual(rc.get_model().kv_cache_dtype, "fp8_e4m3")
self.assertEqual(
rc.get_context().overrides_log(),
[
(
"ModelRunner.configure_kv_cache_dtype",
{"kv_cache_dtype": "fp8_e4m3"},
)
],
)
def test_preserve_config_restores_in_scope_override_without_republish(self):
# An override inside the scope (no republish) mutates the live bags
# and provenance log in place; the scope must restore entry VALUES,
# not just reassign the aliased objects.
self._publish()
rc.get_context().override("srcA", page_size=16)
with rc.get_context().preserve_config():
rc.get_context().override("in-scope", page_size=64)
self.assertEqual(rc.get_schedule().page_size, 64)
self.assertEqual(rc.get_schedule().page_size, 16)
self.assertEqual(
rc.get_context().overrides_log(), [("srcA", {"page_size": 16})]
)
def test_preserve_config_restores_on_exception(self):
target = self._publish()
rc.get_context().override("srcA", page_size=16)
with self.assertRaises(RuntimeError):
with rc.get_context().preserve_config():
rc.get_context().set_server_args(ServerArgs(model_path="dummy"))
raise RuntimeError("nested build failed")
self.assertIs(rc.get_context().server_args, target)
self.assertEqual(rc.get_schedule().page_size, 16)
def test_publish_records_role(self): def test_publish_records_role(self):
rc.publish(ServerArgs(model_path="dummy"), role="scheduler") rc.publish(ServerArgs(model_path="dummy"), role="scheduler")
self.assertEqual(rc.publish_role(), "scheduler") self.assertEqual(rc.publish_role(), "scheduler")