diff --git a/python/sglang/kernels/ops/attention/linear/seg_la.py b/python/sglang/kernels/ops/attention/linear/seg_la.py index 9b4e68d8d..abda8b529 100644 --- a/python/sglang/kernels/ops/attention/linear/seg_la.py +++ b/python/sglang/kernels/ops/attention/linear/seg_la.py @@ -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, ) diff --git a/python/sglang/srt/layers/attention/linear/lightning_backend.py b/python/sglang/srt/layers/attention/linear/lightning_backend.py index 43db9eb46..2a46ae1a1 100644 --- a/python/sglang/srt/layers/attention/linear/lightning_backend.py +++ b/python/sglang/srt/layers/attention/linear/lightning_backend.py @@ -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) diff --git a/test/registered/attention/unittests/lightning/test_triton.py b/test/registered/attention/unittests/lightning/test_triton.py index ee63b84dd..add011285 100644 --- a/test/registered/attention/unittests/lightning/test_triton.py +++ b/test/registered/attention/unittests/lightning/test_triton.py @@ -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):