Fix hybrid linear attention misrouting plain-RadixAttention linear layers to the full backend (Ring-2.5-1T) (#26623)

Co-authored-by: Cheng Wan <54331508+ch-wan@users.noreply.github.com>
This commit is contained in:
Alison Shao
2026-06-02 16:24:49 -07:00
committed by GitHub
co-authored by Cheng Wan
parent 72929c7000
commit 76c9899da7
3 changed files with 10 additions and 21 deletions
@@ -782,22 +782,17 @@ class HybridLinearAttnBackend(AttentionBackend):
self.req_to_token_pool = full_attn_backend.req_to_token_pool self.req_to_token_pool = full_attn_backend.req_to_token_pool
def _is_full_attn( def _is_full_attn(
self, layer: Optional[RadixAttention], layer_id: Optional[int] = None self,
layer: Optional[Union[RadixAttention, RadixLinearAttention]],
layer_id: Optional[int] = None,
) -> bool: ) -> bool:
# Explicit linear-attention subclass → strong linear signal (KDA, GDN, # RadixLinearAttention is unambiguously a linear-attention layer.
# Qwen3-Next, Qwen3.5 main linear layers). # Everything else (including plain RadixAttention) must be classified by
# layer id: models like Bailing/Ring use a plain RadixAttention for their
# linear layers, so an `isinstance(layer, RadixAttention) -> full` shortcut
# would misroute those linear layers to the full-attention backend.
if isinstance(layer, RadixLinearAttention): if isinstance(layer, RadixLinearAttention):
return False return False
# Some hybrid models (Ling-2.5/2.6) wrap their linear layers in plain
# `RadixAttention` rather than `RadixLinearAttention`. Those wrappers
# set `_is_linear_attention=True` on the attn module so we can
# distinguish them from full-attention RadixAttention instances —
# including MTP/NEXTN draft layers, which are full and must default to
# the full-attn path.
if layer is not None and getattr(layer, "_is_linear_attention", False):
return False
if isinstance(layer, RadixAttention):
return True
if layer is not None: if layer is not None:
layer_id = layer.layer_id layer_id = layer.layer_id
@@ -508,12 +508,6 @@ class BailingMoELinearAttention(nn.Module):
quant_config=quant_config, quant_config=quant_config,
prefix=f"{prefix}.attn", prefix=f"{prefix}.attn",
) )
# Marker for HybridLinearAttnBackend._is_full_attn: Bailing wraps
# linear-attention layers in a plain RadixAttention, so the
# dispatcher can't tell from the type alone that this is a linear
# layer (would otherwise default to the full-attn backend, e.g. the
# same way MTP/NEXTN draft layers are routed).
self.attn._is_linear_attention = True
self.group_norm_size = getattr(config, "group_norm_size", 1) self.group_norm_size = getattr(config, "group_norm_size", 1)
self.rms_norm_eps = float(getattr(config, "rms_norm_eps", 1e-5)) self.rms_norm_eps = float(getattr(config, "rms_norm_eps", 1e-5))
@@ -3,7 +3,7 @@
Guards the hybrid linear / full attention dispatcher: Ling-2.5/2.6 Guards the hybrid linear / full attention dispatcher: Ling-2.5/2.6
has 32 layers with `layer_group_size=8`, so layers {7, 15, 23, 31} has 32 layers with `layer_group_size=8`, so layers {7, 15, 23, 31}
are full attention (MLA) and the rest are linear (Lightning seg_la). are full attention (MLA) and the rest are linear (Lightning seg_la).
Runs on the 8-GPU H200 runner with TP=4. Runs nightly on the 8-GPU H200 runner with TP=4.
""" """
import unittest import unittest
@@ -12,7 +12,7 @@ from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
from sglang.test.server_fixtures.default_fixture import DefaultServerBase from sglang.test.server_fixtures.default_fixture import DefaultServerBase
register_cuda_ci(est_time=600, stage="base-c", runner_config="8-gpu-h200") register_cuda_ci(est_time=600, suite="nightly-8-gpu-common", nightly=True)
class TestLing26Flash(GSM8KMixin, DefaultServerBase): class TestLing26Flash(GSM8KMixin, DefaultServerBase):