From 7dc7376697b04c9282d4c9dd621b7bf7343df5ad Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Thu, 4 Jun 2026 17:44:23 -0700 Subject: [PATCH] fix(attn): delegate init_mha_chunk_metadata in HybridLinearAttnBackend (#27316) Co-authored-by: Claude Opus 4.8 (1M context) --- .../attention/hybrid_linear_attn_backend.py | 10 ++ .../unittests/hybrid_linear/README.md | 64 ++++++++ .../unittests/hybrid_linear/__init__.py | 0 .../test_flashinfer_mla_chunk_metadata.py | 145 ++++++++++++++++++ 4 files changed, 219 insertions(+) create mode 100644 test/registered/attention/unittests/hybrid_linear/README.md create mode 100644 test/registered/attention/unittests/hybrid_linear/__init__.py create mode 100644 test/registered/attention/unittests/hybrid_linear/test_flashinfer_mla_chunk_metadata.py diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index 5ffcc3c1d..af3395999 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -807,6 +807,16 @@ class HybridLinearAttnBackend(AttentionBackend): for attn_backend in self.attn_backend_list: attn_backend.init_forward_metadata(forward_batch) + def init_mha_chunk_metadata( + self, forward_batch: ForwardBatch, disable_flashinfer_ragged: bool = False + ): + # Hybrid MLA models (Ring/Ling, Kimi-Linear) resolve this via + # get_attn_backend(), which returns this wrapper; delegate to the + # full-attn backend so its chunked/one-shot prefill metadata is planned. + init = getattr(self.full_attn_backend, "init_mha_chunk_metadata", None) + if init is not None: + init(forward_batch, disable_flashinfer_ragged) + def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int): for attn_backend in self.attn_backend_list: attn_backend.init_cuda_graph_state(max_bs, max_num_tokens) diff --git a/test/registered/attention/unittests/hybrid_linear/README.md b/test/registered/attention/unittests/hybrid_linear/README.md new file mode 100644 index 000000000..9eb203ccb --- /dev/null +++ b/test/registered/attention/unittests/hybrid_linear/README.md @@ -0,0 +1,64 @@ +# Hybrid linear-attention backend unit tests + +These tests cover `HybridLinearAttnBackend` — the wrapper that combines a +**full-attention** backend (DeepSeek-style MLA, e.g. `FlashInferMLAAttnBackend`) +with a **linear-attention** backend (GDN / Mamba2 / KDA / Lightning) for hybrid +models such as Ring/Ling (`bailing_moe_linear`), Qwen3-Next, NemotronH, etc. + +## Why this folder exists (the coverage gap) + +The per-method suites under `../mla`, `../gdn`, `../mamba`, ... each test one +backend in isolation. None of them exercise the **interaction** between the +wrapper and a real MLA full-attention backend during prefill: + +- `../mla/*` drives a **bare** `FlashInferMLAAttnBackend`, never the hybrid + wrapper. It also builds its mock runner with + `disable_chunked_prefix_cache=True` + `flashinfer_mla_disable_ragged=True` + (`kits/attention_unittest/attention_methods/mla_attention.py`), so the + `MHA_ONE_SHOT` / `MHA_CHUNKED_KV` prefill paths are never run. +- `../gdn`, `../kda`, ... *do* wrap in `HybridLinearAttnBackend`, but with + `full_attn_layers=[]` — so the full-attention (MLA) path inside the wrapper is + never reached. + +That gap hid a real production crash: the MLA prefill path plans its flashinfer +ragged wrapper via + +```python +if hasattr(get_attn_backend(), "init_mha_chunk_metadata"): + get_attn_backend().init_mha_chunk_metadata(forward_batch) +``` + +For a hybrid model `get_attn_backend()` returns the **wrapper**. The wrapper did +not expose `init_mha_chunk_metadata`, so the guard was silently False, the +`qo_indptr` / `kv_indptr` were never planned, and flashinfer aborted with: + +``` +ValueError: q.shape[0] (8218) does not match qo_indptr[-1] (800). +``` + +Fix: `HybridLinearAttnBackend.init_mha_chunk_metadata` delegates to the +full-attention backend (`hybrid_linear_attn_backend.py`). + +## Tests + +| File | What it pins | +|---|---| +| `test_flashinfer_mla_chunk_metadata.py` | Wraps a real chunk-KV-enabled `FlashInferMLAAttnBackend` (`full_attn_layers=[0]`) and asserts (1) the wrapper exposes `init_mha_chunk_metadata` and (2) the delegated call plans `qo_indptr[-1]` to the true extend-token count. | + +Run on a CUDA host: + +```bash +FLASHINFER_DISABLE_VERSION_CHECK=1 \ + python -m pytest test/registered/attention/unittests/hybrid_linear/ -v +``` + +## Next work + +- Reusable `kits/attention_unittest/attention_methods/hybrid_linear_attention.py` + helper that composes `full=MLA + linear=` with a + non-empty `full_attn_layers`, then drives a real + `DeepseekV2AttentionMLA.forward_normal_one_shot_core` / + `forward_normal_chunked_kv_core` prefill end-to-end (with numerical reference) + rather than asserting on planned metadata alone. The tiny `kv_lora_rank=32` + MLA config used here is too small for the flashinfer *ragged* prefill kernel's + head-dim constraints, so a true e2e variant needs production-sized head dims. diff --git a/test/registered/attention/unittests/hybrid_linear/__init__.py b/test/registered/attention/unittests/hybrid_linear/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/test/registered/attention/unittests/hybrid_linear/test_flashinfer_mla_chunk_metadata.py b/test/registered/attention/unittests/hybrid_linear/test_flashinfer_mla_chunk_metadata.py new file mode 100644 index 000000000..f8b1d4726 --- /dev/null +++ b/test/registered/attention/unittests/hybrid_linear/test_flashinfer_mla_chunk_metadata.py @@ -0,0 +1,145 @@ +"""Regression test: HybridLinearAttnBackend delegates init_mha_chunk_metadata. + +Hybrid MLA models (Ring/Ling, Kimi-Linear) run DeepSeek-style MLA on their +full-attention layers. Prefill plans the flashinfer ragged wrapper via +``hasattr(get_attn_backend(), "init_mha_chunk_metadata")`` (forward_mha.py). +For a hybrid model ``get_attn_backend()`` is the HybridLinearAttnBackend wrapper; +when it lacked the hook the guard was False, qo_indptr/kv_indptr were never +planned, and flashinfer raised "q.shape[0] does not match qo_indptr[-1]". + +The per-method MLA suite misses this: it drives the bare backend, and its mock +runner sets disable_chunked_prefix_cache=True + flashinfer_mla_disable_ragged=True +so the chunked-MHA path never runs. +""" + +import unittest +from pathlib import Path +from types import SimpleNamespace + +import torch + +from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( + HybridLinearAttnBackend, +) +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.srt.server_args import set_global_server_args_for_scheduler +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.attention_unittest.attention_methods.mla_attention import ( + DEFAULT_KV_LORA_RANK, + DEFAULT_MAX_CONTEXT_LEN, + MLAAttentionCase, + MockMLAModelRunner, + TinyMLAModelConfig, + _make_forward_batch, +) +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-large") + + +_KV_LORA_RANK = DEFAULT_KV_LORA_RANK +_QK_ROPE_HEAD_DIM = 0 + + +class _ChunkKVMLARunner(MockMLAModelRunner): + """MLA mock runner with chunked-prefix-cache (ragged MHA) enabled.""" + + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.server_args.disable_chunked_prefix_cache = False + self.server_args.flashinfer_mla_disable_ragged = False + set_global_server_args_for_scheduler(self.server_args) + + +def _make_case() -> MLAAttentionCase: + return MLAAttentionCase( + name="hybrid_one_shot_prefix", + backend="flashinfer", + forward_mode=ForwardMode.EXTEND, + num_heads=4, + page_size=1, + prefix_lens=(8, 4), + extend_lens=(5, 3), + ) + + +def _build_hybrid_backend(testcase, case: MLAAttentionCase): + model_config = TinyMLAModelConfig( + num_heads=case.num_heads, + kv_lora_rank=_KV_LORA_RANK, + qk_rope_head_dim=_QK_ROPE_HEAD_DIM, + hidden_size=64, + context_len=DEFAULT_MAX_CONTEXT_LEN, + ) + runner = _ChunkKVMLARunner( + case=case, + model_config=model_config, + dtype=torch.float16, + device="cuda", + max_context_len=DEFAULT_MAX_CONTEXT_LEN, + kv_lora_rank=_KV_LORA_RANK, + qk_rope_head_dim=_QK_ROPE_HEAD_DIM, + disable_cuda_graph=True, + disable_piecewise_cuda_graph=True, + runner_batch_size=None, + fp8_kv_cache=False, + ) + try: + from sglang.srt.layers.attention.flashinfer_mla_backend import ( + FlashInferMLAAttnBackend, + ) + + full_backend = FlashInferMLAAttnBackend(runner) + except (AssertionError, ImportError, ModuleNotFoundError) as exc: + testcase.skipTest(f"flashinfer MLA backend is not available: {exc}") + + if not getattr(full_backend, "enable_chunk_kv", False): + testcase.skipTest("chunk-KV path not enabled on this build") + + # full_attn_layers=[0] routes layer 0 through the MLA backend, as a hybrid + # model's full-attention layers do. The linear backend is unused here. + hybrid = HybridLinearAttnBackend( + full_backend, SimpleNamespace(), full_attn_layers=[0] + ) + return runner, full_backend, hybrid + + +@unittest.skipIf(not torch.cuda.is_available(), "CUDA is required") +class TestHybridLinearChunkMetadataDelegation(CustomTestCase): + def test_wrapper_exposes_chunk_metadata_hook(self): + _, full_backend, hybrid = _build_hybrid_backend(self, _make_case()) + self.assertTrue(hasattr(hybrid, "init_mha_chunk_metadata")) + self.assertTrue(hasattr(full_backend, "init_mha_chunk_metadata")) + + def test_delegation_plans_qo_indptr(self): + case = _make_case() + runner, full_backend, hybrid = _build_hybrid_backend(self, case) + + forward_batch = _make_forward_batch( + case, + runner, + max_context_len=DEFAULT_MAX_CONTEXT_LEN, + device="cuda", + ) + forward_batch.num_prefix_chunks = 0 + + bs = case.batch_size + sentinel = 999 + full_backend.qo_indptr[bs] = sentinel + + # disable_flashinfer_ragged=True plans qo_indptr without flashinfer calls. + hybrid.init_mha_chunk_metadata(forward_batch, disable_flashinfer_ragged=True) + + # qo_indptr[-1] must equal the extend (query) token count; a stale value + # here is the root cause of the q.shape[0] != qo_indptr[-1] crash. + planned = full_backend.mha_chunk_kv_cache.qo_indptr[bs].item() + self.assertEqual(planned, case.num_input_tokens) + self.assertNotEqual(planned, sentinel) + + +if __name__ == "__main__": + sys_path_parent = str(Path(__file__).resolve().parents[1]) + import sys + + sys.path.insert(0, sys_path_parent) + unittest.main()