From d13d5c03abea535cdecefeaa176dc2afc9880140 Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Sat, 15 Aug 2026 00:38:00 -0700 Subject: [PATCH] config: decisions keyed on the attention backend read the configured pair `--attention-backend` is one field of three: a launch that sets only `--prefill-attention-backend` or `--decode-attention-backend` leaves the base field at `None`. Seven decisions read that base field alone and therefore answered from a field the operator never set. `attention_backends()` is the pair with the base-field fallback already applied, so each site now asks it for the half it actually needs: - `inkling_common/attn` assembles backend-specific kwargs (rel_bias / score mods) and gates its fused prologue; the backend those describe is the one `self.attn` dispatches to, so `serving_attention_backend()` selects the pair member by `forward_batch.forward_mode`, mirroring `HybridAttnBackend._select_backend` exactly -- draft-extend routes through the prefill branch like the dispatcher does -- and preferring the runner-stamped pair, so a draft runner answers with its own backend. That preference only works if every backend that can enter a ForwardContext carries the stamp, so `DraftBackendFactory._create_backend` now stamps its products with the backend it resolved (draft override first), and the draft-extend conv-sidecar wrapper copies the wrapped backend's stamp -- the replacement backends the spec workers install had no stamp at all and fell back to the target's configured pair. - The chunked-prefix-cache gate is a *prefill* feature -> prefill half. Reading the base field switched the feature off for every prefill-only configuration. - `init_deterministic_inference_config` maps *prefill* knobs (SPLIT_TILE / PREFILL_TRUNCATION_ALIGN) -> prefill half; the map missed and left truncation unset. - `two_batch_overlap` computes extend positions -> prefill half. - mrope's interleaved-rope kernel runs in both phases -> both halves must support triton. This one is not conservative when it misreads: `support_triton(None)` answers **True**, so a `--prefill-attention-backend torch_native` launch took the triton path. - The req-to-token writer has one caller, `alloc_for_extend` -> prefill half; its fallback pays several `.item()` syncs per request, so gating it on the decode half too would send every extend of a mixed launch through the slow path. `get_last_loc` (the spec-decode allocator's helper) keeps the both-halves reading: verify tokens are served by either half depending on `speculative_attention_mode`. - The flashinfer version floor is a guard; it never fired for a launch that pinned flashinfer through a split field. One more site the census found is not converted here: `gpt_oss` derives its `sinks` parameter dtype from the backend, and a single parameter dtype cannot serve a split pair (FA4 asserts bfloat16, trtllm_mha consumes float32), so that one is a behaviour question rather than a config-source one and is fixed in its own PR. `test_split_attention_backend_decisions.py` pins the callable decisions by calling them under a split-only publish, and pins the remaining ones statically -- the file/why map fails if any of them goes back to the base field (reverse-verified). It also asserts the `support_triton(None) is True` trap the sweep exists for. The stamp comes from the constructor, not the request: every factory leaf answers ("effective_name", backend), because several map entries do not build what their key says -- cutedsl_mla draft-extend builds the trtllm-mla backend, "nsa" is a deprecated alias building dsa, and the hybrid-linear entries pick fa3/intel_amx/triton by host, which no static rename table can express (a review catch: on Blackwell the alias stamp reached Inkling's per-forward kwargs assembly, which asserts a concrete kernel name, and crashed the first draft-extend forward). The stamping is pinned by unit tests, not only by a spec e2e: removing the child-stamping loop, stamping an alias from a leaf, or dropping the wrapper copy goes red (reverse-verified), and a static guard walks the factory source asserting no leaf answers an alias name. The child loop states its contract explicitly -- `create_decode_backend` passes `stamps_children=True` because its products are per-step containers by construction, so a container without `attn_backends` raises instead of being silently skipped by a defensive probe. The `_version` invalidation names its contract (autograd's in-place counter: private, chosen because it is the only per-tensor signal that ticks on copy_-style updates; removal fails loudly). The version-floor guard's file joins the pair-reader ratchet, and the one runner-seed chain read sharing the backend's __init__ (`speculative_eagle_topk`) reads the spec bag. --- .../srt/batch_overlap/two_batch_overlap.py | 10 +- python/sglang/srt/entrypoints/engine.py | 2 +- .../layers/attention/trtllm_mha_backend.py | 2 +- .../srt/layers/rotary_embedding/mrope.py | 6 +- python/sglang/srt/managers/scheduler.py | 8 +- python/sglang/srt/mem_cache/allocation.py | 14 +- .../model_runner_components/misc_utils.py | 11 +- .../sglang/srt/models/inkling_common/attn.py | 30 +- python/sglang/srt/speculative/draft_utils.py | 221 ++++++++--- .../test_split_attention_backend_decisions.py | 361 ++++++++++++++++++ 10 files changed, 586 insertions(+), 79 deletions(-) create mode 100644 test/registered/unit/test_split_attention_backend_decisions.py diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index 1e89fba14..ed13cdb2b 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -40,7 +40,11 @@ from sglang.srt.model_executor.forward_batch_info import ( compute_position, ) from sglang.srt.model_executor.forward_context import get_attn_backend -from sglang.srt.runtime_context import get_device, get_exec, get_parallel +from sglang.srt.runtime_context import ( + attention_backends, + get_device, + get_parallel, +) from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip @@ -631,8 +635,10 @@ class TboForwardBatchPreparer: device_field="extend_prefix_lens", sum_field=None, ) + # The prefill half: this computes extend positions. + prefill_backend, _ = attention_backends() _, child_b.extend_start_loc = compute_position( - get_exec().kernel.attention_backend, + prefill_backend, child_b.extend_prefix_lens, child_b.extend_seq_lens, child_b.extend_num_tokens, diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 75d33977b..4943fc2ac 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -1669,7 +1669,7 @@ def _set_envs_and_config(server_args: ServerArgs): # Check flashinfer version if not get_bool_env_var("SGLANG_SKIP_SGL_KERNEL_VERSION_CHECK"): - if server_args.attention_backend == "flashinfer": + if "flashinfer" in server_args.get_attention_backends(): assert_pkg_version( "flashinfer_python", "0.6.17", diff --git a/python/sglang/srt/layers/attention/trtllm_mha_backend.py b/python/sglang/srt/layers/attention/trtllm_mha_backend.py index f5b009407..8b86bebf8 100644 --- a/python/sglang/srt/layers/attention/trtllm_mha_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mha_backend.py @@ -171,7 +171,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): # Speculative decoding # Only support topk <= 1 for now. - self.topk = model_runner.server_args.speculative_eagle_topk or 0 + self.topk = get_spec().speculative_eagle_topk or 0 self.speculative_step_id = speculative_step_id self.target_verify_metadata = {} diff --git a/python/sglang/srt/layers/rotary_embedding/mrope.py b/python/sglang/srt/layers/rotary_embedding/mrope.py index a49246802..dd1d5df95 100644 --- a/python/sglang/srt/layers/rotary_embedding/mrope.py +++ b/python/sglang/srt/layers/rotary_embedding/mrope.py @@ -18,7 +18,7 @@ from sglang.srt.layers.rotary_embedding.yarn import ( yarn_get_mscale_simple, yarn_linear_ramp_mask, ) -from sglang.srt.runtime_context import get_exec +from sglang.srt.runtime_context import attention_backends, get_exec from sglang.srt.utils import ( cpu_has_amx_support, is_cuda, @@ -143,7 +143,9 @@ class MRotaryEmbedding(RotaryEmbedding): last_dim = cos_sin.size()[-1] cos, sin = cos_sin.chunk(2, dim=-1) if self.mrope_interleaved: - if support_triton(get_exec().kernel.attention_backend): + # Runs in prefill and decode: both halves must support triton. + prefill_backend, decode_backend = attention_backends() + if support_triton(prefill_backend) and support_triton(decode_backend): cos = apply_interleaved_rope_triton(cos, self.mrope_section) sin = apply_interleaved_rope_triton(sin, self.mrope_section) else: diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index eb6041cd3..fbe429700 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -28,6 +28,7 @@ from http import HTTPStatus from typing import TYPE_CHECKING, Any, Deque, Dict, List, Optional, Set, Tuple, Union from sglang.srt.runtime_context import ( + attention_backends, get_device, get_disagg, get_exec, @@ -1517,9 +1518,10 @@ class Scheduler( "flashinfer": ("SGLANG_FLASHINFER_PREFILL_SPLIT_TILE_SIZE", 4096), "triton": ("SGLANG_TRITON_PREFILL_TRUNCATION_ALIGN_SIZE", 4096), } - env_var, default_size = backend_sizes.get( - get_exec().kernel.attention_backend, (None, None) - ) + # Both entries are prefill knobs (SPLIT_TILE / PREFILL_TRUNCATION): + # the prefill half decides. + prefill_backend, _ = attention_backends() + env_var, default_size = backend_sizes.get(prefill_backend, (None, None)) self.truncation_align_size = ( get_int_env_var(env_var, default_size) if env_var else None ) diff --git a/python/sglang/srt/mem_cache/allocation.py b/python/sglang/srt/mem_cache/allocation.py index 8fb2e8088..f29b74c84 100644 --- a/python/sglang/srt/mem_cache/allocation.py +++ b/python/sglang/srt/mem_cache/allocation.py @@ -26,7 +26,7 @@ from sglang.srt.mem_cache.common import ( evict_from_tree_cache, ) from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool -from sglang.srt.runtime_context import get_exec, get_parallel +from sglang.srt.runtime_context import attention_backends, get_parallel from sglang.srt.utils import ( is_cpu, is_cuda, @@ -65,7 +65,10 @@ def write_cache_indices( prefix_tensors: list[torch.Tensor], req_to_token_pool: ReqToTokenPool, ): - if support_triton(get_exec().kernel.attention_backend): + # This writer's one caller is `alloc_for_extend`, so the prefill half + # decides; the fallback below pays several `.item()` syncs per request. + prefill_backend, _ = attention_backends() + if support_triton(prefill_backend): prefix_pointers = torch.tensor( [t.data_ptr() for t in prefix_tensors], dtype=torch.uint64, @@ -106,8 +109,11 @@ def get_last_loc( req_pool_indices_tensor: torch.Tensor, prefix_lens_tensor: torch.Tensor, ) -> torch.Tensor: - attn_backend = get_exec().kernel.attention_backend - uses_triton_dispatch = attn_backend not in ("ascend", "torch_native") + prefill_backend, decode_backend = attention_backends() + uses_triton_dispatch = prefill_backend not in ( + "ascend", + "torch_native", + ) and decode_backend not in ("ascend", "torch_native") if _is_hip and uses_triton_dispatch: # HIP-only: the legacy get_last_loc_triton kernel emits a diff --git a/python/sglang/srt/model_executor/model_runner_components/misc_utils.py b/python/sglang/srt/model_executor/model_runner_components/misc_utils.py index 653b3a18c..0f4e86cc2 100644 --- a/python/sglang/srt/model_executor/model_runner_components/misc_utils.py +++ b/python/sglang/srt/model_executor/model_runner_components/misc_utils.py @@ -8,7 +8,11 @@ from sglang.srt.configs.model_config import ( is_deepseek_dsa, is_kimi_k3, ) -from sglang.srt.runtime_context import get_context, get_exec, get_schedule +from sglang.srt.runtime_context import ( + attention_backends, + get_context, + get_schedule, +) from sglang.srt.server_args import CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS if TYPE_CHECKING: @@ -29,10 +33,11 @@ def maybe_disable_chunked_prefix_cache( # model's (often non-MLA) config must not flip the shared setting. if is_draft_worker: return + # Chunked prefix cache is a prefill feature: the prefill half decides. + prefill_backend, _ = attention_backends() if ( not use_mla_backend - or get_exec().kernel.attention_backend - not in CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS + or prefill_backend not in CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS ): if not get_schedule().disable_chunked_prefix_cache: get_context().override( diff --git a/python/sglang/srt/models/inkling_common/attn.py b/python/sglang/srt/models/inkling_common/attn.py index 7df207b06..895258811 100644 --- a/python/sglang/srt/models/inkling_common/attn.py +++ b/python/sglang/srt/models/inkling_common/attn.py @@ -30,9 +30,11 @@ from sglang.srt.models.inkling_common.norm import RMSNorm from sglang.srt.models.inkling_common.sconv import SconvType, ShortConvolution from sglang.srt.models.utils import apply_qk_norm from sglang.srt.runtime_context import ( + attention_backends, get_exec, get_model, get_parallel, + get_spec, ) from sglang.srt.utils import add_prefix, get_current_device_stream_fast @@ -105,6 +107,30 @@ _REL_PROJ_MATMUL_MAX_T = 48 _REL_PROJ_TAU_KERNEL_MAX_T = 32 +def serving_attention_backend(forward_batch: ForwardBatch) -> str: + """The pair member serving this forward. + + Mirrors ``HybridAttnBackend._select_backend`` exactly: decode for + decode/idle, the ``speculative_attention_mode`` half for target-verify, + prefill otherwise -- including draft-extend, which the hybrid dispatcher + routes through its prefill branch. The pair the runner stamped on its + backend wins over the configured one, so a draft runner answers with its + own backend. + """ + from sglang.srt.model_executor.forward_context import get_attn_backend + + backend = get_attn_backend() + configured_prefill, configured_decode = attention_backends() + prefill = backend.prefill_attention_backend_str or configured_prefill + decode = backend.decode_attention_backend_str or configured_decode + mode = forward_batch.forward_mode + if mode.is_decode_or_idle(): + return decode + if mode.is_target_verify(): + return decode if get_spec().speculative_attention_mode == "decode" else prefill + return prefill + + def _rel_proj_kernel_eligible(r: torch.Tensor) -> bool: """rel_proj_small_t input contract: bf16 CUDA, [t, h, d_rel] with a contiguous (h*d_rel) inner block (token rows may be strided), d_rel a @@ -735,7 +761,9 @@ class InklingAttention(nn.Module): apply_log_scaling = log_scaling_tau is not None and not self.is_local - attention_backend = get_exec().kernel.attention_backend + # The kwargs below must describe the backend `self.attn` dispatches + # this forward to. + attention_backend = serving_attention_backend(forward_batch) assert attention_backend in ("fa4", "triton") # The overlap threads a CUDA event into the FA4 sheared-bias kernel, so it # is FA4-only for now. diff --git a/python/sglang/srt/speculative/draft_utils.py b/python/sglang/srt/speculative/draft_utils.py index 821d28874..12bebcc1c 100644 --- a/python/sglang/srt/speculative/draft_utils.py +++ b/python/sglang/srt/speculative/draft_utils.py @@ -40,7 +40,11 @@ class DraftBackendFactory: self.draft_attn_backend = draft_model_runner.draft_attention_backend def _create_backend( - self, backend_name: str, backend_map: dict, error_template: str + self, + backend_name: str, + backend_map: dict, + error_template: str, + stamps_children: bool = False, ): # The split pair with the base-backend fallback already applied. prefill_backend, decode_backend = attention_backends() @@ -54,7 +58,15 @@ class DraftBackendFactory: if backend_type not in backend_map: raise ValueError(error_template.format(backend_type=backend_type)) - return backend_map[backend_type]() + stamp, backend = backend_map[backend_type]() + if backend is not None: + backend.prefill_attention_backend_str = stamp + backend.decode_attention_backend_str = stamp + if stamps_children: + for child in backend.attn_backends: + child.prefill_attention_backend_str = stamp + child.decode_attention_backend_str = stamp + return backend def create_decode_backend(self): # No multi-step draft backend for steps=0 (nospec) or steps=1. @@ -88,6 +100,7 @@ class DraftBackendFactory: "decode_attention_backend", backend_map, "EAGLE is not supported in decode attention backend {backend_type}", + stamps_children=True, ) def create_draft_extend_backend(self): @@ -125,27 +138,41 @@ class DraftBackendFactory: attn_backend_wrapper_for_draft_extend, ) - return attn_backend_wrapper_for_draft_extend(self.draft_model_runner, backend) + wrapped = attn_backend_wrapper_for_draft_extend( + self.draft_model_runner, backend + ) + if wrapped is not backend and wrapped is not None and backend is not None: + wrapped.prefill_attention_backend_str = ( + backend.prefill_attention_backend_str + ) + wrapped.decode_attention_backend_str = backend.decode_attention_backend_str + return wrapped def _create_dsa_decode_backend(self): from sglang.srt.layers.attention.dsa_backend import ( DeepseekSparseAttnMultiStepBackend, ) - return DeepseekSparseAttnMultiStepBackend( - self.draft_model_runner, - self.topk, - self.speculative_num_steps, - seed_dsa_topk_from_draft_extend=self.seed_dsa_topk_from_draft_extend, + return ( + "dsa", + DeepseekSparseAttnMultiStepBackend( + self.draft_model_runner, + self.topk, + self.speculative_num_steps, + seed_dsa_topk_from_draft_extend=self.seed_dsa_topk_from_draft_extend, + ), ) def _create_dsa_prefill_backend(self): from sglang.srt.layers.attention.dsa_backend import DeepseekSparseAttnBackend - return DeepseekSparseAttnBackend( - self.draft_model_runner, - skip_prefill=False, - seed_dsa_topk_from_draft_extend=self.seed_dsa_topk_from_draft_extend, + return ( + "dsa", + DeepseekSparseAttnBackend( + self.draft_model_runner, + skip_prefill=False, + seed_dsa_topk_from_draft_extend=self.seed_dsa_topk_from_draft_extend, + ), ) def _create_flashinfer_decode_backend(self): @@ -154,16 +181,22 @@ class DraftBackendFactory: FlashInferMultiStepDraftBackend, ) - return FlashInferMultiStepDraftBackend( - self.draft_model_runner, self.topk, self.speculative_num_steps + return ( + "flashinfer", + FlashInferMultiStepDraftBackend( + self.draft_model_runner, self.topk, self.speculative_num_steps + ), ) else: from sglang.srt.layers.attention.flashinfer_mla_backend import ( FlashInferMLAMultiStepDraftBackend, ) - return FlashInferMLAMultiStepDraftBackend( - self.draft_model_runner, self.topk, self.speculative_num_steps + return ( + "flashinfer", + FlashInferMLAMultiStepDraftBackend( + self.draft_model_runner, self.topk, self.speculative_num_steps + ), ) def _create_triton_decode_backend(self): @@ -171,8 +204,11 @@ class DraftBackendFactory: TritonMultiStepDraftBackend, ) - return TritonMultiStepDraftBackend( - self.draft_model_runner, self.topk, self.speculative_num_steps + return ( + "triton", + TritonMultiStepDraftBackend( + self.draft_model_runner, self.topk, self.speculative_num_steps + ), ) def _create_intel_amx_decode_backend(self): @@ -180,8 +216,11 @@ class DraftBackendFactory: IntelAMXMultiStepDraftBackend, ) - return IntelAMXMultiStepDraftBackend( - self.draft_model_runner, self.topk, self.speculative_num_steps + return ( + "intel_amx", + IntelAMXMultiStepDraftBackend( + self.draft_model_runner, self.topk, self.speculative_num_steps + ), ) def _create_hybrid_linear_attn_decode_backend(self): @@ -201,8 +240,11 @@ class DraftBackendFactory: def _create_aiter_decode_backend(self): from sglang.srt.layers.attention.aiter_backend import AiterMultiStepDraftBackend - return AiterMultiStepDraftBackend( - self.draft_model_runner, self.topk, self.speculative_num_steps + return ( + "aiter", + AiterMultiStepDraftBackend( + self.draft_model_runner, self.topk, self.speculative_num_steps + ), ) def _create_fa_decode_backend(self, fa_impl_ver: int = 3): @@ -215,11 +257,14 @@ class DraftBackendFactory: MusaFlashAttentionMultiStepBackend as FlashAttentionMultiStepBackend, ) - return FlashAttentionMultiStepBackend( - self.draft_model_runner, - self.topk, - self.speculative_num_steps, - fa_impl_ver=fa_impl_ver, + return ( + f"fa{fa_impl_ver}", + FlashAttentionMultiStepBackend( + self.draft_model_runner, + self.topk, + self.speculative_num_steps, + fa_impl_ver=fa_impl_ver, + ), ) def _create_fa3_decode_backend(self): @@ -233,8 +278,11 @@ class DraftBackendFactory: FlashMLAMultiStepDraftBackend, ) - return FlashMLAMultiStepDraftBackend( - self.draft_model_runner, self.topk, self.speculative_num_steps + return ( + "flashmla", + FlashMLAMultiStepDraftBackend( + self.draft_model_runner, self.topk, self.speculative_num_steps + ), ) def _create_trtllm_mha_decode_backend(self): @@ -242,8 +290,11 @@ class DraftBackendFactory: TRTLLMHAAttnMultiStepDraftBackend, ) - return TRTLLMHAAttnMultiStepDraftBackend( - self.draft_model_runner, self.topk, self.speculative_num_steps + return ( + "trtllm_mha", + TRTLLMHAAttnMultiStepDraftBackend( + self.draft_model_runner, self.topk, self.speculative_num_steps + ), ) def _create_trtllm_mla_decode_backend(self, backend: str = "trtllm-gen"): @@ -256,11 +307,14 @@ class DraftBackendFactory: TRTLLMMLAMultiStepDraftBackend, ) - return TRTLLMMLAMultiStepDraftBackend( - self.draft_model_runner, - self.topk, - self.speculative_num_steps, - backend=backend, + return ( + "trtllm_mla", + TRTLLMMLAMultiStepDraftBackend( + self.draft_model_runner, + self.topk, + self.speculative_num_steps, + backend=backend, + ), ) def _create_cutedsl_mla_decode_backend(self): @@ -273,8 +327,11 @@ class DraftBackendFactory: CuteDslMLAMultiStepDraftBackend, ) - return CuteDslMLAMultiStepDraftBackend( - self.draft_model_runner, self.topk, self.speculative_num_steps + return ( + "cutedsl_mla", + CuteDslMLAMultiStepDraftBackend( + self.draft_model_runner, self.topk, self.speculative_num_steps + ), ) def _create_tokenspeed_mla_decode_backend(self): @@ -287,8 +344,11 @@ class DraftBackendFactory: TokenspeedMLAMultiStepDraftBackend, ) - return TokenspeedMLAMultiStepDraftBackend( - self.draft_model_runner, self.topk, self.speculative_num_steps + return ( + "tokenspeed_mla", + TokenspeedMLAMultiStepDraftBackend( + self.draft_model_runner, self.topk, self.speculative_num_steps + ), ) def _create_ascend_decode_backend(self): @@ -296,8 +356,11 @@ class DraftBackendFactory: AscendAttnMultiStepDraftBackend, ) - return AscendAttnMultiStepDraftBackend( - self.draft_model_runner, self.topk, self.speculative_num_steps + return ( + "ascend", + AscendAttnMultiStepDraftBackend( + self.draft_model_runner, self.topk, self.speculative_num_steps + ), ) def _create_dsv4_decode_backend(self): @@ -307,8 +370,11 @@ class DraftBackendFactory: DeepseekV4AscendMultiStepDraftBackend, ) - return DeepseekV4AscendMultiStepDraftBackend( - self.draft_model_runner, self.topk, self.speculative_num_steps + return ( + "dsv4", + DeepseekV4AscendMultiStepDraftBackend( + self.draft_model_runner, self.topk, self.speculative_num_steps + ), ) elif is_hip(): from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import ( @@ -319,8 +385,11 @@ class DraftBackendFactory: DeepseekV4MultiStepBackend, ) - return DeepseekV4MultiStepBackend( - self.draft_model_runner, self.topk, self.speculative_num_steps + return ( + "dsv4", + DeepseekV4MultiStepBackend( + self.draft_model_runner, self.topk, self.speculative_num_steps + ), ) def _create_flashinfer_prefill_backend(self): @@ -329,28 +398,37 @@ class DraftBackendFactory: FlashInferAttnBackend, ) - return FlashInferAttnBackend(self.draft_model_runner, skip_prefill=False) + return ( + "flashinfer", + FlashInferAttnBackend(self.draft_model_runner, skip_prefill=False), + ) else: from sglang.srt.layers.attention.flashinfer_mla_backend import ( FlashInferMLAAttnBackend, ) - return FlashInferMLAAttnBackend(self.draft_model_runner, skip_prefill=False) + return ( + "flashinfer", + FlashInferMLAAttnBackend(self.draft_model_runner, skip_prefill=False), + ) def _create_triton_prefill_backend(self): from sglang.srt.layers.attention.triton_backend import TritonAttnBackend - return TritonAttnBackend(self.draft_model_runner, skip_prefill=False) + return ( + "triton", + TritonAttnBackend(self.draft_model_runner, skip_prefill=False), + ) def _create_intel_amx_prefill_backend(self): from sglang.srt.layers.attention.intel_amx_backend import IntelAMXAttnBackend - return IntelAMXAttnBackend(self.draft_model_runner) + return ("intel_amx", IntelAMXAttnBackend(self.draft_model_runner)) def _create_aiter_prefill_backend(self): from sglang.srt.layers.attention.aiter_backend import AiterAttnBackend - return AiterAttnBackend(self.draft_model_runner, skip_prefill=False) + return ("aiter", AiterAttnBackend(self.draft_model_runner, skip_prefill=False)) def _create_fa_prefill_backend(self, fa_impl_ver: int = 3): if not is_musa(): @@ -361,8 +439,11 @@ class DraftBackendFactory: from sglang.srt.hardware_backend.musa.attention.flashattention_backend import ( MusaFlashAttentionBackend as FlashAttentionBackend, ) - return FlashAttentionBackend( - self.draft_model_runner, skip_prefill=False, fa_impl_ver=fa_impl_ver + return ( + f"fa{fa_impl_ver}", + FlashAttentionBackend( + self.draft_model_runner, skip_prefill=False, fa_impl_ver=fa_impl_ver + ), ) def _create_fa3_prefill_backend(self): @@ -374,7 +455,10 @@ class DraftBackendFactory: def _create_trtllm_mha_prefill_backend(self): from sglang.srt.layers.attention.trtllm_mha_backend import TRTLLMHAAttnBackend - return TRTLLMHAAttnBackend(self.draft_model_runner, skip_prefill=False) + return ( + "trtllm_mha", + TRTLLMHAAttnBackend(self.draft_model_runner, skip_prefill=False), + ) def _create_trtllm_mla_prefill_backend(self): if not self.draft_model_runner.use_mla_backend: @@ -384,7 +468,10 @@ class DraftBackendFactory: from sglang.srt.layers.attention.trtllm_mla_backend import TRTLLMMLABackend - return TRTLLMMLABackend(self.draft_model_runner, skip_prefill=False) + return ( + "trtllm_mla", + TRTLLMMLABackend(self.draft_model_runner, skip_prefill=False), + ) def _create_tokenspeed_mla_prefill_backend(self): if not self.draft_model_runner.use_mla_backend: @@ -396,19 +483,25 @@ class DraftBackendFactory: TokenspeedMLABackend, ) - return TokenspeedMLABackend(self.draft_model_runner, skip_prefill=False) + return ( + "tokenspeed_mla", + TokenspeedMLABackend(self.draft_model_runner, skip_prefill=False), + ) def _create_ascend_prefill_backend(self): from sglang.srt.hardware_backend.npu.attention.ascend_backend import ( AscendAttnBackend, ) - return AscendAttnBackend(self.draft_model_runner) + return ("ascend", AscendAttnBackend(self.draft_model_runner)) def _create_flashmla_prefill_backend(self): from sglang.srt.layers.attention.flashmla_backend import FlashMLABackend - return FlashMLABackend(self.draft_model_runner, skip_prefill=False) + return ( + "flashmla", + FlashMLABackend(self.draft_model_runner, skip_prefill=False), + ) def _create_dsv4_prefill_backend(self): # On NPU the "dsv4" backend resolves to the Ascend V4 subclass; its @@ -418,17 +511,21 @@ class DraftBackendFactory: ATTENTION_BACKENDS, ) - return ATTENTION_BACKENDS["dsv4"](self.draft_model_runner) + return ("dsv4", ATTENTION_BACKENDS["dsv4"](self.draft_model_runner)) elif is_hip(): from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import ( DeepseekV4HipRadixBackend, ) - return DeepseekV4HipRadixBackend( - self.draft_model_runner, skip_prefill=False + return ( + "dsv4", + DeepseekV4HipRadixBackend(self.draft_model_runner, skip_prefill=False), ) from sglang.srt.layers.attention.deepseek_v4_backend import ( DeepseekV4AttnBackend, ) - return DeepseekV4AttnBackend(self.draft_model_runner, skip_prefill=False) + return ( + "dsv4", + DeepseekV4AttnBackend(self.draft_model_runner, skip_prefill=False), + ) diff --git a/test/registered/unit/test_split_attention_backend_decisions.py b/test/registered/unit/test_split_attention_backend_decisions.py new file mode 100644 index 000000000..3812d658e --- /dev/null +++ b/test/registered/unit/test_split_attention_backend_decisions.py @@ -0,0 +1,361 @@ +"""Decisions keyed on the attention backend must read the configured pair. + +`--attention-backend` is one of three fields: the base one and the two split +ones (`--prefill-attention-backend` / `--decode-attention-backend`). A launch +that sets only a split field leaves the base at `None`, so a decision that reads +`attention_backend` alone answers from a field the operator never set. What that +cost, before the sweep these cases guard: + + - a weight sized for the wrong dtype (gpt-oss `sinks` under trtllm_mha; that + decision is now gone -- the weight stays bfloat16 and the trtllm backend + upcasts at its call site, since at model-build time no config read can say + which backend serves this runner's forwards), + - a prefill feature switched off (chunked prefix cache), + - a triton kernel chosen for a backend that cannot host it (`support_triton(None)` + answers True), in mrope and in the req-to-token writer, + - a version guard that never fires (flashinfer), + - a deterministic-inference knob left unset (prefill truncation align). + +`attention_backends()` is the shared answer: the pair with the base-field +fallback applied. The callable decisions are checked by calling them; the rest +are pinned statically, since reproducing them means building a model or a +scheduler. +""" + +import ast +import unittest +from pathlib import Path + +from sglang.srt.runtime_context import attention_backends, get_context +from sglang.srt.server_args import ServerArgs +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + +import sglang + +_PACKAGE_ROOT = Path(next(iter(sglang.__path__))) / "srt" + +# The decisions this file is about, and which half of the pair each one needs. +# A base-only read here is the regression; the resolution pipeline and the two +# modules that own the config are exempt because "did the operator pin the base +# field?" is a real question *there*. +_PAIR_READERS = { + "models/inkling_common/attn.py": "the half serving the forward (mirrors hybrid dispatch)", + "model_executor/model_runner_components/misc_utils.py": "prefill (chunked prefix cache)", + "layers/rotary_embedding/mrope.py": "both (triton availability)", + "mem_cache/allocation.py": "prefill (req-to-token writer); both (get_last_loc)", + "batch_overlap/two_batch_overlap.py": "prefill (extend positions)", + "managers/scheduler.py": "prefill (truncation align knobs)", + "entrypoints/engine.py": "either half (flashinfer version floor)", +} + + +class TestSplitBackendsReachTheDecisions(CustomTestCase): + def setUp(self): + super().setUp() + self._saved = get_context()._server_args + + def tearDown(self): + if self._saved is not None: + get_context().set_server_args(self._saved) + super().tearDown() + + def _publish(self, **fields): + override = get_context().override_server_args(**fields) + override.install() + self.addCleanup(override.restore) + + def test_the_pair_is_what_a_split_only_launch_configures(self): + self._publish( + attention_backend=None, + prefill_attention_backend="triton", + decode_attention_backend="trtllm_mha", + ) + self.assertEqual(attention_backends(), ("triton", "trtllm_mha")) + + def test_chunked_prefix_cache_follows_the_prefill_backend(self): + from sglang.srt.model_executor.model_runner_components.misc_utils import ( + maybe_disable_chunked_prefix_cache, + ) + from sglang.srt.runtime_context import get_schedule + + # A prefill backend that supports the feature, configured *only* through + # the split field: the gate must leave it on. + self._publish( + attention_backend=None, + prefill_attention_backend="fa3", + decode_attention_backend="triton", + disable_chunked_prefix_cache=False, + ) + maybe_disable_chunked_prefix_cache(use_mla_backend=True, is_draft_worker=False) + self.assertFalse(get_schedule().disable_chunked_prefix_cache) + + # And an unsupported one still switches it off. + self._publish( + attention_backend=None, + prefill_attention_backend="torch_native", + decode_attention_backend="fa3", + disable_chunked_prefix_cache=False, + ) + maybe_disable_chunked_prefix_cache(use_mla_backend=True, is_draft_worker=False) + self.assertTrue(get_schedule().disable_chunked_prefix_cache) + + def test_inkling_selects_the_half_serving_the_forward(self): + from types import SimpleNamespace + + from sglang.srt.model_executor.forward_batch_info import ForwardMode + from sglang.srt.model_executor.forward_context import ( + ForwardContext, + forward_context, + ) + from sglang.srt.models.inkling_common.attn import serving_attention_backend + + self._publish( + attention_backend=None, + prefill_attention_backend="triton", + decode_attention_backend="fa4", + speculative_attention_mode="prefill", + ) + + def batch(mode): + return SimpleNamespace(forward_mode=mode) + + unstamped = SimpleNamespace( + prefill_attention_backend_str=None, decode_attention_backend_str=None + ) + with forward_context(ForwardContext(attn_backend=unstamped)): + self.assertEqual( + serving_attention_backend(batch(ForwardMode.EXTEND)), "triton" + ) + self.assertEqual( + serving_attention_backend(batch(ForwardMode.DECODE)), "fa4" + ) + self.assertEqual(serving_attention_backend(batch(ForwardMode.IDLE)), "fa4") + self.assertEqual( + serving_attention_backend(batch(ForwardMode.TARGET_VERIFY)), "triton" + ) + + # Draft-extend routes through the hybrid dispatcher's prefill branch + # regardless of the spec mode. + with forward_context(ForwardContext(attn_backend=unstamped)): + self.assertEqual( + serving_attention_backend(batch(ForwardMode.DRAFT_EXTEND_V2)), + "triton", + ) + + self._publish( + attention_backend=None, + prefill_attention_backend="triton", + decode_attention_backend="fa4", + speculative_attention_mode="decode", + ) + with forward_context(ForwardContext(attn_backend=unstamped)): + self.assertEqual( + serving_attention_backend(batch(ForwardMode.TARGET_VERIFY)), "fa4" + ) + self.assertEqual( + serving_attention_backend(batch(ForwardMode.DRAFT_EXTEND_V2)), + "triton", + ) + + # The pair the runner stamped on its backend wins over the bags: a + # draft runner serves every phase with its own backend. + stamped = SimpleNamespace( + prefill_attention_backend_str="fa4", decode_attention_backend_str="fa4" + ) + with forward_context(ForwardContext(attn_backend=stamped)): + self.assertEqual( + serving_attention_backend(batch(ForwardMode.EXTEND)), "fa4" + ) + + def test_the_flashinfer_version_guard_sees_a_split_launch(self): + # The launcher runs before any publish, so it asks the record; the + # member and the accessor answer the same pair. + args = ServerArgs.__new__(ServerArgs) + for name, value in ( + ("attention_backend", None), + ("prefill_attention_backend", None), + ("decode_attention_backend", "flashinfer"), + ): + object.__setattr__(args, name, value) + self.assertIn("flashinfer", args.get_attention_backends()) + + def test_support_triton_is_the_regression_being_guarded(self): + from sglang.srt.utils.common import support_triton + + # This is why a base-only read is not merely imprecise: the unset field + # reads as "supported". + self.assertTrue(support_triton(None)) + + def test_no_listed_decision_reads_the_base_field_alone(self): + offenders = [] + for rel, why in _PAIR_READERS.items(): + tree = ast.parse((_PACKAGE_ROOT / rel).read_text()) + for node in ast.walk(tree): + # Any attribute read named `attention_backend` is the base + # field, whatever the base expression is spelled as -- a bag + # chain, a record, or a local alias of either + # (`k = get_exec().kernel; k.attention_backend`). The pair + # helpers are calls, not attributes, so they never match. + if isinstance(node, ast.Attribute) and node.attr == "attention_backend": + offenders.append(f"{rel}:{node.lineno}: base-only read ({why})") + self.assertEqual( + [], + offenders, + "these decisions must read attention_backends() (the pair with the " + "base-field fallback), not the base field:\n" + "\n".join(offenders), + ) + + +class TestDraftFactoryStamping(CustomTestCase): + """The factory's products carry the stamp `serving_attention_backend` + prefers -- removing the child-stamping loop, the `cutedsl_mla` rename, or + the wrapper copy goes red here, not only in a spec e2e.""" + + def setUp(self): + super().setUp() + self._saved = get_context()._server_args + + def tearDown(self): + if self._saved is not None: + get_context().set_server_args(self._saved) + super().tearDown() + + def _publish(self, **fields): + override = get_context().override_server_args(**fields) + override.install() + self.addCleanup(override.restore) + + def _factory(self, draft_backend=None): + from types import SimpleNamespace + + from sglang.srt.speculative.draft_utils import DraftBackendFactory + + runner = SimpleNamespace(draft_attention_backend=draft_backend) + return DraftBackendFactory(runner, topk=1, speculative_num_steps=2) + + def test_a_decode_container_stamps_its_per_step_children(self): + from types import SimpleNamespace + + self._publish(attention_backend="triton") + children = [SimpleNamespace(), SimpleNamespace()] + container = SimpleNamespace(attn_backends=children) + product = self._factory()._create_backend( + "decode_attention_backend", + {"triton": lambda: ("triton", container)}, + "unsupported {backend_type}", + stamps_children=True, + ) + # EAGLE's eager loop puts the children into the ForwardContext + # directly, so an unstamped child answers with the target pair. + for obj in [product, *children]: + self.assertEqual(obj.prefill_attention_backend_str, "triton") + self.assertEqual(obj.decode_attention_backend_str, "triton") + + def test_the_draft_override_wins_over_the_published_pair(self): + from types import SimpleNamespace + + self._publish(attention_backend="triton") + product = self._factory(draft_backend="trtllm_mha")._create_backend( + "decode_attention_backend", + {"trtllm_mha": lambda: ("trtllm_mha", SimpleNamespace(attn_backends=[]))}, + "unsupported {backend_type}", + stamps_children=True, + ) + self.assertEqual(product.prefill_attention_backend_str, "trtllm_mha") + + def test_cutedsl_draft_extend_stamps_the_effective_kernel(self): + # cutedsl_mla only supports decode; the draft-extend map builds the + # trtllm-mla backend, and the *constructor* answers the effective + # name, so the stamp says what actually runs with no rename table -- + # and the conv-sidecar wrapper that enters the ForwardContext must + # answer with the wrapped backend's stamp. + from types import SimpleNamespace + from unittest import mock + + from sglang.srt.layers.attention import attention_registry + + self._publish(attention_backend="triton") + factory = self._factory(draft_backend="cutedsl_mla") + built = SimpleNamespace() + factory._create_trtllm_mla_prefill_backend = lambda: ("trtllm_mla", built) + wrapper = SimpleNamespace() + with mock.patch.object( + attention_registry, + "attn_backend_wrapper_for_draft_extend", + lambda runner, backend: wrapper, + ): + product = factory.create_draft_extend_backend() + self.assertIs(product, wrapper) + self.assertEqual(built.prefill_attention_backend_str, "trtllm_mla") + self.assertEqual(product.prefill_attention_backend_str, "trtllm_mla") + self.assertEqual(product.decode_attention_backend_str, "trtllm_mla") + + def test_a_host_dependent_alias_stamps_the_concrete_kernel(self): + # `hybrid_linear_attn` picks fa3/intel_amx/triton by host inside its + # constructor, so no static rename can say what it builds -- the + # constructor's own answer is the stamp. A stamp that repeats the + # alias crashes Inkling's per-forward kwargs assembly, which asserts + # the name is a concrete kernel. + from types import SimpleNamespace + + self._publish(attention_backend="hybrid_linear_attn") + concrete = SimpleNamespace(attn_backends=[SimpleNamespace()]) + product = self._factory()._create_backend( + "decode_attention_backend", + {"hybrid_linear_attn": lambda: ("triton", concrete)}, + "unsupported {backend_type}", + stamps_children=True, + ) + self.assertEqual(product.prefill_attention_backend_str, "triton") + self.assertEqual( + product.attn_backends[0].decode_attention_backend_str, "triton" + ) + + def test_the_real_map_never_stamps_an_alias(self): + # The factory's real constructors each answer their effective name; + # this pins that no map key with an aliased or host-dependent + # constructor ("nsa", "cutedsl_mla", "hybrid_linear_attn") can leak + # its request name into a stamp: whatever the leaf built, the name it + # answered is a concrete kernel, never one of the alias keys. + import ast as _ast + import inspect + + from sglang.srt.speculative import draft_utils + + tree = _ast.parse(inspect.getsource(draft_utils)) + offenders = [] + for node in _ast.walk(tree): + if not isinstance(node, _ast.FunctionDef): + continue + if not node.name.startswith("_create_") or "_backend" not in node.name: + continue + for ret in _ast.walk(node): + if not isinstance(ret, _ast.Return) or ret.value is None: + continue + # Leaf returns are ("name", ctor(...)); delegations return the + # inner call. A bare backend return would silently miss the + # stamp contract. + if isinstance(ret.value, _ast.Tuple): + name = ret.value.elts[0] + if isinstance(name, _ast.Constant) and name.value in ( + "nsa", + "hybrid_linear_attn", + ): + offenders.append(f"{node.name}: stamps alias {name.value!r}") + elif isinstance(ret.value, _ast.Call): + fn = ret.value.func + is_delegation = isinstance( + fn, _ast.Attribute + ) and fn.attr.startswith("_create_") + if not is_delegation: + offenders.append( + f"{node.name}: returns a bare backend (no effective name)" + ) + self.assertEqual([], offenders, "\n".join(offenders)) + + +if __name__ == "__main__": + unittest.main()