Batch CP shared-KV index work for bs>1 fast paths

The bs>1 path needs index top-k, shared-index prepare, current-index compact, and current-slot compose to consume flattened batch descriptors instead of falling back to per-request or per-segment Python/Torch work. This change wires SGLang to the new TAI batch prepare kernels, keeps fallback explicit, and records the remaining HiCache/load-backup gaps in the bs>1 workstream docs.

Constraint: CP shared-KV bs>1 must reuse fast paths rather than adding slow batch-only fallbacks

Constraint: No new collective operations were introduced

Rejected: Leave current-only cp_index as Python slice/cat | it keeps per-segment overhead in the short-extend bs>1 case

Rejected: Infer max segment lengths from CUDA descriptor tensors | .item() would add CPU synchronization on the hot path

Confidence: medium

Scope-risk: moderate

Directive: Do not remove the explicit fallback warnings without verifying the corresponding TAI symbols are present in production

Tested: local py_compile for touched SGLang files

Tested: remote g0034 test_nsa_cp_utils.py passed, 53 tests

Tested: remote g0034 test_fill_current_index_page_slots_uses_tai_kernel_when_available passed

Not-tested: full ETE bs>1 traffic with HiCache load/backup and draft/EAGLE enabled
This commit is contained in:
laoyao0822
2026-06-03 08:45:32 +08:00
parent d36f62a3cd
commit 19dcd6c4dc
6 changed files with 1132 additions and 113 deletions
+363 -44
View File
@@ -1203,6 +1203,15 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
def get_page_table_64(self):
return logical_pages
def get_page_table_1(self):
return torch.empty((2, 1000), dtype=torch.int32)
def get_page_table_1(self):
return torch.empty((2, 1000), dtype=torch.int32)
def get_page_table_1(self):
return torch.empty((2, 1000), dtype=torch.int32)
def fake_materialize(forward_batch, layer_id, logical_page_table):
materialize_calls.append((layer_id, logical_page_table))
return materialized_index, dense_pages
@@ -1304,6 +1313,9 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
def get_page_table_64(self):
return logical_pages
def get_page_table_1(self):
return torch.empty((2, 1000), dtype=torch.int32)
def fake_materialize(forward_batch, layer_id, logical_page_table):
materialize_calls.append((layer_id, logical_page_table))
return materialized_index, dense_pages
@@ -1329,6 +1341,7 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
"batch_idx": batch_idx,
"kv_len": kv_len,
"actual_seq_q": actual_seq_q,
"cp_index": cp_index,
"q": q_fp8.flatten().tolist(),
"weights": weights.flatten().tolist(),
"actual_seq_q_tensor": actual_seq_q_tensor,
@@ -1338,17 +1351,20 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
"current_index_kv": current_index_kv,
}
)
return torch.full(
(actual_seq_q, 2),
len(topk_calls),
dtype=torch.int32,
)
rows = int(q_fp8.shape[0])
return torch.arange(1, rows + 1, dtype=torch.int32).view(rows, 1).repeat(1, 2)
indexer._maybe_materialize_shared_index_buffer = fake_materialize
indexer._get_topk_ragged_with_cp = fake_get_topk
forward_batch = SimpleNamespace(
batch_size=2,
forward_mode=SimpleNamespace(
is_extend_without_speculative=lambda: True,
),
extend_prefix_lens_cpu=[0, 0],
extend_seq_lens_cpu=[1000, 1000],
seq_lens_cpu=torch.tensor([1000, 1000], dtype=torch.int64),
nsa_cp_metadata=NSAContextParallelMetadata(
batch_size=2,
kv_len_prev=100,
@@ -1357,8 +1373,8 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
actual_seq_q_next=1,
actual_seq_q_prev_cu_tensor=torch.tensor([0, 2], dtype=torch.int32),
actual_seq_q_next_cu_tensor=torch.tensor([0, 1], dtype=torch.int32),
request_kv_len_prev=[100, 300],
request_kv_len_next=[200, 400],
request_kv_len_prev=[2, 1],
request_kv_len_next=[3, 4],
request_actual_seq_q_prev=[2, 1],
request_actual_seq_q_next=[1, 3],
),
@@ -1378,32 +1394,26 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
self.assertEqual(len(materialize_calls), 1)
self.assertIs(materialize_calls[0][1], logical_pages)
self.assertEqual(len(topk_calls), 1)
self.assertEqual(topk_calls[0]["batch_idx"], 0)
self.assertEqual(topk_calls[0]["actual_seq_q"], 7)
self.assertEqual(
[
(
call["batch_idx"],
call["kv_len"],
call["actual_seq_q"],
call["q"],
call["weights"],
call["actual_seq_q_tensor"].tolist(),
call["actual_seq_q_cu_tensor"].tolist(),
)
for call in topk_calls
],
[
(0, 100, 2, [0.0, 1.0], [100.0, 101.0], [2], [0, 2]),
(0, 200, 1, [2.0], [102.0], [1], [0, 1]),
(1, 300, 1, [3.0], [103.0], [1], [0, 1]),
(1, 400, 3, [4.0, 5.0, 6.0], [104.0, 105.0, 106.0], [3], [0, 3]),
],
topk_calls[0]["cp_index"],
[(0, 0, 2), (0, 2, 3), (1, 0, 1), (1, 1, 4)],
)
self.assertEqual(topk_calls[0]["q"], [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0])
self.assertEqual(
topk_calls[0]["weights"],
[100.0, 101.0, 102.0, 103.0, 104.0, 105.0, 106.0],
)
self.assertIsNone(topk_calls[0]["actual_seq_q_tensor"])
self.assertIsNone(topk_calls[0]["actual_seq_q_cu_tensor"])
self.assertTrue(all(call["shared_index_buffer"] is materialized_index for call in topk_calls))
self.assertTrue(all(call["shared_block_tables"] is dense_pages for call in topk_calls))
self.assertTrue(all(call["current_index_kv"] is None for call in topk_calls))
self.assertEqual(
result.tolist(),
[[1, 1], [1, 1], [2, 2], [3, 3], [4, 4], [4, 4], [4, 4]],
[[1, 1], [2, 2], [3, 3], [4, 4], [5, 5], [6, 6], [7, 7]],
)
def test_indexer_in_seq_cp_pair_batch_materializes_partial_current_index_reuse_once(self):
@@ -1413,7 +1423,7 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
indexer = object.__new__(Indexer)
indexer.index_topk = 2
current_index_kv = (torch.tensor([1]), torch.tensor([2]))
current_index_kv = (torch.arange(7), torch.arange(7))
logical_pages = torch.tensor([[1, 2, 3], [4, 5, 6]], dtype=torch.int32)
materialized_index = torch.tensor([11], dtype=torch.int32)
dense_pages = torch.tensor([[1, 2, 3], [4, 5, 6]], dtype=torch.int32)
@@ -1424,6 +1434,9 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
def get_page_table_64(self):
return logical_pages
def get_page_table_1(self):
return torch.empty((2, 512), dtype=torch.int32)
def fake_materialize(
forward_batch,
layer_id,
@@ -1458,23 +1471,33 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
topk_calls.append(
{
"batch_idx": batch_idx,
"actual_seq_q": actual_seq_q,
"cp_index": cp_index,
"q_rows": int(q_fp8.shape[0]),
"current_index_kv": current_index_kv,
"shared_index_buffer": shared_index_buffer,
"shared_block_tables": shared_block_tables,
"actual_seq_q_cu_tensor": actual_seq_q_cu_tensor,
}
)
return torch.full((actual_seq_q, 2), len(topk_calls), dtype=torch.int32)
rows = int(q_fp8.shape[0])
return torch.arange(1, rows + 1, dtype=torch.int32).view(rows, 1).repeat(1, 2)
indexer._maybe_materialize_shared_index_buffer = fake_materialize
indexer._get_topk_ragged_with_cp = fake_get_topk
forward_batch = SimpleNamespace(
batch_size=2,
forward_mode=SimpleNamespace(
is_extend_without_speculative=lambda: True,
),
extend_prefix_lens_cpu=[64, 64],
extend_seq_lens_cpu=[3, 4],
seq_lens_cpu=torch.tensor([67, 68], dtype=torch.int64),
nsa_cp_metadata=NSAContextParallelMetadata(
batch_size=2,
request_kv_len_prev=[100, 300],
request_kv_len_next=[200, 400],
request_kv_len_prev=[2, 1],
request_kv_len_next=[3, 4],
request_actual_seq_q_prev=[2, 1],
request_actual_seq_q_next=[1, 3],
),
@@ -1492,7 +1515,13 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
self.assertEqual(len(materialize_calls), 1)
self.assertIs(materialize_calls[0]["logical_page_table"], logical_pages)
self.assertIs(materialize_calls[0]["current_index_kv"], current_index_kv)
self.assertEqual(len(topk_calls), 4)
self.assertEqual(len(topk_calls), 1)
self.assertEqual(topk_calls[0]["actual_seq_q"], 7)
self.assertEqual(
topk_calls[0]["cp_index"],
[(0, 0, 2), (0, 2, 3), (1, 0, 1), (1, 1, 4)],
)
self.assertEqual(topk_calls[0]["q_rows"], 7)
self.assertTrue(
all(call["current_index_kv"] is None for call in topk_calls)
)
@@ -1502,14 +1531,11 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
self.assertTrue(
all(call["shared_block_tables"] is dense_pages for call in topk_calls)
)
self.assertEqual([call["batch_idx"] for call in topk_calls], [0, 0, 1, 1])
self.assertEqual(
[call["actual_seq_q_cu_tensor"].tolist() for call in topk_calls],
[[0, 2], [0, 1], [0, 1], [0, 3]],
)
self.assertEqual([call["batch_idx"] for call in topk_calls], [0])
self.assertIsNone(topk_calls[0]["actual_seq_q_cu_tensor"])
self.assertEqual(
result.tolist(),
[[1, 1], [1, 1], [2, 2], [3, 3], [4, 4], [4, 4], [4, 4]],
[[1, 1], [2, 2], [3, 3], [4, 4], [5, 5], [6, 6], [7, 7]],
)
def test_indexer_in_seq_cp_pair_batch_reuses_current_index_without_materialize(self):
@@ -1519,7 +1545,7 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
indexer = object.__new__(Indexer)
indexer.index_topk = 2
current_index_kv = (torch.tensor([1]), torch.tensor([2]))
current_index_kv = (torch.arange(7), torch.arange(7))
topk_calls = []
class Mode:
@@ -1552,13 +1578,17 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
topk_calls.append(
{
"batch_idx": batch_idx,
"actual_seq_q": actual_seq_q,
"cp_index": cp_index,
"q_rows": int(q_fp8.shape[0]),
"current_index_kv": current_index_kv,
"shared_index_buffer": shared_index_buffer,
"shared_block_tables": shared_block_tables,
"actual_seq_q_cu_tensor": actual_seq_q_cu_tensor,
}
)
return torch.full((actual_seq_q, 2), len(topk_calls), dtype=torch.int32)
rows = int(q_fp8.shape[0])
return torch.arange(1, rows + 1, dtype=torch.int32).view(rows, 1).repeat(1, 2)
indexer._maybe_materialize_shared_index_buffer = fake_materialize
indexer._get_topk_ragged_with_cp = fake_get_topk
@@ -1587,7 +1617,7 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
current_index_kv=current_index_kv,
)
self.assertEqual(len(topk_calls), 4)
self.assertEqual(len(topk_calls), 1)
self.assertTrue(
all(call["current_index_kv"] is current_index_kv for call in topk_calls)
)
@@ -1597,16 +1627,305 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
self.assertTrue(
all(call["shared_block_tables"] is None for call in topk_calls)
)
self.assertEqual([call["batch_idx"] for call in topk_calls], [0, 0, 1, 1])
self.assertEqual(topk_calls[0]["batch_idx"], 0)
self.assertEqual(topk_calls[0]["actual_seq_q"], 7)
self.assertEqual(
[call["actual_seq_q_cu_tensor"].tolist() for call in topk_calls],
[[0, 2], [0, 1], [0, 1], [0, 3]],
topk_calls[0]["cp_index"],
[(0, 1, 3), (0, 2, 3), (1, 3, 4), (1, 1, 4)],
)
self.assertEqual(topk_calls[0]["q_rows"], 7)
self.assertIsNone(topk_calls[0]["actual_seq_q_cu_tensor"])
self.assertEqual(
result.tolist(),
[[1, 1], [1, 1], [2, 2], [3, 3], [4, 4], [4, 4], [4, 4]],
[[1, 1], [2, 2], [3, 3], [4, 4], [5, 5], [6, 6], [7, 7]],
)
def test_indexer_ragged_cp_index_current_batch_does_not_materialize(self):
import contextlib
import torch
from types import SimpleNamespace
from unittest.mock import patch
from sglang.srt.layers.attention.nsa import nsa_indexer
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
indexer = object.__new__(Indexer)
indexer.index_topk = 2
indexer._with_real_sm_count = lambda: contextlib.nullcontext()
deep_gemm_calls = []
def fake_logits(q_fp8, kv_fp8, weights, ks, ke, clean_logits=False):
deep_gemm_calls.append(
{
"q_rows": int(q_fp8.shape[0]),
"kv_rows": int(kv_fp8[0].shape[0]),
"weights_rows": int(weights.shape[0]),
"ks": ks.tolist(),
"ke": ke.tolist(),
}
)
return torch.zeros((int(q_fp8.shape[0]), 8), dtype=torch.float32)
class Metadata:
def get_page_table_64(self):
raise AssertionError("current cp_index path must not materialize index pages")
def topk_transform(self, logits, topk, **kwargs):
return (
torch.arange(1, int(logits.shape[0]) + 1, dtype=torch.int32)
.view(-1, 1)
.repeat(1, topk)
)
forward_batch = SimpleNamespace(
token_to_kv_pool=SimpleNamespace(page_size=64),
seq_lens_cpu=torch.tensor([3, 4], dtype=torch.int64),
extend_seq_lens_cpu=[3, 4],
)
q_fp8 = torch.empty((7, 1), dtype=torch.float32)
weights = torch.empty((7, 1, 1), dtype=torch.float32)
current_index_kv = (
torch.arange(7, dtype=torch.uint8).view(7, 1),
torch.arange(7, dtype=torch.float32).view(7, 1),
)
with patch.object(
nsa_indexer,
"deep_gemm",
SimpleNamespace(fp8_mqa_logits=fake_logits),
):
result = Indexer._get_topk_ragged_with_cp(
indexer,
forward_batch,
layer_id=7,
q_fp8=q_fp8,
weights=weights,
metadata=Metadata(),
kv_len=0,
actual_seq_q=7,
cp_index=[(0, 1, 3), (0, 2, 3), (1, 3, 4), (1, 1, 4)],
current_index_kv=current_index_kv,
)
self.assertEqual(result.tolist(), [[1, 1], [2, 2], [3, 3], [4, 4], [5, 5], [6, 6], [7, 7]])
self.assertEqual(len(deep_gemm_calls), 1)
self.assertEqual(deep_gemm_calls[0]["q_rows"], 7)
self.assertEqual(deep_gemm_calls[0]["weights_rows"], 7)
self.assertEqual(deep_gemm_calls[0]["kv_rows"], 14)
self.assertEqual(deep_gemm_calls[0]["ks"], [0, 0, 3, 6, 10, 10, 10])
self.assertEqual(deep_gemm_calls[0]["ke"], [2, 3, 6, 10, 12, 13, 14])
def test_indexer_ragged_cp_index_shared_batch_uses_tai_prepare_once(self):
import contextlib
import torch
from types import SimpleNamespace
from unittest.mock import patch
from sglang.srt.layers.attention.nsa import index_buf_accessor, nsa_indexer
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
indexer = object.__new__(Indexer)
indexer.index_topk = 2
indexer._with_real_sm_count = lambda: contextlib.nullcontext()
prepare_calls = []
deep_gemm_calls = []
def fake_prepare(**kwargs):
prepare_calls.append(kwargs)
total_kv_len = int(kwargs["total_kv_len"])
total_q_count = int(kwargs["total_q_count"])
return (
torch.zeros((total_kv_len, 1), dtype=torch.uint8),
torch.zeros((total_kv_len,), dtype=torch.float32),
torch.tensor([0, 0, 3, 6, 10, 10, 10], dtype=torch.int32),
torch.tensor([2, 3, 3, 4, 2, 3, 4], dtype=torch.int32),
)
def fake_logits(q_fp8, kv_fp8, weights, ks, ke, clean_logits=False):
deep_gemm_calls.append(
{
"q_rows": int(q_fp8.shape[0]),
"kv_rows": int(kv_fp8[0].shape[0]),
"weights_rows": int(weights.shape[0]),
"ks": ks.tolist(),
"ke": ke.tolist(),
}
)
return torch.zeros((int(q_fp8.shape[0]), 8), dtype=torch.float32)
class Metadata:
def topk_transform(self, logits, topk, **kwargs):
return (
torch.arange(1, int(logits.shape[0]) + 1, dtype=torch.int32)
.view(-1, 1)
.repeat(1, topk)
)
forward_batch = SimpleNamespace(
token_to_kv_pool=SimpleNamespace(page_size=64, index_head_dim=1),
seq_lens_cpu=torch.tensor([3, 4], dtype=torch.int64),
extend_seq_lens_cpu=[3, 4],
)
q_fp8 = torch.empty((7, 1), dtype=torch.float32)
weights = torch.empty((7, 1, 1), dtype=torch.float32)
shared_index_buffer = torch.zeros((8, 264), dtype=torch.uint8)
shared_block_tables = torch.arange(8, dtype=torch.int64).view(2, 4)
with patch.object(
nsa_indexer,
"try_tai_prepare_cp_mqa_index_batch",
side_effect=fake_prepare,
create=True,
), patch.object(
index_buf_accessor.GetK,
"execute",
side_effect=AssertionError("batched path must not call per-segment GetK"),
), patch.object(
index_buf_accessor.GetS,
"execute",
side_effect=AssertionError("batched path must not call per-segment GetS"),
), patch.object(
nsa_indexer,
"deep_gemm",
SimpleNamespace(fp8_mqa_logits=fake_logits),
):
result = Indexer._get_topk_ragged_with_cp(
indexer,
forward_batch,
layer_id=7,
q_fp8=q_fp8,
weights=weights,
metadata=Metadata(),
kv_len=0,
actual_seq_q=7,
cp_index=[(0, 1, 3), (0, 2, 3), (1, 3, 4), (1, 1, 4)],
shared_index_buffer=shared_index_buffer,
shared_block_tables=shared_block_tables,
)
self.assertEqual(result.tolist(), [[1, 1], [2, 2], [3, 3], [4, 4], [5, 5], [6, 6], [7, 7]])
self.assertEqual(len(prepare_calls), 1)
call = prepare_calls[0]
self.assertIs(call["index_buffer"], shared_index_buffer)
self.assertIs(call["block_tables"], shared_block_tables)
self.assertEqual(call["batch_indices"].tolist(), [0, 0, 1, 1])
self.assertEqual(call["kv_lens"].tolist(), [3, 3, 4, 4])
self.assertEqual(call["q_starts"].tolist(), [1, 2, 3, 1])
self.assertEqual(call["q_lens"].tolist(), [2, 1, 1, 3])
self.assertEqual(call["k_bases"].tolist(), [0, 3, 6, 10])
self.assertEqual(call["q_bases"].tolist(), [0, 2, 3, 4])
self.assertEqual(call["total_kv_len"], 14)
self.assertEqual(call["total_q_count"], 7)
self.assertEqual(call["max_kv_len"], 4)
self.assertEqual(call["max_q_len"], 3)
self.assertEqual(len(deep_gemm_calls), 1)
self.assertEqual(deep_gemm_calls[0]["q_rows"], 7)
self.assertEqual(deep_gemm_calls[0]["weights_rows"], 7)
self.assertEqual(deep_gemm_calls[0]["kv_rows"], 14)
self.assertEqual(deep_gemm_calls[0]["ks"], [0, 0, 3, 6, 10, 10, 10])
self.assertEqual(deep_gemm_calls[0]["ke"], [2, 3, 6, 10, 12, 13, 14])
def test_indexer_ragged_cp_index_current_batch_uses_tai_compact_once(self):
import contextlib
import torch
from types import SimpleNamespace
from unittest.mock import patch
from sglang.srt.layers.attention.nsa import nsa_indexer
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
indexer = object.__new__(Indexer)
indexer.index_topk = 2
indexer._with_real_sm_count = lambda: contextlib.nullcontext()
prepare_calls = []
deep_gemm_calls = []
def fake_prepare(**kwargs):
prepare_calls.append(kwargs)
total_kv_len = int(kwargs["total_kv_len"])
return (
torch.zeros((total_kv_len, 1), dtype=torch.uint8),
torch.zeros((total_kv_len,), dtype=torch.float32),
torch.tensor([0, 0, 3, 6, 10, 10, 10], dtype=torch.int32),
torch.tensor([2, 3, 3, 4, 2, 3, 4], dtype=torch.int32),
)
def fake_logits(q_fp8, kv_fp8, weights, ks, ke, clean_logits=False):
deep_gemm_calls.append(
{
"q_rows": int(q_fp8.shape[0]),
"kv_rows": int(kv_fp8[0].shape[0]),
"ks": ks.tolist(),
"ke": ke.tolist(),
}
)
return torch.zeros((int(q_fp8.shape[0]), 8), dtype=torch.float32)
class Metadata:
def get_page_table_64(self):
raise AssertionError("current cp_index path must not materialize index pages")
def topk_transform(self, logits, topk, **kwargs):
return (
torch.arange(1, int(logits.shape[0]) + 1, dtype=torch.int32)
.view(-1, 1)
.repeat(1, topk)
)
forward_batch = SimpleNamespace(
token_to_kv_pool=SimpleNamespace(page_size=64, index_head_dim=1),
seq_lens_cpu=torch.tensor([3, 4], dtype=torch.int64),
extend_seq_lens_cpu=[3, 4],
)
current_index_kv = (
torch.arange(7, dtype=torch.uint8).view(7, 1),
torch.arange(7, dtype=torch.float32).view(7, 1),
)
with patch.object(
nsa_indexer,
"try_tai_prepare_cp_mqa_current_index_batch",
side_effect=fake_prepare,
create=True,
), patch.object(
nsa_indexer,
"deep_gemm",
SimpleNamespace(fp8_mqa_logits=fake_logits),
):
result = Indexer._get_topk_ragged_with_cp(
indexer,
forward_batch,
layer_id=7,
q_fp8=torch.empty((7, 1), dtype=torch.float32),
weights=torch.empty((7, 1, 1), dtype=torch.float32),
metadata=Metadata(),
kv_len=0,
actual_seq_q=7,
cp_index=[(0, 1, 3), (0, 2, 3), (1, 3, 4), (1, 1, 4)],
current_index_kv=current_index_kv,
)
self.assertEqual(result.tolist(), [[1, 1], [2, 2], [3, 3], [4, 4], [5, 5], [6, 6], [7, 7]])
self.assertEqual(len(prepare_calls), 1)
call = prepare_calls[0]
self.assertIs(call["current_index_k"], current_index_kv[0])
self.assertIs(call["current_index_scale"], current_index_kv[1])
self.assertEqual(call["current_bases"].tolist(), [0, 0, 3, 3])
self.assertEqual(call["kv_lens"].tolist(), [3, 3, 4, 4])
self.assertEqual(call["q_starts"].tolist(), [1, 2, 3, 1])
self.assertEqual(call["q_lens"].tolist(), [2, 1, 1, 3])
self.assertEqual(call["k_bases"].tolist(), [0, 3, 6, 10])
self.assertEqual(call["q_bases"].tolist(), [0, 2, 3, 4])
self.assertEqual(call["total_kv_len"], 14)
self.assertEqual(call["total_q_count"], 7)
self.assertEqual(len(deep_gemm_calls), 1)
self.assertEqual(deep_gemm_calls[0]["kv_rows"], 14)
self.assertEqual(deep_gemm_calls[0]["ks"], [0, 0, 3, 6, 10, 10, 10])
self.assertEqual(deep_gemm_calls[0]["ke"], [2, 3, 6, 10, 12, 13, 14])
def test_indexer_in_seq_cp_pair_skips_materialize_when_current_index_reused(self):
import torch
@@ -1191,6 +1191,55 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertEqual(current_mask.tolist(), [[False, True, True, False, False]])
self.assertEqual(mixed_locs.tolist(), [[4, 12, 13, -1, -1]])
def test_fill_current_index_page_slots_uses_tai_kernel_when_available(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
class FakeKernels:
calls = []
@staticmethod
def fill_current_index_page_slots(*args, **kwargs):
FakeKernels.calls.append((args, kwargs))
dense_page_buffer = args[0]
dense_page_buffer[2, 0] = 99
return dense_page_buffer
dense_page_buffer = torch.zeros((3, 32), dtype=torch.uint8)
current_k = torch.ones((2, 4), dtype=torch.uint8)
current_scale = torch.ones((2, 1), dtype=torch.float32)
current_locs = torch.tensor([8, 9], dtype=torch.int64)
page_inverse = torch.tensor([0, -1, 2], dtype=torch.int64)
with patch.object(
runtime,
"_tai_materialize_runtime_enabled",
return_value=True,
), patch.object(
runtime,
"_load_tai_materialize_kernels",
return_value=FakeKernels,
):
result = runtime.fill_current_index_page_slots(
dense_page_buffer=dense_page_buffer,
current_index_k=current_k,
current_index_scale=current_scale,
current_locs=current_locs,
page_inverse=page_inverse,
page_size=4,
index_head_dim=4,
)
self.assertIs(result, dense_page_buffer)
self.assertEqual(int(result[2, 0]), 99)
self.assertEqual(len(FakeKernels.calls), 1)
args, kwargs = FakeKernels.calls[0]
self.assertIs(args[0], dense_page_buffer)
self.assertTrue(torch.equal(args[1], current_k))
self.assertTrue(torch.equal(args[2], current_scale))
self.assertTrue(torch.equal(args[3], current_locs))
self.assertTrue(torch.equal(args[4], page_inverse))
self.assertEqual(kwargs, {"page_size": 4, "index_head_dim": 4})
def test_tai_current_slot_fill_is_skipped_when_sparse_page_self_test_fails(self):
from sglang.srt.environ import envs
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime