Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -624,7 +624,7 @@ def _get_k_and_s_triton(
:param page_indices: (num_pages,), int32/int64
:param seq_lens: tensor of sequence lens, int64
:param seq_len_sum: sum of all sequence len, int32
:param seq_len_sum: max of sequence len, int32
:param max_seq_len: max of sequence len, int32
:param page_size: int, typically 64
:param index_head_dim: int, typically 128
:return: tuple of (k_out, s_out) where
Expand Down
12 changes: 9 additions & 3 deletions python/sglang/srt/layers/attention/nsa/nsa_indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,6 +97,11 @@ def get_indexer_seq_len_cpu(self) -> torch.Tensor:
Return: seq lens for each batch.
"""

def get_indexer_seq_len(self) -> torch.Tensor:
"""
Return: seq lens for each batch.
"""

def get_nsa_extend_len_cpu(self) -> List[int]:
"""
Return: extend seq lens for each batch.
Expand Down Expand Up @@ -538,11 +543,12 @@ def _get_topk_ragged(

ks, ke = metadata.get_indexer_kvcache_range()

seq_len_sum = forward_batch.seq_lens_sum
max_seq_len = torch.max(forward_batch.seq_lens_cpu).item()
indexer_seq_lens_cpu = metadata.get_indexer_seq_len_cpu()
seq_len_sum = torch.sum(indexer_seq_lens_cpu).item()
max_seq_len = torch.max(indexer_seq_lens_cpu).item()
k_fp8, k_scale = forward_batch.token_to_kv_pool.get_index_k_scale_buffer(
layer_id,
forward_batch.seq_lens,
metadata.get_indexer_seq_len(),
block_tables,
seq_len_sum,
max_seq_len,
Expand Down
8 changes: 8 additions & 0 deletions python/sglang/srt/layers/attention/nsa_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -138,6 +138,8 @@ class NSAMetadata:
indexer_k_start_end: Optional[Tuple[torch.Tensor, torch.Tensor]] = None
# seq lens for each batch.
indexer_seq_lens_cpu: Optional[torch.Tensor] = None
# seq lens for each batch.
indexer_seq_lens: Optional[torch.Tensor] = None
# batch index for each token.
token_to_batch_idx: Optional[torch.Tensor] = None

Expand Down Expand Up @@ -194,6 +196,9 @@ def get_cu_seqlens_k(self) -> torch.Tensor:
def get_indexer_kvcache_range(self) -> Tuple[torch.Tensor, torch.Tensor]:
return self.attn_metadata.indexer_k_start_end

def get_indexer_seq_len(self) -> torch.Tensor:
return self.attn_metadata.indexer_seq_lens

def get_indexer_seq_len_cpu(self) -> torch.Tensor:
return self.attn_metadata.indexer_seq_lens_cpu

Expand Down Expand Up @@ -404,6 +409,7 @@ def init_forward_metadata(self, forward_batch: ForwardBatch):
bs_idx_cpu = None
# seq_len_cpu of selected sequences
indexer_seq_lens_cpu = forward_batch.seq_lens_cpu
indexer_seq_lens = forward_batch.seq_lens

if forward_batch.forward_mode.is_decode_or_idle():
extend_seq_lens_cpu = [1] * batch_size
Expand Down Expand Up @@ -504,6 +510,7 @@ def init_forward_metadata(self, forward_batch: ForwardBatch):
)
)
indexer_seq_lens_cpu = indexer_seq_lens_cpu[bs_idx_cpu]
indexer_seq_lens = indexer_seq_lens[bs_idx]
cache_seqlens_int32 = cache_seqlens_int32[bs_idx]
cu_seqlens_k = compute_cu_seqlens(cache_seqlens_int32)
max_seqlen_k = (
Expand Down Expand Up @@ -641,6 +648,7 @@ def init_forward_metadata(self, forward_batch: ForwardBatch):
topk_indices_offset=topk_indices_offset,
indexer_k_start_end=indexer_k_start_end,
indexer_seq_lens_cpu=indexer_seq_lens_cpu,
indexer_seq_lens=indexer_seq_lens,
token_to_batch_idx=token_to_batch_idx,
)
self.forward_metadata = metadata
Expand Down
4 changes: 4 additions & 0 deletions test/registered/kernels/test_nsa_indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -132,6 +132,10 @@ def get_indexer_seq_len_cpu(self) -> torch.Tensor:
"""Return: seq lens for each batch."""
return torch.tensor(self.seq_lens, dtype=torch.int32, device="cpu")

def get_indexer_seq_len(self) -> torch.Tensor:
"""Return: seq lens for each batch."""
return torch.tensor(self.seq_lens, dtype=torch.int32, device=self.device)

def get_nsa_extend_len_cpu(self) -> List[int]:
"""
Return: extend seq lens for each batch.
Expand Down
Loading