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:
Cheng Wan
2026-08-05 19:32:24 -07:00
committed by GitHub
parent 64eeb153df
commit 4ea227fa91
22 changed files with 199 additions and 103 deletions
@@ -398,23 +398,16 @@ def attn_backend_wrapper(runner: "ModelRunner", full_attn_backend: "AttentionBac
allowed = {"triton", "trtllm_mha", "flashinfer"}
else:
allowed = {"triton", "trtllm_mha", "fa4"}
attn_be = runner.server_args.attention_backend
prefill_be = runner.server_args.prefill_attention_backend
decode_be = runner.server_args.decode_attention_backend
# When using split prefill/decode backends, check each individually
if prefill_be and decode_be:
assert prefill_be in allowed and decode_be in allowed, (
f"Only {allowed} backends are supported on Blackwell GPUs for hybrid GDN models. "
f"Got prefill={prefill_be}, decode={decode_be}."
)
else:
assert attn_be in allowed, (
f"Only {allowed} backends are supported on Blackwell GPUs for hybrid GDN models. "
f"Got attention_backend={attn_be}."
)
prefill_be = runner.prefill_attention_backend_str
decode_be = runner.decode_attention_backend_str
assert prefill_be in allowed and decode_be in allowed, (
f"Only {allowed} backends are supported on Blackwell GPUs for hybrid GDN models. "
f"Got prefill={prefill_be}, decode={decode_be}."
)
elif is_npu():
assert (
runner.server_args.attention_backend == "ascend"
runner.prefill_attention_backend_str == "ascend"
and runner.decode_attention_backend_str == "ascend"
), "ascend backend is the only supported backend on NPU for hybrid GDN models, use --attention-backend ascend to specify the backend."
logger.info(f"Using hybrid linear attention backend for hybrid GDN models.")
linear_attn_backend = GDNAttnBackend(runner)
@@ -319,9 +319,8 @@ class FlashInferAttnBackend(AttentionBackend):
self.decode_kv_access = self.kv_cache_quant_method.resolve_attention_access(
"decode", "flashinfer"
)
prefill_backend, decode_backend = (
model_runner.server_args.get_attention_backends()
)
prefill_backend = model_runner.prefill_attention_backend_str
decode_backend = model_runner.decode_attention_backend_str
if self.__class__ is FlashInferAttnBackend:
if prefill_backend == "flashinfer":
self._check_kv_attention_access("prefill", self.prefill_kv_access)
+5
View File
@@ -310,6 +310,7 @@ class TpModelWorker(BaseTpWorker):
memory_pool_config: Optional[MemoryPoolConfig] = None,
is_multi_layer_eagle: bool = False,
context_length: Optional[int] = None,
draft_attention_backend: Optional[str] = None,
):
# Parse args
self.server_args = server_args
@@ -325,6 +326,8 @@ class TpModelWorker(BaseTpWorker):
# Draft worker: target's effective context length; the draft runs at
# absolute target positions. None keeps server_args.context_length.
self.context_length = context_length
# Draft worker: the attention backend the algorithm resolved for it.
self.draft_attention_backend = draft_attention_backend
# MTP model runners
self.model_runner_list: List[ModelRunner] = []
@@ -459,6 +462,7 @@ class TpModelWorker(BaseTpWorker):
req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
memory_pool_config=self.memory_pool_config,
draft_attention_backend=self.draft_attention_backend,
draft_model_idx=0 if self.is_multi_layer_eagle else None,
)
@@ -479,6 +483,7 @@ class TpModelWorker(BaseTpWorker):
req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
memory_pool_config=self.memory_pool_config,
draft_attention_backend=self.draft_attention_backend,
draft_model_idx=i,
)
)
@@ -924,7 +924,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
ret.extend_prefix_lens = extend_prefix_lens
ret.extend_num_tokens = batch.extend_num_tokens
positions, ret.extend_start_loc = compute_position(
model_runner.server_args.attention_backend,
model_runner.prefill_attention_backend_str,
ret.extend_prefix_lens,
ret.extend_seq_lens,
ret.extend_num_tokens,
@@ -112,6 +112,7 @@ from sglang.srt.model_executor.model_runner_components.attention_backend_setup i
build_attention_backends,
configure_aux_hidden_state_capture,
get_attention_backend,
resolve_attention_backend_strs,
)
from sglang.srt.model_executor.model_runner_components.cuda_graph_setup import (
capture_cuda_graphs,
@@ -261,6 +262,24 @@ class ModelRunnerOutput:
indexer_topk_output: Optional[TopkCaptureOutput] = None
def resolve_draft_attention_backend(
*,
draft_attention_backend: Optional[str],
server_args: ServerArgs,
is_draft_worker: bool,
) -> Optional[str]:
"""The attention backend a runner uses because it is a draft runner.
``None`` for a target runner. For a draft: the backend the algorithm that
built it resolved (the supported-backend fallback in
``build_draft_tp_worker``), else ``--speculative-draft-attention-backend``.
It belongs to the runner, not the process: target and draft coexist.
"""
if not is_draft_worker:
return None
return draft_attention_backend or server_args.speculative_draft_attention_backend
class ModelRunner:
"""ModelRunner runs the forward passes of the models."""
@@ -277,6 +296,7 @@ class ModelRunner:
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] = None,
memory_pool_config: Optional[MemoryPoolConfig] = None,
draft_model_idx: Optional[int] = None,
draft_attention_backend: Optional[str] = None,
):
# Parse args
self.mem_fraction_static = mem_fraction_static
@@ -293,6 +313,11 @@ class ModelRunner:
self.dist_port = nccl_port
self.server_args = server_args
self.is_draft_worker = is_draft_worker
self.draft_attention_backend = resolve_draft_attention_backend(
draft_attention_backend=draft_attention_backend,
server_args=server_args,
is_draft_worker=is_draft_worker,
)
# This runner's own load format, resolved before anything keys off it:
# the remote-instance transfer engine is initialized at the top of
# initialize(), long before the weights are loaded.
@@ -902,12 +927,15 @@ class ModelRunner:
dflash_target_layer_ids=self.spec_aux_config.dflash_target_layer_ids,
is_dspark=self.spec_algorithm.is_dspark(),
)
# Resolve before building: backends read the pair off the runner while
# they construct (the FlashInfer KV-access check).
resolved = resolve_attention_backend_strs(model_runner=self)
self.prefill_attention_backend_str = resolved.prefill
self.decode_attention_backend_str = resolved.decode
backends = build_attention_backends(model_runner=self)
self.attn_backend = backends.attn_backend
self.decode_attn_backend = backends.decode_attn_backend
self.decode_attn_backend_group = backends.decode_attn_backend_group
self.prefill_attention_backend_str = backends.prefill_attention_backend_str
self.decode_attention_backend_str = backends.decode_attention_backend_str
if self.server_args.dcp_size > 1 and get_parallel().dcp_replicate_q_proj:
self._prepare_replicated_q_proj()
@@ -1061,7 +1089,9 @@ class ModelRunner:
)
maybe_trigger_remote_instance_nccl_send_group(
server_args=self.server_args, tp_rank=self.ps.tp_rank
server_args=self.server_args,
tp_rank=self.ps.tp_rank,
load_format=draft_load_format,
)
with self._load_format_scope(draft_load_format):
@@ -1248,9 +1278,7 @@ class ModelRunner:
if spec_algorithm is not None
else False
),
speculative_draft_attention_backend=getattr(
self.server_args, "speculative_draft_attention_backend", None
),
speculative_draft_attention_backend=self.draft_attention_backend,
)
)
# This runner's OWN resolved dtype string (target or draft). Attention
@@ -17,7 +17,6 @@ from sglang.srt.utils import init_cublas
if TYPE_CHECKING:
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__)
@@ -73,8 +72,13 @@ def build_attention_backends(*, model_runner: ModelRunner) -> AttentionBackends:
if model_runner.device in ("cuda", "musa"):
init_cublas()
resolved = _resolve_attention_backend_strs(
server_args=server_args, is_draft_worker=model_runner.is_draft_worker
# Already resolved and stamped on the runner before this call.
resolved = ResolvedAttentionBackendStr(
prefill=model_runner.prefill_attention_backend_str,
decode=model_runner.decode_attention_backend_str,
is_draft_override=bool(
model_runner.is_draft_worker and model_runner.draft_attention_backend
),
)
if server_args.enable_pdmux:
@@ -140,10 +144,7 @@ def get_attention_backend(
*, model_runner: ModelRunner, init_new_workspace: bool = False
) -> AttentionBackend:
"""Init attention kernel backend."""
resolved = _resolve_attention_backend_strs(
server_args=model_runner.server_args,
is_draft_worker=model_runner.is_draft_worker,
)
resolved = resolve_attention_backend_strs(model_runner=model_runner)
return _build_resolved_backend(
model_runner=model_runner,
resolved=resolved,
@@ -151,10 +152,18 @@ def get_attention_backend(
)
def _resolve_attention_backend_strs(
*, server_args: ServerArgs, is_draft_worker: bool
def resolve_attention_backend_strs(
*, model_runner: ModelRunner
) -> ResolvedAttentionBackendStr:
draft_attn_backend = server_args.speculative_draft_attention_backend
"""The (prefill, decode) backends this runner runs.
A draft runner's backend is its own (``ModelRunner.draft_attention_backend``):
target and draft coexist in one process, so it cannot come from the
process-wide config.
"""
server_args = model_runner.server_args
is_draft_worker = model_runner.is_draft_worker
draft_attn_backend = model_runner.draft_attention_backend
if is_draft_worker and draft_attn_backend:
logger.warning(f"Overriding draft attention backend to {draft_attn_backend}.")
# Single backend for all draft modes (no prefill/decode split).
@@ -218,7 +227,7 @@ def _build_resolved_backend(
else:
attn_backend = _build_backend_from_str(
model_runner=model_runner,
backend_str=model_runner.server_args.attention_backend,
backend_str=resolved.prefill,
init_new_workspace=init_new_workspace,
)
return attn_backend
@@ -1425,7 +1425,7 @@ class DFlashWorkerV2(BaseSpecWorker):
"DFLASH prefill expected out_cache_loc, but got None."
)
positions, _ = compute_position(
self.model_runner.server_args.attention_backend,
self.model_runner.prefill_attention_backend_str,
draft_seq_lens,
ctx_lens,
int(sum(batch.extend_lens)),
+2 -1
View File
@@ -39,7 +39,8 @@ class DraftBackendFactory:
self.topk = topk
self.speculative_num_steps = speculative_num_steps
self.seed_dsa_topk_from_draft_extend = seed_dsa_topk_from_draft_extend
self.draft_attn_backend = server_args.speculative_draft_attention_backend
# The draft runner's own backend, not the process-wide config.
self.draft_attn_backend = draft_model_runner.draft_attention_backend
def _create_backend(
self, backend_name: str, backend_map: dict, error_template: str
@@ -9,7 +9,6 @@ import torch
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.server_args import ServerArgs
from sglang.srt.speculative.dflash_info import DFlashVerifyInput
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
@@ -63,26 +62,6 @@ def _resolve_draft_attention_backend_fallback(
return draft_backend
def draft_server_args_overrides(draft_backend) -> dict:
"""The fields a draft variant must carry: its attention backend.
Backend selection reads them off the config object the draft runner holds --
``speculative_draft_attention_backend`` in ``_resolve_attention_backend_strs``
and ``configure_kv_cache_dtype``, ``attention_backend`` in the non-hybrid
branch of the backend build, and the split pair must not shadow either with
the target's. ``disable_chunked_prefix_cache`` is the target's resolved gate,
which lives in the bags only: publishing the variant re-projects the bags
from it, so the value has to travel on the variant.
"""
return dict(
speculative_draft_attention_backend=draft_backend,
prefill_attention_backend=None,
decode_attention_backend=None,
attention_backend=draft_backend,
disable_chunked_prefix_cache=get_schedule().disable_chunked_prefix_cache,
)
def build_draft_tp_worker(
*,
server_args: ServerArgs,
@@ -101,23 +80,17 @@ def build_draft_tp_worker(
server_args=server_args, algo_label=algo_label
)
)
draft_server_args = server_args.derive(
"draft_worker.build", **draft_server_args_overrides(draft_backend)
draft_worker = TpModelWorker(
server_args=server_args,
gpu_id=gpu_id,
ps=ps,
nccl_port=nccl_port,
is_draft_worker=True,
# The draft runs at absolute target positions.
context_length=target_model_config.context_len,
draft_attention_backend=draft_backend,
)
# The draft's layers must resolve config from the draft's own bags.
with get_context().preserve_config():
get_context().set_server_args(draft_server_args)
draft_worker = TpModelWorker(
server_args=draft_server_args,
gpu_id=gpu_id,
ps=ps,
nccl_port=nccl_port,
is_draft_worker=True,
# The draft runs at absolute target positions.
context_length=target_model_config.context_len,
)
draft_model_runner = draft_worker.model_runner
draft_worker.draft_runner = draft_model_runner
return DraftWorkerBundle(
@@ -444,7 +444,7 @@ class DSparkWorkerV2(BaseSpecWorker):
batch.prefix_lens, dtype=torch.int32, device=device
)
positions, _ = compute_position(
self.model_runner.server_args.attention_backend,
self.model_runner.prefill_attention_backend_str,
draft_seq_lens,
ctx_lens,
int(sum(batch.extend_lens)),
@@ -322,6 +322,11 @@ class MockModelRunner(ModelRunner):
self.dtype = dtype
self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto"
# This runner's own resolved backends (production stamps these in
# ModelRunner.initialize); a draft runner would carry its own.
self.prefill_attention_backend_str = case.backend
self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None
self.gpu_id = 0
self.canary_manager = None
self.page_size = case.page_size
@@ -293,6 +293,11 @@ class DSAMockModelRunner(ModelRunner):
# `set_mla_kv_buffer` does the quantize on the way in.
self.kv_cache_dtype = torch.float8_e4m3fn if fp8_kv_cache else dtype
self.kv_cache_dtype_str = "auto"
# This runner's own resolved backends (production stamps these in
# ModelRunner.initialize); a draft runner would carry its own.
self.prefill_attention_backend_str = case.backend
self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None
# For TARGET_VERIFY / DRAFT_EXTEND, the DSA backend uses
# `self.speculative_num_draft_tokens` to size `seqlens_expanded`
# (`dsa_backend.py:482-486,510-515`). When zero, deep_gemm's
@@ -332,6 +332,11 @@ class MockDSV4ModelRunner:
self.dtype = dtype
self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto"
# This runner's own resolved backends (production stamps these in
# ModelRunner.initialize); a draft runner would carry its own.
self.prefill_attention_backend_str = case.backend
self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None
self.gpu_id = 0
self.canary_manager = None
self.page_size = case.page_size
@@ -322,6 +322,11 @@ class DualChunkMockModelRunner(ModelRunner):
self.dtype = dtype
self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto"
# This runner's own resolved backends (production stamps these in
# ModelRunner.initialize); a draft runner would carry its own.
self.prefill_attention_backend_str = case.backend
self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None
self.gpu_id = 0
self.canary_manager = None
self.page_size = case.page_size
@@ -217,6 +217,11 @@ class MockGDNModelRunner(ModelRunner):
self.dtype = dtype
self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto"
# This runner's own resolved backends (production stamps these in
# ModelRunner.initialize); a draft runner would carry its own.
self.prefill_attention_backend_str = case.backend
self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None
@@ -222,6 +222,11 @@ class MockKDAModelRunner(ModelRunner):
self.dtype = dtype
self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto"
# This runner's own resolved backends (production stamps these in
# ModelRunner.initialize); a draft runner would carry its own.
self.prefill_attention_backend_str = case.backend
self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None
@@ -230,6 +230,11 @@ class MockLightningModelRunner(ModelRunner):
self.dtype = dtype
self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto"
# This runner's own resolved backends (production stamps these in
# ModelRunner.initialize); a draft runner would carry its own.
self.prefill_attention_backend_str = case.backend
self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None
@@ -315,6 +315,11 @@ class MockMamba2ModelRunner(ModelRunner):
self.dtype = dtype
self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto"
# This runner's own resolved backends (production stamps these in
# ModelRunner.initialize); a draft runner would carry its own.
self.prefill_attention_backend_str = case.backend
self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None
self.gpu_id = 0
self.ps = ParallelState.trivial()
self.canary_manager = None
@@ -233,6 +233,11 @@ class MockMLAModelRunner(ModelRunner):
# does the BF16->FP8 cast on the way in.
self.kv_cache_dtype = torch.float8_e4m3fn if fp8_kv_cache else dtype
self.kv_cache_dtype_str = "fp8_e4m3" if fp8_kv_cache else "auto"
# This runner's own resolved backends (production stamps these in
# ModelRunner.initialize); a draft runner would carry its own.
self.prefill_attention_backend_str = case.backend
self.decode_attention_backend_str = case.backend
self.draft_attention_backend = None
self.gpu_id = 0
self.canary_manager = None
self.page_size = case.page_size
@@ -75,6 +75,7 @@ class TestKVCacheQuantRegistry(CustomTestCase):
runner = object.__new__(ModelRunner)
runner.server_args = SimpleNamespace(kv_cache_dtype="fp4_e2m1")
runner.draft_attention_backend = None
with self.assertRaisesRegex(ValueError, "fp4_mx_block16"):
runner.configure_kv_cache_dtype()
@@ -51,17 +51,15 @@ class TestChunkedPrefixCacheGate(CustomTestCase):
get_context().set_server_args(sa) # what a later republish would do
self.assertFalse(get_schedule().disable_chunked_prefix_cache)
def test_draft_variant_fields_carry_the_gate(self):
# Publishing the draft variant re-projects the bags from it, so the
# gate — which lives in the bags only — has to travel on the variant.
from sglang.srt.speculative.draft_worker_common import (
draft_server_args_overrides,
)
def test_the_gate_survives_a_draft_build(self):
# The draft build no longer publishes a config of its own, so the gate
# the target resolved stays in the bags for the rest of the process.
self._seed(attention_backend="triton")
maybe_disable_chunked_prefix_cache(use_mla_backend=True, is_draft_worker=False)
fields = draft_server_args_overrides(draft_backend="fa3")
self.assertTrue(fields["disable_chunked_prefix_cache"])
self.assertTrue(get_schedule().disable_chunked_prefix_cache)
maybe_disable_chunked_prefix_cache(use_mla_backend=False, is_draft_worker=True)
self.assertTrue(get_schedule().disable_chunked_prefix_cache)
if __name__ == "__main__":
@@ -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 ---------------------