diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index 637f0d7ca..56c2be07e 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -36,6 +36,7 @@ from sglang.srt.layers.attention.nsa.utils import ( from sglang.srt.layers.attention.utils import ( concat_mla_absorb_q_general, mla_quantize_and_rope_for_fp8, + seqlens_expand_triton, ) from sglang.srt.layers.dp_attention import get_attention_tp_size from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode @@ -434,24 +435,11 @@ class NativeSparseAttnBackend( extend_seq_lens_cpu = [self.speculative_num_draft_tokens] * batch_size forward_batch.extend_seq_lens_cpu = extend_seq_lens_cpu - seqlens_int32_cpu = [ - self.speculative_num_draft_tokens + kv_len - for kv_len in forward_batch.seq_lens_cpu.tolist() - ] - seqlens_expanded = torch.cat( - [ - torch.arange( - kv_len - qo_len + 1, - kv_len + 1, - dtype=torch.int32, - device=device, - ) - for qo_len, kv_len in zip( - extend_seq_lens_cpu, - seqlens_int32_cpu, - strict=True, - ) - ] + seqlens_expanded = seqlens_expand_triton( + torch.tensor(extend_seq_lens_cpu, dtype=torch.int32, device=device), + cache_seqlens_int32, + self.speculative_num_draft_tokens * batch_size, + self.speculative_num_draft_tokens, ) page_table = torch.repeat_interleave( page_table, repeats=self.speculative_num_draft_tokens, dim=0 @@ -474,20 +462,12 @@ class NativeSparseAttnBackend( dtype=torch.int32, device=device, ) - seqlens_expanded = torch.cat( - [ - torch.arange( - kv_len - qo_len + 1, - kv_len + 1, - dtype=torch.int32, - device=device, - ) - for qo_len, kv_len in zip( - forward_batch.extend_seq_lens_cpu, - forward_batch.seq_lens_cpu.tolist(), - strict=True, - ) - ] + + seqlens_expanded = seqlens_expand_triton( + forward_batch.extend_seq_lens, + cache_seqlens_int32, + sum(extend_seq_lens_cpu), + self.speculative_num_draft_tokens, ) if forward_batch.forward_mode.is_draft_extend_v2(): # DRAFT_EXTEND_V2: V2 worker pre-fills draft KV cache with ALL speculated @@ -1005,24 +985,13 @@ class NativeSparseAttnBackend( metadata.page_table_1[:, :max_seqlen_k].copy_(page_indices) extend_seq_lens_cpu = [self.speculative_num_draft_tokens] * bs - seqlens_int32_cpu = [ - self.speculative_num_draft_tokens + kv_len - for kv_len in seq_lens_cpu.tolist() - ] - seqlens_expanded = torch.cat( - [ - torch.arange( - kv_len - qo_len + 1, - kv_len + 1, - dtype=torch.int32, - device=self.device, - ) - for qo_len, kv_len in zip( - extend_seq_lens_cpu, - seqlens_int32_cpu, - strict=True, - ) - ] + seqlens_expanded = seqlens_expand_triton( + torch.tensor( + extend_seq_lens_cpu, dtype=torch.int32, device=self.device + ), + cache_seqlens, + self.speculative_num_draft_tokens * bs, + self.speculative_num_draft_tokens, ) metadata.nsa_seqlens_expanded.copy_(seqlens_expanded) nsa_cache_seqlens = compute_nsa_seqlens( @@ -1048,20 +1017,11 @@ class NativeSparseAttnBackend( page_indices ) - seqlens_expanded = torch.cat( - [ - torch.arange( - kv_len - qo_len + 1, - kv_len + 1, - dtype=torch.int32, - device=self.device, - ) - for qo_len, kv_len in zip( - extend_seq_lens_cpu, - seq_lens_cpu.tolist(), - strict=True, - ) - ] + seqlens_expanded = seqlens_expand_triton( + extend_seq_lens, + cache_seqlens, + sum(extend_seq_lens_cpu), + self.speculative_num_draft_tokens, ) metadata.nsa_seqlens_expanded[: seqlens_expanded.shape[0]].copy_( seqlens_expanded diff --git a/python/sglang/srt/layers/attention/utils.py b/python/sglang/srt/layers/attention/utils.py index 44d5edaaf..d679025dd 100644 --- a/python/sglang/srt/layers/attention/utils.py +++ b/python/sglang/srt/layers/attention/utils.py @@ -286,6 +286,67 @@ def pad_sequence_with_mask( return B, output, attn_mask +@triton.jit +def seqlens_expand_kernel( + extend_seq_lens_ptr, # [N] + seq_lens_ptr, # [N] + offsets_ptr, # [N+1] + output_ptr, # [sum(extend_seq_lens)] + N, + BLOCK: tl.constexpr, +): + pid = tl.program_id(0) + + if pid >= N: + return + + qo_len = tl.load(extend_seq_lens_ptr + pid) + kv_len = tl.load(seq_lens_ptr + pid) + + start = kv_len - qo_len + 1 + out_offset = tl.load(offsets_ptr + pid) + + offs = tl.arange(0, BLOCK) + mask = offs < qo_len + + values = start + offs + tl.store(output_ptr + out_offset + offs, values, mask=mask) + + +def seqlens_expand_triton( + extend_seq_lens: torch.Tensor, + seq_lens: torch.Tensor, + total_len: int, + max_q_len: int, +): + """ + extend_seq_lens: [N], int32, CUDA + seq_lens: [N], int32, CUDA + """ + assert extend_seq_lens.is_cuda + assert seq_lens.is_cuda + + N = extend_seq_lens.numel() + + offsets = torch.zeros(N + 1, device=extend_seq_lens.device, dtype=torch.int32) + offsets[1:] = torch.cumsum(extend_seq_lens, dim=0) + output = torch.empty(total_len, device=extend_seq_lens.device, dtype=torch.int32) + + BLOCK = triton.next_power_of_2(max_q_len) + grid = (N,) + + seqlens_expand_kernel[grid]( + extend_seq_lens, + seq_lens, + offsets, + output, + N, + BLOCK=BLOCK, + ) + + return output + + # When num_kv_heads=1, we have tensors with degenerate strides, # For example, as below, where we have stride[-3] == stride[-2]: # - shape: [num_pages, 1, 64, 128]