[FIX] Prevent Lightning Attention extra-buffer mamba state corruption (#29973)
Co-authored-by: 范坤 <fankun@U-Q0542J2D-0225.local>
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -4,6 +4,7 @@ from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.attention.linear.seg_la import SegLaMeta, seg_la_fwd
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
@@ -156,6 +157,82 @@ class TestTritonLightningBackendCorrectness(CustomTestCase):
|
||||
with self.subTest(case=case.name, layout=layout):
|
||||
run_lightning_attention_case(self, case, loc_layout=layout)
|
||||
|
||||
def test_seg_la_prefill_tracks_extra_buffer_state(self):
|
||||
torch.manual_seed(20260703)
|
||||
device = "cuda"
|
||||
dtype = torch.float32
|
||||
batch_size = 1
|
||||
num_heads = 2
|
||||
head_dim = 128
|
||||
extend_len = 96
|
||||
track_len = 64
|
||||
active_slot = 0
|
||||
track_slot = 1
|
||||
untouched_slot = 2
|
||||
|
||||
q = torch.randn(extend_len, num_heads, head_dim, dtype=dtype, device=device)
|
||||
k = torch.randn(extend_len, num_heads, head_dim, dtype=dtype, device=device)
|
||||
v = torch.randn(extend_len, num_heads, head_dim, dtype=dtype, device=device)
|
||||
slopes = torch.tensor([0.5, 0.25], dtype=torch.float32, device=device)
|
||||
|
||||
s = torch.empty(
|
||||
3, num_heads, head_dim, head_dim, dtype=torch.float32, device=device
|
||||
)
|
||||
s[active_slot] = torch.randn_like(s[active_slot]) * 0.05
|
||||
s[track_slot].fill_(12345.0)
|
||||
s[untouched_slot].fill_(54321.0)
|
||||
initial_state = s[active_slot].clone()
|
||||
untouched_state = s[untouched_slot].clone()
|
||||
|
||||
meta = SegLaMeta(
|
||||
batch_size=batch_size,
|
||||
max_q_length=None,
|
||||
q_offsets=torch.tensor([0, extend_len], dtype=torch.int32, device=device),
|
||||
s_offsets=torch.tensor([active_slot], dtype=torch.int32, device=device),
|
||||
q_lengths=torch.tensor([extend_len], dtype=torch.int32, device=device),
|
||||
s_scales=torch.tensor([True], dtype=torch.bool, device=device),
|
||||
mask=None,
|
||||
)
|
||||
|
||||
seg_la_fwd(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
s=s,
|
||||
decay_scales=slopes,
|
||||
meta=meta,
|
||||
track_lens=torch.tensor([track_len], dtype=torch.int32, device=device),
|
||||
track_state_indices=torch.tensor(
|
||||
[track_slot], dtype=torch.int64, device=device
|
||||
),
|
||||
decouple=True,
|
||||
)
|
||||
|
||||
expected = initial_state.float()
|
||||
expected_at_track = None
|
||||
decay = torch.exp(-slopes)
|
||||
for token_idx in range(extend_len):
|
||||
for head_idx in range(num_heads):
|
||||
expected[head_idx] = expected[head_idx] * decay[head_idx] + torch.outer(
|
||||
k[token_idx, head_idx], v[token_idx, head_idx]
|
||||
)
|
||||
if token_idx + 1 == track_len:
|
||||
expected_at_track = expected.clone()
|
||||
|
||||
torch.testing.assert_close(
|
||||
s[track_slot],
|
||||
expected_at_track,
|
||||
atol=3e-2,
|
||||
rtol=3e-2,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
s[active_slot],
|
||||
expected,
|
||||
atol=3e-2,
|
||||
rtol=3e-2,
|
||||
)
|
||||
torch.testing.assert_close(s[untouched_slot], untouched_state)
|
||||
|
||||
def test_runner_mode_cuda_graph_decode_cases(self):
|
||||
for case in self.CUDA_GRAPH_CASES:
|
||||
with self.subTest(case=case.name, backend=case.backend):
|
||||
|
||||
Reference in New Issue
Block a user