[FIX] Prevent Lightning Attention extra-buffer mamba state corruption (#29973)

Co-authored-by: 范坤 <fankun@U-Q0542J2D-0225.local>
This commit is contained in:
Fan Kun
2026-07-22 22:01:29 +08:00
committed by GitHub
co-authored by 范坤
parent 21065bc862
commit 004df6b520
3 changed files with 208 additions and 6 deletions
@@ -221,11 +221,14 @@ def seg_la_p_kernel(
q_lengths,
s_scales,
decay_scales,
track_lens,
track_state_indices,
HEAD_DIM: tl.constexpr,
K_SPLIT_DIM: tl.constexpr,
V_SPLIT_DIM: tl.constexpr,
BLOCK: tl.constexpr,
EVEN: tl.constexpr,
TRACK_STATE: tl.constexpr,
):
bid = tl.program_id(0)
hid = tl.program_id(1)
@@ -241,6 +244,11 @@ def seg_la_p_kernel(
q_offset = tl.load(q_offsets + bid)
s_offset = tl.load(s_offsets + bid)
decay_scale = -tl.load(decay_scales + hid)
track_len = 0
track_state_offset = -1
if TRACK_STATE:
track_len = tl.load(track_lens + bid)
track_state_offset = tl.load(track_state_indices + bid)
offs_b = tl.arange(0, BLOCK)
offs_k = tl.arange(0, K_SPLIT_DIM)
@@ -288,6 +296,15 @@ def seg_la_p_kernel(
+ (offs_k[:, None] * HEAD_DIM + offs_v[None, :])
)
state = tl.load(s_ptrs, mask=s_scale > 0).to(tl.float32)
if TRACK_STATE:
track_s_ptrs = (
S
+ track_state_offset * stride_s
+ hid * HEAD_DIM * HEAD_DIM
+ kid * HEAD_DIM * K_SPLIT_DIM
+ vid * V_SPLIT_DIM
+ (offs_k[:, None] * HEAD_DIM + offs_v[None, :])
)
for n in range(0, q_length, BLOCK):
n = tl.multiple_of(n, BLOCK)
@@ -330,6 +347,15 @@ def seg_la_p_kernel(
o = tl.dot(q, state) * block_decay * softmax_scale + o
state = state * block_decay + tl.dot(k, v)
if TRACK_STATE:
should_track = (
(track_state_offset >= 0) & (track_len > 0) & (n + b == track_len)
)
tl.store(
track_s_ptrs,
state.to(S.dtype.element_ty),
mask=should_track,
)
if EVEN:
tl.store(out_ptrs + n * H * HEAD_DIM, o.to(Out.dtype.element_ty))
@@ -663,6 +689,8 @@ def seg_la_fwd(
meta,
caches=None,
cache_indices=None,
track_lens=None,
track_state_indices=None,
softmax_scale=None,
decouple=False,
):
@@ -779,11 +807,18 @@ def seg_la_fwd(
meta.q_lengths,
meta.s_scales,
decay_scales,
track_lens if track_lens is not None else meta.q_lengths,
(
track_state_indices
if track_state_indices is not None
else meta.s_offsets
),
HEAD_DIM=HEAD_DIM,
K_SPLIT_DIM=K_SPLIT_DIM,
V_SPLIT_DIM=V_SPLIT_DIM,
BLOCK=BLOCK,
EVEN=EVEN,
TRACK_STATE=track_lens is not None and track_state_indices is not None,
num_warps=num_warps,
num_stages=num_stages,
)
@@ -15,7 +15,7 @@ from sglang.srt.layers.attention.linear.linear_metadata import (
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.runtime_context import get_parallel
from sglang.srt.runtime_context import get_parallel, get_server_args
logger = logging.getLogger(__name__)
@@ -101,6 +101,11 @@ class LightningAttentionBackend(MambaAttnBackendBase):
forward_batch.forward_mode,
forward_batch.spec_info,
forward_batch.seq_lens_cpu if not in_capture else None,
num_padding=(
0 if in_capture else getattr(forward_batch, "num_padding", None)
),
in_capture=in_capture,
mamba_track_indices=getattr(forward_batch, "mamba_track_indices", None),
)
self.forward_metadata = BailingLinearMetadata.prepare_decode(
metadata.query_start_loc,
@@ -108,6 +113,7 @@ class LightningAttentionBackend(MambaAttnBackendBase):
bs,
forward_batch.seq_lens,
)
self.forward_metadata.mamba_track_indices = metadata.mamba_track_indices
def init_forward_metadata(self, forward_batch: ForwardBatch):
metadata = self._forward_metadata(forward_batch)
@@ -116,6 +122,12 @@ class LightningAttentionBackend(MambaAttnBackendBase):
metadata.mamba_cache_indices,
forward_batch,
)
self.forward_metadata.mamba_track_indices = metadata.mamba_track_indices
self.forward_metadata.track_ssm_h_src = metadata.track_ssm_h_src
self.forward_metadata.track_ssm_h_dst = metadata.track_ssm_h_dst
self.forward_metadata.track_ssm_final_src = metadata.track_ssm_final_src
self.forward_metadata.track_ssm_final_dst = metadata.track_ssm_final_dst
self.forward_metadata.has_mamba_track_mask = metadata.has_mamba_track_mask
@staticmethod
def _build_slope_tensor(
@@ -249,6 +261,8 @@ class LightningAttentionBackend(MambaAttnBackendBase):
mask=None,
temp_cache=None,
intermediate_state_indices=None,
track_lens=None,
track_state_indices=None,
):
q_offsets = metadata.query_start_loc
@@ -270,10 +284,63 @@ class LightningAttentionBackend(MambaAttnBackendBase):
meta=seg_meta,
caches=temp_cache,
cache_indices=intermediate_state_indices,
track_lens=track_lens,
track_state_indices=track_state_indices,
decouple=True,
)
return hidden
def _prepare_seg_la_track_store(self, forward_batch, metadata):
if (
self.linear_backend != "seg_la"
or not metadata.has_mamba_track_mask
or metadata.num_prefills == 0
or forward_batch.mamba_track_mask is None
):
return None, None
h_dst = metadata.track_ssm_h_dst
if h_dst is None or h_dst.numel() == 0:
return None, None
mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size
num_prefills = metadata.num_prefills
track_mask = forward_batch.mamba_track_mask[:num_prefills]
extend_lens = forward_batch.extend_seq_lens[:num_prefills]
prefix_lens = forward_batch.extend_prefix_lens[:num_prefills]
track_seqlens = forward_batch.mamba_track_seqlens[:num_prefills]
lens_to_track = track_seqlens - prefix_lens
boundary_lens = (
lens_to_track // mamba_cache_chunk_size
) * mamba_cache_chunk_size
track_rows = (track_mask & (boundary_lens < extend_lens)).nonzero(
as_tuple=True
)[0]
if track_rows.numel() == 0:
return None, None
if h_dst.numel() != track_rows.numel():
raise RuntimeError(
"seg_la mamba track metadata mismatch: "
f"{h_dst.numel()} destination slots for {track_rows.numel()} rows"
)
track_lens = torch.zeros(
(metadata.batch_size,),
dtype=torch.int32,
device=metadata.mamba_cache_indices.device,
)
track_state_indices = torch.full(
(metadata.batch_size,),
-1,
dtype=h_dst.dtype,
device=metadata.mamba_cache_indices.device,
)
track_lens[track_rows] = boundary_lens[track_rows].to(torch.int32)
track_state_indices[track_rows] = h_dst
return track_lens, track_state_indices
def forward_extend(
self,
q: torch.Tensor,
@@ -315,6 +382,11 @@ class LightningAttentionBackend(MambaAttnBackendBase):
if forward_batch.forward_mode.is_target_verify()
else None
)
track_lens, track_state_indices = (
(None, None)
if forward_batch.forward_mode.is_target_verify()
else self._prepare_seg_la_track_store(forward_batch, metadata)
)
o = self._linear_attention_entry(
q,
k,
@@ -329,6 +401,8 @@ class LightningAttentionBackend(MambaAttnBackendBase):
else None
),
intermediate_state_indices=intermediate_state_indices,
track_lens=track_lens,
track_state_indices=track_state_indices,
)
else:
raise ValueError(
@@ -340,11 +414,20 @@ class LightningAttentionBackend(MambaAttnBackendBase):
and forward_batch.mamba_track_mask is not None
):
# save mamba cache for extra buffer
mamba_track_mask = forward_batch.mamba_track_mask
mamba_track_indices = forward_batch.mamba_track_indices
dst_masked = mamba_track_indices[mamba_track_mask]
src_masked = metadata.mamba_cache_indices[mamba_track_mask]
ssm_states[dst_masked] = ssm_states[src_masked]
if self.linear_backend == "seg_la":
if (
metadata.track_ssm_final_dst is not None
and metadata.track_ssm_final_dst.numel() > 0
):
ssm_states[metadata.track_ssm_final_dst] = ssm_states[
metadata.track_ssm_final_src
]
else:
mamba_track_mask = forward_batch.mamba_track_mask
mamba_track_indices = forward_batch.mamba_track_indices
dst_masked = mamba_track_indices[mamba_track_mask]
src_masked = metadata.mamba_cache_indices[mamba_track_mask]
ssm_states[dst_masked] = ssm_states[src_masked]
return o.view(-1, layer.tp_q_head_num * layer.v_head_dim)
@@ -380,4 +463,11 @@ class LightningAttentionBackend(MambaAttnBackendBase):
raise ValueError(
f"linear backend: {self.linear_backend} is not support for now"
)
self._track_mamba_state_decode(
forward_batch,
mamba_cache_params.conv[0],
ssm_states,
cache_indices,
)
return o.view(-1, layer.tp_q_head_num * layer.v_head_dim)