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"} allowed = {"triton", "trtllm_mha", "flashinfer"}
else: else:
allowed = {"triton", "trtllm_mha", "fa4"} allowed = {"triton", "trtllm_mha", "fa4"}
attn_be = runner.server_args.attention_backend prefill_be = runner.prefill_attention_backend_str
prefill_be = runner.server_args.prefill_attention_backend decode_be = runner.decode_attention_backend_str
decode_be = runner.server_args.decode_attention_backend assert prefill_be in allowed and decode_be in allowed, (
# When using split prefill/decode backends, check each individually f"Only {allowed} backends are supported on Blackwell GPUs for hybrid GDN models. "
if prefill_be and decode_be: f"Got prefill={prefill_be}, decode={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}."
)
elif is_npu(): elif is_npu():
assert ( 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." ), "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.") logger.info(f"Using hybrid linear attention backend for hybrid GDN models.")
linear_attn_backend = GDNAttnBackend(runner) linear_attn_backend = GDNAttnBackend(runner)
@@ -319,9 +319,8 @@ class FlashInferAttnBackend(AttentionBackend):
self.decode_kv_access = self.kv_cache_quant_method.resolve_attention_access( self.decode_kv_access = self.kv_cache_quant_method.resolve_attention_access(
"decode", "flashinfer" "decode", "flashinfer"
) )
prefill_backend, decode_backend = ( prefill_backend = model_runner.prefill_attention_backend_str
model_runner.server_args.get_attention_backends() decode_backend = model_runner.decode_attention_backend_str
)
if self.__class__ is FlashInferAttnBackend: if self.__class__ is FlashInferAttnBackend:
if prefill_backend == "flashinfer": if prefill_backend == "flashinfer":
self._check_kv_attention_access("prefill", self.prefill_kv_access) 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, memory_pool_config: Optional[MemoryPoolConfig] = None,
is_multi_layer_eagle: bool = False, is_multi_layer_eagle: bool = False,
context_length: Optional[int] = None, context_length: Optional[int] = None,
draft_attention_backend: Optional[str] = None,
): ):
# Parse args # Parse args
self.server_args = server_args self.server_args = server_args
@@ -325,6 +326,8 @@ class TpModelWorker(BaseTpWorker):
# Draft worker: target's effective context length; the draft runs at # Draft worker: target's effective context length; the draft runs at
# absolute target positions. None keeps server_args.context_length. # absolute target positions. None keeps server_args.context_length.
self.context_length = 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 # MTP model runners
self.model_runner_list: List[ModelRunner] = [] self.model_runner_list: List[ModelRunner] = []
@@ -459,6 +462,7 @@ class TpModelWorker(BaseTpWorker):
req_to_token_pool=self.req_to_token_pool, req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
memory_pool_config=self.memory_pool_config, 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, 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, req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
memory_pool_config=self.memory_pool_config, memory_pool_config=self.memory_pool_config,
draft_attention_backend=self.draft_attention_backend,
draft_model_idx=i, draft_model_idx=i,
) )
) )
@@ -924,7 +924,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
ret.extend_prefix_lens = extend_prefix_lens ret.extend_prefix_lens = extend_prefix_lens
ret.extend_num_tokens = batch.extend_num_tokens ret.extend_num_tokens = batch.extend_num_tokens
positions, ret.extend_start_loc = compute_position( 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_prefix_lens,
ret.extend_seq_lens, ret.extend_seq_lens,
ret.extend_num_tokens, ret.extend_num_tokens,
@@ -112,6 +112,7 @@ from sglang.srt.model_executor.model_runner_components.attention_backend_setup i
build_attention_backends, build_attention_backends,
configure_aux_hidden_state_capture, configure_aux_hidden_state_capture,
get_attention_backend, get_attention_backend,
resolve_attention_backend_strs,
) )
from sglang.srt.model_executor.model_runner_components.cuda_graph_setup import ( from sglang.srt.model_executor.model_runner_components.cuda_graph_setup import (
capture_cuda_graphs, capture_cuda_graphs,
@@ -261,6 +262,24 @@ class ModelRunnerOutput:
indexer_topk_output: Optional[TopkCaptureOutput] = None 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: class ModelRunner:
"""ModelRunner runs the forward passes of the models.""" """ModelRunner runs the forward passes of the models."""
@@ -277,6 +296,7 @@ class ModelRunner:
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] = None, token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] = None,
memory_pool_config: Optional[MemoryPoolConfig] = None, memory_pool_config: Optional[MemoryPoolConfig] = None,
draft_model_idx: Optional[int] = None, draft_model_idx: Optional[int] = None,
draft_attention_backend: Optional[str] = None,
): ):
# Parse args # Parse args
self.mem_fraction_static = mem_fraction_static self.mem_fraction_static = mem_fraction_static
@@ -293,6 +313,11 @@ class ModelRunner:
self.dist_port = nccl_port self.dist_port = nccl_port
self.server_args = server_args self.server_args = server_args
self.is_draft_worker = is_draft_worker 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: # This runner's own load format, resolved before anything keys off it:
# the remote-instance transfer engine is initialized at the top of # the remote-instance transfer engine is initialized at the top of
# initialize(), long before the weights are loaded. # 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, dflash_target_layer_ids=self.spec_aux_config.dflash_target_layer_ids,
is_dspark=self.spec_algorithm.is_dspark(), 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) backends = build_attention_backends(model_runner=self)
self.attn_backend = backends.attn_backend self.attn_backend = backends.attn_backend
self.decode_attn_backend = backends.decode_attn_backend self.decode_attn_backend = backends.decode_attn_backend
self.decode_attn_backend_group = backends.decode_attn_backend_group 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: if self.server_args.dcp_size > 1 and get_parallel().dcp_replicate_q_proj:
self._prepare_replicated_q_proj() self._prepare_replicated_q_proj()
@@ -1061,7 +1089,9 @@ class ModelRunner:
) )
maybe_trigger_remote_instance_nccl_send_group( 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): with self._load_format_scope(draft_load_format):
@@ -1248,9 +1278,7 @@ class ModelRunner:
if spec_algorithm is not None if spec_algorithm is not None
else False else False
), ),
speculative_draft_attention_backend=getattr( speculative_draft_attention_backend=self.draft_attention_backend,
self.server_args, "speculative_draft_attention_backend", None
),
) )
) )
# This runner's OWN resolved dtype string (target or draft). Attention # 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: if TYPE_CHECKING:
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -73,8 +72,13 @@ def build_attention_backends(*, model_runner: ModelRunner) -> AttentionBackends:
if model_runner.device in ("cuda", "musa"): if model_runner.device in ("cuda", "musa"):
init_cublas() init_cublas()
resolved = _resolve_attention_backend_strs( # Already resolved and stamped on the runner before this call.
server_args=server_args, is_draft_worker=model_runner.is_draft_worker 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: if server_args.enable_pdmux:
@@ -140,10 +144,7 @@ def get_attention_backend(
*, model_runner: ModelRunner, init_new_workspace: bool = False *, model_runner: ModelRunner, init_new_workspace: bool = False
) -> AttentionBackend: ) -> AttentionBackend:
"""Init attention kernel backend.""" """Init attention kernel backend."""
resolved = _resolve_attention_backend_strs( resolved = resolve_attention_backend_strs(model_runner=model_runner)
server_args=model_runner.server_args,
is_draft_worker=model_runner.is_draft_worker,
)
return _build_resolved_backend( return _build_resolved_backend(
model_runner=model_runner, model_runner=model_runner,
resolved=resolved, resolved=resolved,
@@ -151,10 +152,18 @@ def get_attention_backend(
) )
def _resolve_attention_backend_strs( def resolve_attention_backend_strs(
*, server_args: ServerArgs, is_draft_worker: bool *, model_runner: ModelRunner
) -> ResolvedAttentionBackendStr: ) -> 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: if is_draft_worker and draft_attn_backend:
logger.warning(f"Overriding draft attention backend to {draft_attn_backend}.") logger.warning(f"Overriding draft attention backend to {draft_attn_backend}.")
# Single backend for all draft modes (no prefill/decode split). # Single backend for all draft modes (no prefill/decode split).
@@ -218,7 +227,7 @@ def _build_resolved_backend(
else: else:
attn_backend = _build_backend_from_str( attn_backend = _build_backend_from_str(
model_runner=model_runner, model_runner=model_runner,
backend_str=model_runner.server_args.attention_backend, backend_str=resolved.prefill,
init_new_workspace=init_new_workspace, init_new_workspace=init_new_workspace,
) )
return attn_backend return attn_backend
@@ -1425,7 +1425,7 @@ class DFlashWorkerV2(BaseSpecWorker):
"DFLASH prefill expected out_cache_loc, but got None." "DFLASH prefill expected out_cache_loc, but got None."
) )
positions, _ = compute_position( positions, _ = compute_position(
self.model_runner.server_args.attention_backend, self.model_runner.prefill_attention_backend_str,
draft_seq_lens, draft_seq_lens,
ctx_lens, ctx_lens,
int(sum(batch.extend_lens)), int(sum(batch.extend_lens)),
+2 -1
View File
@@ -39,7 +39,8 @@ class DraftBackendFactory:
self.topk = topk self.topk = topk
self.speculative_num_steps = speculative_num_steps self.speculative_num_steps = speculative_num_steps
self.seed_dsa_topk_from_draft_extend = seed_dsa_topk_from_draft_extend 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( def _create_backend(
self, backend_name: str, backend_map: dict, error_template: str 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.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.managers.tp_worker import TpModelWorker
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, get_schedule
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.dflash_info import DFlashVerifyInput from sglang.srt.speculative.dflash_info import DFlashVerifyInput
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2 from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
@@ -63,26 +62,6 @@ def _resolve_draft_attention_backend_fallback(
return draft_backend 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( def build_draft_tp_worker(
*, *,
server_args: ServerArgs, server_args: ServerArgs,
@@ -101,23 +80,17 @@ def build_draft_tp_worker(
server_args=server_args, algo_label=algo_label server_args=server_args, algo_label=algo_label
) )
) )
draft_server_args = server_args.derive( draft_worker = TpModelWorker(
"draft_worker.build", **draft_server_args_overrides(draft_backend) 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_model_runner = draft_worker.model_runner
draft_worker.draft_runner = draft_model_runner draft_worker.draft_runner = draft_model_runner
return DraftWorkerBundle( return DraftWorkerBundle(
@@ -444,7 +444,7 @@ class DSparkWorkerV2(BaseSpecWorker):
batch.prefix_lens, dtype=torch.int32, device=device batch.prefix_lens, dtype=torch.int32, device=device
) )
positions, _ = compute_position( positions, _ = compute_position(
self.model_runner.server_args.attention_backend, self.model_runner.prefill_attention_backend_str,
draft_seq_lens, draft_seq_lens,
ctx_lens, ctx_lens,
int(sum(batch.extend_lens)), int(sum(batch.extend_lens)),
@@ -322,6 +322,11 @@ class MockModelRunner(ModelRunner):
self.dtype = dtype self.dtype = dtype
self.kv_cache_dtype = dtype self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto" 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.gpu_id = 0
self.canary_manager = None self.canary_manager = None
self.page_size = case.page_size self.page_size = case.page_size
@@ -293,6 +293,11 @@ class DSAMockModelRunner(ModelRunner):
# `set_mla_kv_buffer` does the quantize on the way in. # `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 = torch.float8_e4m3fn if fp8_kv_cache else dtype
self.kv_cache_dtype_str = "auto" 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 # For TARGET_VERIFY / DRAFT_EXTEND, the DSA backend uses
# `self.speculative_num_draft_tokens` to size `seqlens_expanded` # `self.speculative_num_draft_tokens` to size `seqlens_expanded`
# (`dsa_backend.py:482-486,510-515`). When zero, deep_gemm's # (`dsa_backend.py:482-486,510-515`). When zero, deep_gemm's
@@ -332,6 +332,11 @@ class MockDSV4ModelRunner:
self.dtype = dtype self.dtype = dtype
self.kv_cache_dtype = dtype self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto" 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.gpu_id = 0
self.canary_manager = None self.canary_manager = None
self.page_size = case.page_size self.page_size = case.page_size
@@ -322,6 +322,11 @@ class DualChunkMockModelRunner(ModelRunner):
self.dtype = dtype self.dtype = dtype
self.kv_cache_dtype = dtype self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto" 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.gpu_id = 0
self.canary_manager = None self.canary_manager = None
self.page_size = case.page_size self.page_size = case.page_size
@@ -217,6 +217,11 @@ class MockGDNModelRunner(ModelRunner):
self.dtype = dtype self.dtype = dtype
self.kv_cache_dtype = dtype self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto" 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.gpu_id = 0
self.ps = ParallelState.trivial() self.ps = ParallelState.trivial()
self.canary_manager = None self.canary_manager = None
@@ -222,6 +222,11 @@ class MockKDAModelRunner(ModelRunner):
self.dtype = dtype self.dtype = dtype
self.kv_cache_dtype = dtype self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto" 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.gpu_id = 0
self.ps = ParallelState.trivial() self.ps = ParallelState.trivial()
self.canary_manager = None self.canary_manager = None
@@ -230,6 +230,11 @@ class MockLightningModelRunner(ModelRunner):
self.dtype = dtype self.dtype = dtype
self.kv_cache_dtype = dtype self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto" 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.gpu_id = 0
self.ps = ParallelState.trivial() self.ps = ParallelState.trivial()
self.canary_manager = None self.canary_manager = None
@@ -315,6 +315,11 @@ class MockMamba2ModelRunner(ModelRunner):
self.dtype = dtype self.dtype = dtype
self.kv_cache_dtype = dtype self.kv_cache_dtype = dtype
self.kv_cache_dtype_str = "auto" 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.gpu_id = 0
self.ps = ParallelState.trivial() self.ps = ParallelState.trivial()
self.canary_manager = None self.canary_manager = None
@@ -233,6 +233,11 @@ class MockMLAModelRunner(ModelRunner):
# does the BF16->FP8 cast on the way in. # 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 = torch.float8_e4m3fn if fp8_kv_cache else dtype
self.kv_cache_dtype_str = "fp8_e4m3" if fp8_kv_cache else "auto" 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.gpu_id = 0
self.canary_manager = None self.canary_manager = None
self.page_size = case.page_size self.page_size = case.page_size
@@ -75,6 +75,7 @@ class TestKVCacheQuantRegistry(CustomTestCase):
runner = object.__new__(ModelRunner) runner = object.__new__(ModelRunner)
runner.server_args = SimpleNamespace(kv_cache_dtype="fp4_e2m1") runner.server_args = SimpleNamespace(kv_cache_dtype="fp4_e2m1")
runner.draft_attention_backend = None
with self.assertRaisesRegex(ValueError, "fp4_mx_block16"): with self.assertRaisesRegex(ValueError, "fp4_mx_block16"):
runner.configure_kv_cache_dtype() runner.configure_kv_cache_dtype()
@@ -51,17 +51,15 @@ class TestChunkedPrefixCacheGate(CustomTestCase):
get_context().set_server_args(sa) # what a later republish would do get_context().set_server_args(sa) # what a later republish would do
self.assertFalse(get_schedule().disable_chunked_prefix_cache) self.assertFalse(get_schedule().disable_chunked_prefix_cache)
def test_draft_variant_fields_carry_the_gate(self): def test_the_gate_survives_a_draft_build(self):
# Publishing the draft variant re-projects the bags from it, so the # The draft build no longer publishes a config of its own, so the gate
# gate — which lives in the bags only — has to travel on the variant. # the target resolved stays in the bags for the rest of the process.
from sglang.srt.speculative.draft_worker_common import (
draft_server_args_overrides,
)
self._seed(attention_backend="triton") self._seed(attention_backend="triton")
maybe_disable_chunked_prefix_cache(use_mla_backend=True, is_draft_worker=False) maybe_disable_chunked_prefix_cache(use_mla_backend=True, is_draft_worker=False)
fields = draft_server_args_overrides(draft_backend="fa3") self.assertTrue(get_schedule().disable_chunked_prefix_cache)
self.assertTrue(fields["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__": if __name__ == "__main__":
@@ -12,12 +12,17 @@ import unittest
from types import SimpleNamespace from types import SimpleNamespace
from sglang.srt.managers.scheduler import Scheduler 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 ( from sglang.srt.model_executor.model_runner_components.load_model_utils import (
build_load_config, build_load_config,
) )
from sglang.srt.runtime_context import get_context, get_model 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.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -102,27 +107,66 @@ class TestDraftPerRunnerConfig(CustomTestCase):
build_load_config(load_format="dummy", **common).load_format, "dummy" 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): def _runner(self, *, is_draft_worker, draft_attention_backend=None):
self._seed(disable_chunked_prefix_cache=False) runner = ModelRunner.__new__(ModelRunner)
fields = draft_server_args_overrides("triton") 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") def test_the_draft_backend_applies_to_the_draft_runner_only(self):
self.assertEqual(fields["speculative_draft_attention_backend"], "triton") self._seed(attention_backend="fa3")
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_variant_carries_the_targets_resolved_gate(self): draft = resolve_attention_backend_strs(
"""Publishing the variant re-projects the bags, so the gate travels.""" model_runner=self._runner(
self._seed(disable_chunked_prefix_cache=False) is_draft_worker=True, draft_attention_backend="triton"
get_context().override("test.gate", disable_chunked_prefix_cache=True) )
self.assertTrue(
draft_server_args_overrides("triton")["disable_chunked_prefix_cache"]
) )
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 --------------------- # -- the scheduler hands over the process's own config ---------------------