config: the draft runner carries its own attention backend
`build_draft_tp_worker` built a `ServerArgs` variant whose only job was to make four config reads answer with the draft's backend instead of the target's, and published it for the duration of the build so the bags agreed. The backend is a per-runner fact — target and draft coexist in one process — so it moves onto the runner, and the variant and the construction-time publish both go away. `ModelRunner` takes `draft_attention_backend` and resolves the runner's effective value once (`resolve_draft_attention_backend`: the algorithm's resolved backend, else `--speculative-draft-attention-backend`, else None for a target runner); `TpModelWorker` threads it to both runner constructions. `resolve_attention_backend_strs` reads it off the runner, and `ModelRunner` stamps the resolved pair *before* building backends so a backend can read it while it constructs — which is what the FlashInfer KV-access check needs now that it no longer asks the config. `configure_kv_cache_dtype` and the draft backend factory read the runner too. One latent bug falls out: the non-hybrid branch of the backend build ignored the resolved pair and re-read `server_args.attention_backend`, which is why the variant had to set that field as well as the split pair. It now uses the value that was resolved for the runner. `draft_server_args_overrides` and the `preserve_config()` publish switch are deleted; with them goes the last production `ServerArgs.derive` outside pre-publish config building, and the last construction-time publish. The chunked-prefix gate the target resolved simply stays in the bags, since nothing re-projects them.
This commit is contained in:
@@ -12,12 +12,17 @@ import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
from sglang.srt.managers.scheduler import Scheduler
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
from sglang.srt.model_executor.model_runner import (
|
||||
ModelRunner,
|
||||
resolve_draft_attention_backend,
|
||||
)
|
||||
from sglang.srt.model_executor.model_runner_components.attention_backend_setup import (
|
||||
resolve_attention_backend_strs,
|
||||
)
|
||||
from sglang.srt.model_executor.model_runner_components.load_model_utils import (
|
||||
build_load_config,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_context, get_model
|
||||
from sglang.srt.speculative.draft_worker_common import draft_server_args_overrides
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
@@ -102,27 +107,66 @@ class TestDraftPerRunnerConfig(CustomTestCase):
|
||||
build_load_config(load_format="dummy", **common).load_format, "dummy"
|
||||
)
|
||||
|
||||
# -- the variant left in the dflash / dspark path carries backends only ----
|
||||
# -- the attention backend is per-runner, not a config variant -------------
|
||||
|
||||
def test_the_draft_variant_carries_the_backend_family_only(self):
|
||||
self._seed(disable_chunked_prefix_cache=False)
|
||||
fields = draft_server_args_overrides("triton")
|
||||
def _runner(self, *, is_draft_worker, draft_attention_backend=None):
|
||||
runner = ModelRunner.__new__(ModelRunner)
|
||||
runner.server_args = get_context().server_args
|
||||
runner.is_draft_worker = is_draft_worker
|
||||
runner.draft_attention_backend = draft_attention_backend
|
||||
return runner
|
||||
|
||||
self.assertEqual(fields["attention_backend"], "triton")
|
||||
self.assertEqual(fields["speculative_draft_attention_backend"], "triton")
|
||||
self.assertIsNone(fields["prefill_attention_backend"])
|
||||
self.assertIsNone(fields["decode_attention_backend"])
|
||||
self.assertNotIn("context_length", fields)
|
||||
self.assertNotIn("load_format", fields)
|
||||
self.assertNotIn("skip_tokenizer_init", fields)
|
||||
def test_the_draft_backend_applies_to_the_draft_runner_only(self):
|
||||
self._seed(attention_backend="fa3")
|
||||
|
||||
def test_the_variant_carries_the_targets_resolved_gate(self):
|
||||
"""Publishing the variant re-projects the bags, so the gate travels."""
|
||||
self._seed(disable_chunked_prefix_cache=False)
|
||||
get_context().override("test.gate", disable_chunked_prefix_cache=True)
|
||||
self.assertTrue(
|
||||
draft_server_args_overrides("triton")["disable_chunked_prefix_cache"]
|
||||
draft = resolve_attention_backend_strs(
|
||||
model_runner=self._runner(
|
||||
is_draft_worker=True, draft_attention_backend="triton"
|
||||
)
|
||||
)
|
||||
self.assertEqual((draft.prefill, draft.decode), ("triton", "triton"))
|
||||
self.assertTrue(draft.is_draft_override)
|
||||
|
||||
target = resolve_attention_backend_strs(
|
||||
model_runner=self._runner(is_draft_worker=False)
|
||||
)
|
||||
self.assertEqual((target.prefill, target.decode), ("fa3", "fa3"))
|
||||
|
||||
def test_an_unresolved_draft_falls_back_to_the_config_field(self):
|
||||
"""The v2 workers pass no backend: --speculative-draft-attention-backend."""
|
||||
server_args = self._seed(
|
||||
attention_backend="fa3", speculative_draft_attention_backend="triton"
|
||||
)
|
||||
|
||||
def effective(*, is_draft_worker, passed=None):
|
||||
return resolve_draft_attention_backend(
|
||||
draft_attention_backend=passed,
|
||||
server_args=server_args,
|
||||
is_draft_worker=is_draft_worker,
|
||||
)
|
||||
|
||||
self.assertEqual(effective(is_draft_worker=True), "triton")
|
||||
self.assertEqual(effective(is_draft_worker=True, passed="fa3"), "fa3")
|
||||
self.assertIsNone(effective(is_draft_worker=False))
|
||||
|
||||
draft = resolve_attention_backend_strs(
|
||||
model_runner=self._runner(
|
||||
is_draft_worker=True,
|
||||
draft_attention_backend=effective(is_draft_worker=True),
|
||||
)
|
||||
)
|
||||
self.assertEqual((draft.prefill, draft.decode), ("triton", "triton"))
|
||||
|
||||
def test_the_target_keeps_its_split_pair(self):
|
||||
self._seed(
|
||||
attention_backend="fa3",
|
||||
prefill_attention_backend="flashinfer",
|
||||
decode_attention_backend="fa3",
|
||||
)
|
||||
target = resolve_attention_backend_strs(
|
||||
model_runner=self._runner(is_draft_worker=False)
|
||||
)
|
||||
self.assertEqual((target.prefill, target.decode), ("flashinfer", "fa3"))
|
||||
|
||||
# -- the scheduler hands over the process's own config ---------------------
|
||||
|
||||
|
||||
Reference in New Issue
Block a user