Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 7 additions & 4 deletions .claude/skills/sglang-runtime-context/SKILL.md
Original file line number Diff line number Diff line change
Expand Up @@ -237,10 +237,13 @@ Never module-skip a test "until the migration settles" — seed the context inst
## Hard-won pitfalls (check these before/while refactoring)

- **Moving code drops first-line guards**: early returns (`if self.is_draft_worker: return`)
are the easiest thing to lose when relocating a method body. Only drafts built through
`build_draft_tp_worker()` get private bags (a preserved publish of the rewritten copy);
drafts constructed directly with `is_draft_worker=True` skip publish and **share the
target's bags** — a draft-side write there poisons the target.
are the easiest thing to lose when relocating a method body. Every draft is built
under a preserved publish of its own config: the scheduler makes the copy with
`draft_server_args_copy()` (seeded from the resolved config, so load-time overrides
carry) and publishes it around the worker factory, and `build_draft_tp_worker()`
nests the same shape for dflash / dspark. The publish ends when construction does —
anything the draft reads later (`alloc_memory_pool`, `init_attention_backends`,
cuda-graph capture) is back on the target's bags.
- **Registry-completeness timing**: a gate that consults an extensible list is only correct
after the registrars ran (platform `init_backend()` at module import). See "load-time vs
resolution-time".
Expand Down
30 changes: 13 additions & 17 deletions python/sglang/srt/managers/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -868,31 +868,27 @@ def maybe_init_draft_worker(self):
self.external_corpus_manager = None
return

from sglang.srt.speculative.draft_worker_common import (
draft_server_args_copy,
)

