[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)
|
||||
|
||||
Reference in New Issue
Block a user