[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,
|
q_lengths,
|
||||||
s_scales,
|
s_scales,
|
||||||
decay_scales,
|
decay_scales,
|
||||||
|
track_lens,
|
||||||
|
track_state_indices,
|
||||||
HEAD_DIM: tl.constexpr,
|
HEAD_DIM: tl.constexpr,
|
||||||
K_SPLIT_DIM: tl.constexpr,
|
K_SPLIT_DIM: tl.constexpr,
|
||||||
V_SPLIT_DIM: tl.constexpr,
|
V_SPLIT_DIM: tl.constexpr,
|
||||||
BLOCK: tl.constexpr,
|
BLOCK: tl.constexpr,
|
||||||
EVEN: tl.constexpr,
|
EVEN: tl.constexpr,
|
||||||
|
TRACK_STATE: tl.constexpr,
|
||||||
):
|
):
|
||||||
bid = tl.program_id(0)
|
bid = tl.program_id(0)
|
||||||
hid = tl.program_id(1)
|
hid = tl.program_id(1)
|
||||||
@@ -241,6 +244,11 @@ def seg_la_p_kernel(
|
|||||||
q_offset = tl.load(q_offsets + bid)
|
q_offset = tl.load(q_offsets + bid)
|
||||||
s_offset = tl.load(s_offsets + bid)
|
s_offset = tl.load(s_offsets + bid)
|
||||||
decay_scale = -tl.load(decay_scales + hid)
|
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_b = tl.arange(0, BLOCK)
|
||||||
offs_k = tl.arange(0, K_SPLIT_DIM)
|
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, :])
|
+ (offs_k[:, None] * HEAD_DIM + offs_v[None, :])
|
||||||
)
|
)
|
||||||
state = tl.load(s_ptrs, mask=s_scale > 0).to(tl.float32)
|
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):
|
for n in range(0, q_length, BLOCK):
|
||||||
n = tl.multiple_of(n, 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
|
o = tl.dot(q, state) * block_decay * softmax_scale + o
|
||||||
|
|
||||||
state = state * block_decay + tl.dot(k, v)
|
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:
|
if EVEN:
|
||||||
tl.store(out_ptrs + n * H * HEAD_DIM, o.to(Out.dtype.element_ty))
|
tl.store(out_ptrs + n * H * HEAD_DIM, o.to(Out.dtype.element_ty))
|
||||||
@@ -663,6 +689,8 @@ def seg_la_fwd(
|
|||||||
meta,
|
meta,
|
||||||
caches=None,
|
caches=None,
|
||||||
cache_indices=None,
|
cache_indices=None,
|
||||||
|
track_lens=None,
|
||||||
|
track_state_indices=None,
|
||||||
softmax_scale=None,
|
softmax_scale=None,
|
||||||
decouple=False,
|
decouple=False,
|
||||||
):
|
):
|
||||||
@@ -779,11 +807,18 @@ def seg_la_fwd(
|
|||||||
meta.q_lengths,
|
meta.q_lengths,
|
||||||
meta.s_scales,
|
meta.s_scales,
|
||||||
decay_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,
|
HEAD_DIM=HEAD_DIM,
|
||||||
K_SPLIT_DIM=K_SPLIT_DIM,
|
K_SPLIT_DIM=K_SPLIT_DIM,
|
||||||
V_SPLIT_DIM=V_SPLIT_DIM,
|
V_SPLIT_DIM=V_SPLIT_DIM,
|
||||||
BLOCK=BLOCK,
|
BLOCK=BLOCK,
|
||||||
EVEN=EVEN,
|
EVEN=EVEN,
|
||||||
|
TRACK_STATE=track_lens is not None and track_state_indices is not None,
|
||||||
num_warps=num_warps,
|
num_warps=num_warps,
|
||||||
num_stages=num_stages,
|
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.layers.radix_attention import RadixAttention
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -101,6 +101,11 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
forward_batch.forward_mode,
|
forward_batch.forward_mode,
|
||||||
forward_batch.spec_info,
|
forward_batch.spec_info,
|
||||||
forward_batch.seq_lens_cpu if not in_capture else None,
|
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(
|
self.forward_metadata = BailingLinearMetadata.prepare_decode(
|
||||||
metadata.query_start_loc,
|
metadata.query_start_loc,
|
||||||
@@ -108,6 +113,7 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
bs,
|
bs,
|
||||||
forward_batch.seq_lens,
|
forward_batch.seq_lens,
|
||||||
)
|
)
|
||||||
|
self.forward_metadata.mamba_track_indices = metadata.mamba_track_indices
|
||||||
|
|
||||||
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
||||||
metadata = self._forward_metadata(forward_batch)
|
metadata = self._forward_metadata(forward_batch)
|
||||||
@@ -116,6 +122,12 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
metadata.mamba_cache_indices,
|
metadata.mamba_cache_indices,
|
||||||
forward_batch,
|
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
|
@staticmethod
|
||||||
def _build_slope_tensor(
|
def _build_slope_tensor(
|
||||||
@@ -249,6 +261,8 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
mask=None,
|
mask=None,
|
||||||
temp_cache=None,
|
temp_cache=None,
|
||||||
intermediate_state_indices=None,
|
intermediate_state_indices=None,
|
||||||
|
track_lens=None,
|
||||||
|
track_state_indices=None,
|
||||||
):
|
):
|
||||||
q_offsets = metadata.query_start_loc
|
q_offsets = metadata.query_start_loc
|
||||||
|
|
||||||
@@ -270,10 +284,63 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
meta=seg_meta,
|
meta=seg_meta,
|
||||||
caches=temp_cache,
|
caches=temp_cache,
|
||||||
cache_indices=intermediate_state_indices,
|
cache_indices=intermediate_state_indices,
|
||||||
|
track_lens=track_lens,
|
||||||
|
track_state_indices=track_state_indices,
|
||||||
decouple=True,
|
decouple=True,
|
||||||
)
|
)
|
||||||
return hidden
|
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(
|
def forward_extend(
|
||||||
self,
|
self,
|
||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
@@ -315,6 +382,11 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
if forward_batch.forward_mode.is_target_verify()
|
if forward_batch.forward_mode.is_target_verify()
|
||||||
else None
|
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(
|
o = self._linear_attention_entry(
|
||||||
q,
|
q,
|
||||||
k,
|
k,
|
||||||
@@ -329,6 +401,8 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
else None
|
else None
|
||||||
),
|
),
|
||||||
intermediate_state_indices=intermediate_state_indices,
|
intermediate_state_indices=intermediate_state_indices,
|
||||||
|
track_lens=track_lens,
|
||||||
|
track_state_indices=track_state_indices,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -340,6 +414,15 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
and forward_batch.mamba_track_mask is not None
|
and forward_batch.mamba_track_mask is not None
|
||||||
):
|
):
|
||||||
# save mamba cache for extra buffer
|
# save mamba cache for extra buffer
|
||||||
|
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_mask = forward_batch.mamba_track_mask
|
||||||
mamba_track_indices = forward_batch.mamba_track_indices
|
mamba_track_indices = forward_batch.mamba_track_indices
|
||||||
dst_masked = mamba_track_indices[mamba_track_mask]
|
dst_masked = mamba_track_indices[mamba_track_mask]
|
||||||
@@ -380,4 +463,11 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"linear backend: {self.linear_backend} is not support for now"
|
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)
|
return o.view(-1, layer.tp_q_head_num * layer.v_head_dim)
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ from pathlib import Path
|
|||||||
|
|
||||||
import torch
|
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.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
@@ -156,6 +157,82 @@ class TestTritonLightningBackendCorrectness(CustomTestCase):
|
|||||||
with self.subTest(case=case.name, layout=layout):
|
with self.subTest(case=case.name, layout=layout):
|
||||||
run_lightning_attention_case(self, case, loc_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):
|
def test_runner_mode_cuda_graph_decode_cases(self):
|
||||||
for case in self.CUDA_GRAPH_CASES:
|
for case in self.CUDA_GRAPH_CASES:
|
||||||
with self.subTest(case=case.name, backend=case.backend):
|
with self.subTest(case=case.name, backend=case.backend):
|
||||||
|
|||||||
Reference in New Issue
Block a user