# Launch a draft worker for speculative decoding
draft_worker_kwargs = dict(
draft_server_args = draft_server_args_copy(
server_args=self.server_args,
target_model_config=self.tp_worker.model_runner.model_config,
)
draft_worker_kwargs = dict(
server_args=draft_server_args,
gpu_id=self.ps.gpu_id,
ps=self.ps,
nccl_port=self.nccl_port,
target_worker=self.tp_worker,
)

if get_spec().speculative_draft_load_format is not None:
# Write the draft load_format onto server_args (not just the bag):
# the draft worker is built from a copy of self.server_args and
# build_load_config reads server_args.load_format, so a bag-only
# override would be ignored and the draft would load in the target's
# format.
self.server_args.override(
"scheduler.draft_load_format",
load_format=get_spec().speculative_draft_load_format,
)
logger.info(
f"Using draft model load_format: '{get_spec().speculative_draft_load_format}'"
)

DraftWorkerClass = self.spec_algorithm.create_worker(self.server_args)
self.draft_worker = DraftWorkerClass(**draft_worker_kwargs)
DraftWorkerClass = self.spec_algorithm.create_worker(draft_server_args)
with get_context().preserve_config():
get_context().set_server_args(draft_server_args)
self.draft_worker = DraftWorkerClass(**draft_worker_kwargs)

if self.spec_algorithm.is_ngram():
from sglang.srt.speculative.external_corpus_manager import (
Expand Down
41 changes: 40 additions & 1 deletion python/sglang/srt/speculative/draft_worker_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
from sglang.srt.runtime_context import get_context, get_schedule
from sglang.srt.runtime_context import get_context, get_schedule, get_spec
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.dflash_info import DFlashVerifyInput
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
Expand Down Expand Up @@ -61,6 +61,13 @@ def _resolve_draft_attention_backend_fallback(
return draft_backend


def _draft_load_format_fields() -> dict:
draft_load_format = get_spec().speculative_draft_load_format
if draft_load_format is None:
return {}
return dict(load_format=draft_load_format)


def draft_server_args_overrides(target_model_config, draft_backend) -> dict:
"""Pre-publish field adjustments for a draft ``ServerArgs`` copy.

Expand All @@ -78,7 +85,39 @@ def draft_server_args_overrides(target_model_config, draft_backend) -> dict:
attention_backend=draft_backend,
context_length=target_model_config.context_len,
disable_chunked_prefix_cache=get_schedule().disable_chunked_prefix_cache,
**_draft_load_format_fields(),
)


def draft_server_args_copy(server_args: ServerArgs, target_model_config) -> ServerArgs:
"""A draft-only ``ServerArgs`` for the workers that build their own draft.

Starts from the config the process resolved, not from the pristine seed:
the copy is published while the draft builds, and load-time overrides made
before this point (the chunked-prefix gate, the SM100 GDN prefill default)
are part of what the draft's layers must see. On top of that,
``context_length`` follows the target (the draft reads target KV) and
``load_format`` follows ``--speculative-draft-load-format``. The target's
own instance is untouched.
"""
draft_load_format = get_spec().speculative_draft_load_format
if draft_load_format is not None:
logger.info(f"Using draft model load_format: '{draft_load_format}'")

resolved = {}
for _source, fields in get_context().overrides_log():
resolved.update(fields)

draft_server_args = deepcopy(server_args)
draft_server_args.override(
"draft_worker.copy",
**{
**resolved,
"context_length": target_model_config.context_len,
**_draft_load_format_fields(),
},
)
return draft_server_args


def build_draft_tp_worker(
Expand Down
12 changes: 0 additions & 12 deletions python/sglang/srt/speculative/eagle_worker_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -264,12 +264,6 @@ def init_token_map(self):
self.hot_token_id = None
elif get_spec().speculative_token_map is not None:
self.hot_token_id = load_token_map(get_spec().speculative_token_map)
self.server_args.override(
"eagle_worker.hot_token_map",
json_model_override_args=(
f'{{"hot_vocab_size": {len(self.hot_token_id)}}}'
),
)
else:
self.hot_token_id = None

Expand Down Expand Up @@ -1010,12 +1004,6 @@ def __init__(
server_args.speculative_algorithm
)

# Override the context length of the draft model to be the same as the target model.
server_args.override(
"spec_worker.match_target_context_length",
context_length=target_worker.model_runner.model_config.context_len,
)

self._draft_worker = EagleDraftWorker(
server_args,
gpu_id,
Expand Down
6 changes: 0 additions & 6 deletions python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -679,12 +679,6 @@ def __init__(
self.req_to_token_pool, self.token_to_kv_pool_allocator = (
target_worker.get_memory_pool()
)
# Match the draft context length to the target (assistant reads target KV).
server_args.override(
"spec_worker.match_target_context_length",
context_length=target_worker.model_runner.model_config.context_len,
)

self._draft_worker = FrozenKVMTPDraftWorker(
server_args,
gpu_id,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -907,12 +907,6 @@ def __init__(
server_args.speculative_algorithm
)

# Override the context length of the draft model to be the same as the target model.
server_args.override(
"spec_worker.match_target_context_length",
context_length=target_worker.model_runner.model_config.context_len,
)

self._draft_worker = MultiLayerEagleDraftWorker(
server_args,
gpu_id,
Expand Down
6 changes: 0 additions & 6 deletions python/sglang/srt/speculative/standalone_worker_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -150,12 +150,6 @@ def __init__(
server_args.speculative_algorithm
)

# Override the context length of the draft model to be the same as the target model.
server_args.override(
"spec_worker.match_target_context_length",
context_length=target_worker.model_runner.model_config.context_len,
)

# Create our custom draft worker that doesn't share embeddings/lm_head
self._draft_worker = StandaloneDraftWorker(
server_args,
Expand Down
84 changes: 84 additions & 0 deletions test/registered/unit/spec/test_draft_server_args_copy.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,84 @@
"""The draft's ServerArgs is a copy; the target's stays as the launcher left it.

Regression: the v2 spec workers wrote the draft's context_length (and the
scheduler the draft's load_format) onto the ServerArgs instance they share with
the target worker, so every later reader of that instance saw draft values.
"""

import unittest
from types import SimpleNamespace

from sglang.srt.runtime_context import get_context
from sglang.srt.speculative.draft_worker_common import (
draft_server_args_copy,
draft_server_args_overrides,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase

register_cpu_ci(est_time=5, suite="base-a-test-cpu")

TARGET_MODEL_CONFIG = SimpleNamespace(context_len=4096)


class TestDraftServerArgsCopy(CustomTestCase):
def _seed(self, **fields):
override = get_context().override_server_args(**fields)
server_args = override.install()
self.addCleanup(override.restore)
return server_args

def test_the_draft_context_length_follows_the_target(self):
target = self._seed(context_length=None)
draft = draft_server_args_copy(target, TARGET_MODEL_CONFIG)
self.assertEqual(draft.context_length, 4096)

def test_the_target_instance_is_left_alone(self):
target = self._seed(context_length=None, load_format="auto")
draft = draft_server_args_copy(target, TARGET_MODEL_CONFIG)
self.assertIsNot(draft, target)
self.assertIsNone(target.context_length)
self.assertEqual(target.load_format, "auto")

def test_the_draft_load_format_applies_only_when_configured(self):
target = self._seed(load_format="auto", speculative_draft_load_format="dummy")
self.assertEqual(
draft_server_args_copy(target, TARGET_MODEL_CONFIG).load_format, "dummy"
)
self.assertEqual(target.load_format, "auto")

target = self._seed(load_format="auto")
self.assertEqual(
draft_server_args_copy(target, TARGET_MODEL_CONFIG).load_format, "auto"
)

def test_load_time_overrides_reach_the_draft(self):
target = self._seed(disable_chunked_prefix_cache=False)
# What the target runner resolved before the draft is built — e.g. the
# chunked-prefix gate for an attention backend that cannot serve it.
get_context().override("test.gate", disable_chunked_prefix_cache=True)

draft = draft_server_args_copy(target, TARGET_MODEL_CONFIG)
self.assertTrue(draft.disable_chunked_prefix_cache)
self.assertFalse(target.disable_chunked_prefix_cache)

def test_the_draft_specific_fields_win_over_the_resolved_ones(self):
target = self._seed(context_length=None, load_format="auto")
get_context().override("test.late", context_length=128, load_format="npcache")

draft = draft_server_args_copy(target, TARGET_MODEL_CONFIG)
self.assertEqual(draft.context_length, 4096)

def test_the_built_draft_overrides_carry_the_load_format_too(self):
self._seed(speculative_draft_load_format="dummy")
fields = draft_server_args_overrides(TARGET_MODEL_CONFIG, "triton")
self.assertEqual(fields["load_format"], "dummy")

self._seed()
self.assertNotIn(
"load_format", draft_server_args_overrides(TARGET_MODEL_CONFIG, "triton")
)


if __name__ == "__main__":
unittest.main()
Loading
Loading