Protect CP shared-KV cache-hit correctness under batched FP8 reuse

Cache-hit GSM8K regressions only appeared after the second pass reused request-specific suffix pages, so this change adds fail-fast transfer validation, masks stale rectangular page-table tails, and extends CUDA/unit coverage across FP8 CP shared-KV write, load, top-k, and materialization paths. The temporary ledger records eliminated hypotheses to prevent re-debugging the same L2 and persistent-cache paths.\n\nConstraint: CP shared KV stores physical pages but scheduler-visible semantics must remain valid-token/page-bounded.\nConstraint: bs>1 FP8 prefill must preserve existing CP shared-KV fast paths without silent fallback.\nRejected: Blame raw HiCache L2 load without tests | L2 KV and index backup/load/materialize roundtrips pass on remote CUDA.\nRejected: Disable current/partial reuse broadly | hides the cache-hit contract regression and costs performance.\nConfidence: medium\nScope-risk: moderate\nDirective: Do not weaken CP shared-KV fail-fast or rectangular-tail masking without rerunning second-pass cache-hit accuracy tests.\nTested: remote CUDA pytest for fused FP8 MLA store, fused persistent index store, L2-loaded FP8 KV materialize, L2-loaded index materialize, ragged top-k offset, TAI batched index MQA prepare.\nTested: local py_compile for touched test files and git diff --check.\nNot-tested: full second-pass GSM8K accuracy after these diagnostic tests; root cause remains under investigation.
This commit is contained in:
laoyao0822
2026-06-07 13:26:49 +08:00
parent b17976b60d
commit f75ffff8d9
10 changed files with 2645 additions and 1 deletions
@@ -4,8 +4,12 @@ import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import torch
import torch.distributed as dist
from sglang.srt.disaggregation.base import KVPoll
from sglang.srt.disaggregation.prefill import PrefillBootstrapQueue
from sglang.srt.disaggregation.utils import poll_and_all_reduce_attn_cp_tp_group
from sglang.srt.disaggregation.utils import ReqToMetadataIdxAllocator
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -25,6 +29,9 @@ class FakeSender:
if self.should_fail:
raise RuntimeError("boom")
def poll(self):
return KVPoll.Success
class TestPrefillBootstrapQueue(CustomTestCase):
def _make_req(self, rid, bootstrap_room, origin_input_ids, sender):
@@ -124,6 +131,70 @@ class TestPrefillBootstrapQueue(CustomTestCase):
self.assertEqual(skipped.disagg_kv_sender.init_calls, [])
self.assertEqual(checked.disagg_kv_sender.init_calls, [(3, 0)])
def test_poll_consensus_debug_fails_before_shape_mismatch_hang(self):
reduce_calls = []
def fake_all_reduce(tensor, op, group):
reduce_calls.append((tensor.dtype, int(tensor.numel()), op, group))
# The new debug guard uses an int64 [queue_len, queue_hash] scalar
# vector before building the uint8 per-request poll tensor. Simulate
# another rank having a different queue length and assert that we
# fail fast before reaching the old variable-length uint8 all_reduce.
if tensor.dtype == torch.int64 and int(tensor.numel()) == 2:
if op == dist.ReduceOp.MIN:
tensor[0] = 1
elif op == dist.ReduceOp.MAX:
tensor[0] = 2
return
raise AssertionError(
"poll tensor all_reduce should not run after queue mismatch"
)
with (
patch(
"sglang.srt.disaggregation.utils.envs.SGLANG_CP_SHARED_KV_BS_GT1_DEBUG.get",
return_value=True,
),
patch("sglang.srt.disaggregation.utils.dist.all_reduce", fake_all_reduce),
):
with self.assertRaisesRegex(RuntimeError, "poll_queue.*inflight"):
poll_and_all_reduce_attn_cp_tp_group(
[FakeSender(), FakeSender()],
MagicMock(name="cp_group"),
MagicMock(name="tp_group"),
debug_label="inflight",
debug_ids=["rid-a", "rid-b"],
)
self.assertTrue(reduce_calls)
def test_poll_consensus_debug_disabled_preserves_old_collective_shape(self):
reduce_calls = []
def fake_all_reduce(tensor, op, group):
reduce_calls.append((tensor.dtype, int(tensor.numel()), op, group))
with (
patch(
"sglang.srt.disaggregation.utils.envs.SGLANG_CP_SHARED_KV_BS_GT1_DEBUG.get",
return_value=False,
),
patch("sglang.srt.disaggregation.utils.dist.all_reduce", fake_all_reduce),
):
polls = poll_and_all_reduce_attn_cp_tp_group(
[FakeSender(), FakeSender()],
MagicMock(name="cp_group"),
MagicMock(name="tp_group"),
debug_label="inflight",
debug_ids=["rid-a", "rid-b"],
)
self.assertEqual(polls, [KVPoll.Success, KVPoll.Success])
self.assertEqual(
[(dtype, size) for dtype, size, _op, _group in reduce_calls],
[(torch.uint8, 2), (torch.uint8, 2)],
)
if __name__ == "__main__":
unittest.main()
@@ -14,6 +14,37 @@ register_cpu_ci(est_time=1, suite="stage-a-test-cpu")
class TestNSATopkTransform(unittest.TestCase):
@unittest.skipIf(not torch.cuda.is_available(), "CUDA is required")
def test_ragged_topk_transform_offsets_are_request_relative_with_row_starts(self):
from sgl_kernel import fast_topk_transform_ragged_fused
# CP bs>1 compacts each request segment into a temporary K buffer.
# The selected compact column must be translated back to the normal
# request-concatenated ragged layout as:
# request_base + (selected_compact_col - row_start)
# not request_base + selected_compact_col.
columns = 5000
logits = torch.zeros((2, columns), device="cuda", dtype=torch.float32)
logits[0, 2999] = 10.0
logits[1, 1000 + 2999] = 10.0
lengths = torch.tensor([3000, 3000], device="cuda", dtype=torch.int32)
row_starts = torch.tensor([0, 1000], device="cuda", dtype=torch.int32)
request_offsets = torch.tensor(
[0, 100000], device="cuda", dtype=torch.int32
)
out = fast_topk_transform_ragged_fused(
score=logits,
lengths=lengths,
topk_indices_offset=request_offsets,
topk=2048,
row_starts=row_starts,
)
torch.cuda.synchronize()
self.assertEqual(int(out[0, 0].item()), 2999)
self.assertEqual(int(out[1, 0].item()), 100000 + 2999)
def test_paged_topk_transform_raises_when_fused_output_is_not_from_page_table(self):
page_table = torch.tensor(
[
@@ -1316,6 +1316,31 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
self.assertIsNone(result)
def test_fp8_ragged_mla_defers_cp_materialize_to_flattened_path(self):
from pathlib import Path
source = (
Path(__file__).resolve().parents[4]
/ "python/sglang/srt/layers/attention/nsa_backend.py"
).read_text()
paged_gate = (
"if (\n"
" forward_batch.uses_cp_shared_kv\n"
" and topk_transform_method == TopkTransformMethod.PAGED\n"
" ):"
)
self.assertIn(
paged_gate,
source,
"RAGGED topk indices are flattened request/KV coordinates, not raw KV logical locs; "
"generic CP MLA materialize must stay gated to PAGED topk.",
)
ragged_start = source.index(" if topk_transform_method == TopkTransformMethod.RAGGED:")
ragged_end = source.index(" attn_output = self._forward_flashmla_sparse", ragged_start)
ragged_source = source[ragged_start:ragged_end]
self.assertIn("page_table_1_flattened", ragged_source)
self.assertIn("materialize_shared_token_kv_buffer", ragged_source)
@unittest.skipIf(not torch.cuda.is_available(), "CUDA is required")
def test_tai_current_slot_fill_sparse_page_self_test_passes_on_installed_kernel(
self,
@@ -1474,6 +1499,47 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
[(0, 2), (4, 5)],
)
def test_valid_page_mask_prevents_stale_rectangular_tail_remap(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
page_size = 4
layout = CpSharedKVLayout(page_size=page_size, cp_size=1, cp_rank=0)
kv_cache = torch.zeros((64, 1, 1), dtype=torch.float32)
# Row 0 is valid for 3 pages. The 4th column simulates stale req_to_token
# rectangular tail that duplicates logical page 1. Without masking,
# build_slot_page_inverse scatter_ maps page 1 to that stale slot.
raw_logical_pages = torch.tensor(
[
[1, 2, 5, 1],
[9, 10, 0, 0],
],
dtype=torch.int64,
)
valid_logical_pages = runtime.mask_batch_logical_pages_to_valid_lengths(
raw_logical_pages,
seq_lens_cpu=[12, 8],
page_size=page_size,
)
self.assertEqual(valid_logical_pages.tolist(), [[1, 2, 5, 0], [9, 10, 0, 0]])
slot_remap = runtime.build_shared_token_kv_slot_remap(
kv_cache=kv_cache,
logical_locs=torch.tensor([4], dtype=torch.int64),
remap_logical_pages=valid_logical_pages,
layout=layout,
page_size=page_size,
)
dense_loc = runtime.remap_logical_locs_to_slot_dense_locs(
torch.tensor([4], dtype=torch.int64),
page_inverse=slot_remap.page_inverse,
page_size=page_size,
)
self.assertEqual(dense_loc.tolist(), [4])
def test_materialize_batch_prefix_span_and_reuse_current_kv_page_slots(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
@@ -3157,6 +3223,747 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
rtol=0,
)
@unittest.skipIf(not torch.cuda.is_available(), "CUDA required")
def test_fp8_fused_mla_store_matches_fallback_for_bs5_compute_padding_rows(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
from sglang.srt.layers.attention.nsa.quant_k_cache import (
quantize_k_cache_separate,
)
from sglang.srt.layers.attention.nsa.utils import (
NSAContextParallelMetadata,
build_batch_page_aligned_in_seq_split_plan,
get_cp_shared_kv_local_out_cache_loc,
select_cp_local_valid_rows_for_cache_write,
)
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
build_in_seq_page_compute_owners,
)
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
from sglang.srt.mem_cache.utils import set_mla_kv_buffer_triton
try:
runtime._load_tai_fused_mla_store_kernel()
except Exception as exc:
self.skipTest(f"TAI fused MLA store kernel unavailable: {exc}")
page_size = 64
cp_size = 8
prefix_lens = [640, 640, 640, 640, 640]
extend_lens = [95, 81, 140, 66, 86]
def build_valid_locs():
chunks = []
logical_page_cursor = 0
for req_id, (extend_len, prefix_len) in enumerate(
zip(extend_lens, prefix_lens)
):
owners = build_in_seq_page_compute_owners(
extend_len=int(extend_len),
extend_prefix_len=int(prefix_len),
page_size=page_size,
cp_size=cp_size,
)
self.assertIsNotNone(owners)
remaining = int(extend_len)
req_chunks = []
for owner in owners:
logical_page = (
int(owner)
+ 1
+ cp_size * (req_id * (len(owners) + 4) + logical_page_cursor)
)
page_locs = torch.arange(
logical_page * page_size,
(logical_page + 1) * page_size,
dtype=torch.int64,
)
valid_rows = min(page_size, max(remaining, 0))
if valid_rows > 0:
req_chunks.append(page_locs[:valid_rows])
remaining -= valid_rows
logical_page_cursor += 1
self.assertEqual(remaining, 0)
chunks.append(torch.cat(req_chunks, dim=0))
return torch.cat(chunks, dim=0)
out_cache_loc = build_valid_locs().to(device="cuda")
for rank in range(cp_size):
layout = CpSharedKVLayout(
page_size=page_size,
cp_size=cp_size,
cp_rank=rank,
)
plan = build_batch_page_aligned_in_seq_split_plan(
extend_lens=extend_lens,
prefix_lens=prefix_lens,
page_size=page_size,
cp_size=cp_size,
cp_rank=rank,
)
forward_batch = SimpleNamespace(
uses_cp_shared_kv=True,
cp_shared_kv_layout=layout,
nsa_cp_metadata=NSAContextParallelMetadata(
batch_size=len(extend_lens),
batch_plan=plan,
page_aligned=True,
page_size=page_size,
extend_prefix_len=prefix_lens[0],
),
out_cache_loc=out_cache_loc.clone(),
token_to_kv_pool=SimpleNamespace(page_size=page_size),
)
with patch(
"sglang.srt.layers.attention.nsa.utils.is_nsa_prefill_cp_round_robin_split",
return_value=False,
):
logical_locs = get_cp_shared_kv_local_out_cache_loc(forward_batch)
if logical_locs.numel() == 0:
continue
local_compute_rows = sum(
int(x) for x in plan.request_compute_rank_local_tokens
)
torch.manual_seed(20260607 + rank)
k_nope = (
torch.randn(
(local_compute_rows, 1, 512),
device="cuda",
dtype=torch.bfloat16,
)
* 2.0
) + 0.25
latent_cache = torch.randn(
(local_compute_rows, 576), device="cuda", dtype=torch.bfloat16
)
k_rope = latent_cache[:, 512:].unsqueeze(1)
self.assertFalse(k_rope.is_contiguous())
valid_k_nope = select_cp_local_valid_rows_for_cache_write(
forward_batch, k_nope
)
valid_k_rope = select_cp_local_valid_rows_for_cache_write(
forward_batch, k_rope
)
self.assertEqual(int(valid_k_nope.shape[0]), int(logical_locs.numel()))
physical_locs = layout.logical_locs_to_physical(logical_locs)
capacity_tokens = int(physical_locs.max().item()) + page_size + 1
expected = torch.zeros(
(capacity_tokens, 1, 656), dtype=torch.uint8, device="cuda"
)
nope_part, rope_part = quantize_k_cache_separate(
valid_k_nope, valid_k_rope
)
set_mla_kv_buffer_triton(expected, physical_locs, nope_part, rope_part)
class FakePool:
def __init__(self):
self.nsa_kv_cache_store_fp8 = True
self.page_size = page_size
self.start_layer = 0
self.kv_buffer = [
torch.zeros(
(capacity_tokens, 1, 656),
dtype=torch.uint8,
device="cuda",
)
]
pool = FakePool()
layer = SimpleNamespace(layer_id=0)
with patch.object(
runtime, "cp_shared_kv_tai_fused_mla_store_enabled", return_value=True
), patch.object(runtime, "cp_shared_kv_debug_enabled", return_value=False):
used = runtime.try_tai_fused_mla_store(
token_to_kv_pool=pool,
layer=layer,
layout=layout,
logical_locs=logical_locs,
k_nope=valid_k_nope,
k_rope=valid_k_rope,
)
torch.cuda.synchronize()
self.assertTrue(used)
torch.testing.assert_close(pool.kv_buffer[0], expected, atol=0, rtol=0)
@unittest.skipIf(not torch.cuda.is_available(), "CUDA required")
def test_fp8_index_fused_store_persistent_pages_survive_bs5_materialize(self):
from sglang.jit_kernel.fused_store_index_cache import (
can_use_nsa_fused_store,
fused_store_index_k_cache,
)
from sglang.srt.layers.attention.nsa import (
cp_shared_kv_runtime as runtime,
index_buf_accessor,
)
from sglang.srt.layers.attention.nsa.triton_kernel import act_quant
from sglang.srt.layers.attention.nsa.utils import (
NSAContextParallelMetadata,
build_batch_page_aligned_in_seq_split_plan,
get_cp_shared_kv_local_out_cache_loc,
select_cp_local_valid_rows_for_cache_write,
)
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
build_in_seq_page_compute_owners,
)
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
if not hasattr(torch, "float8_e4m3fn"):
self.skipTest("torch.float8_e4m3fn unavailable")
if not can_use_nsa_fused_store(torch.bfloat16, torch.int64, 64):
self.skipTest("NSA fused index store JIT unavailable")
page_size = 64
cp_size = 8
batch_size = 5
index_head_dim = 128
page_bytes = page_size * (index_head_dim + 4)
prefix_lens = [640, 640, 640, 640, 640]
extend_lens = [95, 81, 140, 66, 86]
request_pages: list[list[int]] = []
request_valid_locs: list[torch.Tensor] = []
all_valid_locs: list[torch.Tensor] = []
next_generation = 32
for req_id, (prefix_len, extend_len) in enumerate(
zip(prefix_lens, extend_lens)
):
owners = build_in_seq_page_compute_owners(
extend_len=extend_len,
extend_prefix_len=prefix_len,
page_size=page_size,
cp_size=cp_size,
)
self.assertIsNotNone(owners)
req_pages = []
req_locs = []
remaining = int(extend_len)
for owner in owners:
logical_page = int(owner) + 1 + cp_size * next_generation
next_generation += 1
req_pages.append(logical_page)
valid_rows = min(page_size, remaining)
page_locs = torch.arange(
logical_page * page_size,
logical_page * page_size + valid_rows,
dtype=torch.int64,
)
req_locs.append(page_locs)
remaining -= valid_rows
self.assertEqual(remaining, 0, f"request {req_id} loc construction")
request_pages.append(req_pages)
request_valid_locs.append(torch.cat(req_locs, dim=0))
all_valid_locs.append(request_valid_locs[-1])
max_request_pages = max(len(pages) for pages in request_pages)
logical_pages = torch.zeros(
(batch_size, max_request_pages), dtype=torch.int64, device="cuda"
)
for req_id, pages in enumerate(request_pages):
logical_pages[req_id, : len(pages)] = torch.tensor(
pages, dtype=torch.int64, device="cuda"
)
out_cache_loc = torch.cat(all_valid_locs, dim=0).to(device="cuda")
max_logical_page = max(max(pages) for pages in request_pages)
physical_page_capacity = (max_logical_page - 1) // cp_size + 3
rank_dense_buffers = []
expected_by_loc: dict[int, torch.Tensor] = {}
for rank in range(cp_size):
layout = CpSharedKVLayout(
page_size=page_size,
cp_size=cp_size,
cp_rank=rank,
)
plan = build_batch_page_aligned_in_seq_split_plan(
extend_lens=extend_lens,
prefix_lens=prefix_lens,
page_size=page_size,
cp_size=cp_size,
cp_rank=rank,
)
forward_batch = SimpleNamespace(
uses_cp_shared_kv=True,
cp_shared_kv_layout=layout,
nsa_cp_metadata=NSAContextParallelMetadata(
batch_size=batch_size,
batch_plan=plan,
page_aligned=True,
page_size=page_size,
extend_prefix_len=prefix_lens[0],
),
out_cache_loc=out_cache_loc.clone(),
token_to_kv_pool=SimpleNamespace(page_size=page_size),
)
with patch(
"sglang.srt.layers.attention.nsa.utils.is_nsa_prefill_cp_round_robin_split",
return_value=False,
):
logical_locs = get_cp_shared_kv_local_out_cache_loc(forward_batch)
local_rows = sum(int(x) for x in plan.request_compute_rank_local_tokens)
torch.manual_seed(20260607 + rank)
local_key = (
torch.randn(
(local_rows, index_head_dim),
dtype=torch.bfloat16,
device="cuda",
)
* 2.0
+ 0.125
).contiguous()
valid_key = select_cp_local_valid_rows_for_cache_write(
forward_batch,
local_key,
).contiguous()
self.assertEqual(int(valid_key.shape[0]), int(logical_locs.numel()))
page_buffer = torch.zeros(
(physical_page_capacity, page_bytes),
dtype=torch.uint8,
device="cuda",
)
if logical_locs.numel() > 0:
physical_locs = layout.logical_locs_to_physical(logical_locs)
fused_store_index_k_cache(
valid_key,
page_buffer,
physical_locs.contiguous(),
page_size=page_size,
)
for loc, row in zip(logical_locs.cpu().tolist(), valid_key):
expected_by_loc[int(loc)] = row.detach()
slot_remap = runtime.build_shared_paged_buffer_slot_remap(
page_buffer,
logical_pages,
layout,
)
with patch.object(
runtime, "_all_reduce_materialized_buffer", _identity_all_reduce
), patch.object(
runtime, "_try_tai_materialize_shared_pages", return_value=None
), patch.object(
runtime, "_try_tai_ipc_materialize_paged_buffer_page_slots_into",
return_value=False,
):
dense_page_buffer, dense_pages = runtime.materialize_shared_paged_buffer(
page_buffer,
logical_pages,
layout,
slot_remap=slot_remap,
)
rank_dense_buffers.append(dense_page_buffer.to(torch.int16))
merged_dense_buffer = (
torch.stack(rank_dense_buffers, dim=0).sum(dim=0).to(torch.uint8)
)
class FakePool:
page_size = 64
index_head_dim = 128
quant_block_size = 128
device = "cuda"
# Cross-check fallback and fused quantization on the same valid rows so
# this test fails on row/order/page corruption, not on an unrelated
# scale-format policy difference.
for req_id, (extend_len, req_locs) in enumerate(
zip(extend_lens, request_valid_locs)
):
req_page_count = len(request_pages[req_id])
page_indices = dense_pages[req_id, :req_page_count].contiguous()
k_u8 = index_buf_accessor.GetK.execute(
FakePool,
merged_dense_buffer,
seq_len=int(extend_len),
page_indices=page_indices,
)
scale_u8 = index_buf_accessor.GetS.execute(
FakePool,
merged_dense_buffer,
seq_len=int(extend_len),
page_indices=page_indices,
)
stored_deq = (
k_u8.contiguous().view(torch.float8_e4m3fn).float()
* scale_u8.contiguous().view(torch.float32).view(-1, 1)
)
expected_key = torch.stack(
[expected_by_loc[int(loc)] for loc in req_locs.tolist()],
dim=0,
).to(device="cuda")
ref_fp8, ref_scale = act_quant(
expected_key.contiguous(),
block_size=128,
scale_fmt="ue8m0",
)
ref_deq = ref_fp8.float() * ref_scale.float().view(-1, 1)
torch.testing.assert_close(
stored_deq,
expected_key.float(),
rtol=0.20,
atol=0.65,
)
torch.testing.assert_close(
stored_deq,
ref_deq,
rtol=0.35,
atol=0.90,
)
def test_cp8_index_partial_current_compose_matches_rank_merged_reference_for_bs5(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
build_in_seq_page_compute_owners,
)
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
page_size = 64
cp_size = 8
batch_size = 5
prefix_len = 640
prefix_pages = prefix_len // page_size
extend_lens = [95, 81, 140, 66, 86]
index_head_dim = 8
scale_bytes = 4
page_bytes = page_size * index_head_dim + page_size * scale_bytes
pages_per_request = prefix_pages + 3
shared_prefix_pages = list(range(1, prefix_pages + 1))
next_owner_generation = 32
logical_rows = []
current_locs_by_rank = [[] for _ in range(cp_size)]
current_k_by_rank = [[] for _ in range(cp_size)]
current_scale_by_rank = [[] for _ in range(cp_size)]
expected = torch.zeros(
(batch_size * pages_per_request + 1, page_bytes), dtype=torch.uint8
)
scale_offset = page_size * index_head_dim
def fill_index_row(buffer, dense_page, offset, k_row, scale_row):
buffer[dense_page, offset * index_head_dim : (offset + 1) * index_head_dim] = k_row
scale = scale_row.reshape(1, -1).view(torch.uint8).reshape(-1)
start = scale_offset + offset * scale_bytes
buffer[dense_page, start : start + scale_bytes] = scale
current_pages_by_req = []
for req_id, extend_len in enumerate(extend_lens):
owners = build_in_seq_page_compute_owners(
extend_len=extend_len,
extend_prefix_len=prefix_len,
page_size=page_size,
cp_size=cp_size,
)
current_pages = []
for owner in owners:
logical_page = int(owner) + 1 + cp_size * next_owner_generation
next_owner_generation += 1
current_pages.append(logical_page)
current_pages_by_req.append(current_pages)
logical_rows.append(
shared_prefix_pages
+ current_pages
+ [0] * (pages_per_request - prefix_pages - len(current_pages))
)
logical_pages = torch.tensor(logical_rows, dtype=torch.int64)
max_logical_page = int(logical_pages.max().item())
physical_capacity = (max_logical_page - 1) // cp_size + 2
prefix_page_bytes = {}
for logical_page in shared_prefix_pages:
row = ((torch.arange(page_bytes, dtype=torch.int64) + logical_page * 17) % 251).to(torch.uint8)
prefix_page_bytes[logical_page] = row
for req_id, current_pages in enumerate(current_pages_by_req):
remaining = int(extend_lens[req_id])
for page_offset, logical_page in enumerate(current_pages):
dense_page = req_id * pages_per_request + prefix_pages + page_offset + 1
valid_rows = min(page_size, remaining)
owner = int((logical_page - 1) % cp_size)
for offset in range(valid_rows):
logical_loc = logical_page * page_size + offset
k_row = ((torch.arange(index_head_dim, dtype=torch.int64) + req_id * 31 + page_offset * 7 + offset) % 253).to(torch.uint8)
scale_row = torch.tensor(
[[req_id * 1000.0 + page_offset * 100.0 + offset + 0.5]],
dtype=torch.float32,
)
current_locs_by_rank[owner].append(logical_loc)
current_k_by_rank[owner].append(k_row)
current_scale_by_rank[owner].append(scale_row.reshape(1))
fill_index_row(expected, dense_page, offset, k_row, scale_row)
remaining -= valid_rows
self.assertEqual(remaining, 0)
for req_id in range(batch_size):
for prefix_slot, logical_page in enumerate(shared_prefix_pages):
dense_page = req_id * pages_per_request + prefix_slot + 1
expected[dense_page] = prefix_page_bytes[logical_page]
prefix_slot_spans = runtime.build_batch_prefix_slot_spans(
logical_pages=logical_pages,
prefix_lens_cpu=[prefix_len] * batch_size,
page_size=page_size,
)
current_slot_spans = runtime.build_batch_current_slot_spans(
logical_pages=logical_pages,
prefix_lens_cpu=[prefix_len] * batch_size,
extend_lens_cpu=extend_lens,
page_size=page_size,
)
rank_outputs = []
for rank in range(cp_size):
layout = CpSharedKVLayout(page_size=page_size, cp_size=cp_size, cp_rank=rank)
page_buffer = torch.zeros((physical_capacity, page_bytes), dtype=torch.uint8)
for logical_page, row in prefix_page_bytes.items():
if int((logical_page - 1) % cp_size) == rank:
physical_page = int((logical_page - 1) // cp_size) + 1
page_buffer[physical_page] = row
slot_remap = runtime.build_shared_paged_buffer_slot_remap(
page_buffer,
logical_pages,
layout,
)
current_locs = torch.tensor(current_locs_by_rank[rank], dtype=torch.int64)
if current_k_by_rank[rank]:
current_k = torch.stack(current_k_by_rank[rank], dim=0)
current_scale = torch.stack(current_scale_by_rank[rank], dim=0)
else:
current_k = torch.empty((0, index_head_dim), dtype=torch.uint8)
current_scale = torch.empty((0, 1), dtype=torch.float32)
with patch.object(runtime, "_all_reduce_materialized_buffer_range", _identity_all_reduce):
dense_page_buffer, dense_pages = runtime.materialize_prefix_and_reuse_current_index_page_slots(
page_buffer=page_buffer,
current_index_k=current_k,
current_index_scale=current_scale,
current_locs=current_locs,
slot_remap=slot_remap,
layout=layout,
page_size=page_size,
index_head_dim=index_head_dim,
prefix_pages=0,
prefix_slot_spans=prefix_slot_spans,
current_slot_spans=current_slot_spans,
layer_id=0,
)
expected_dense_pages = logical_pages.clone()
flat_positive = expected_dense_pages.reshape(-1) > 0
expected_dense_pages.reshape(-1)[flat_positive] = torch.arange(
1,
int(expected_dense_pages.numel()) + 1,
dtype=expected_dense_pages.dtype,
)[flat_positive]
self.assertEqual(dense_pages.tolist(), expected_dense_pages.tolist())
rank_outputs.append(dense_page_buffer.to(torch.int16))
merged = torch.stack(rank_outputs, dim=0).sum(dim=0).to(torch.uint8)
torch.testing.assert_close(merged, expected, atol=0, rtol=0)
def test_cp8_kv_partial_current_keeps_remote_current_valid_locs_after_reduce(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
page_size = 4
cp_size = 8
# Two current pages: logical page 33 is rank0-owned, logical page 42 is
# rank1-owned. After the current slot all-reduce both valid pages are
# present on every rank, so rank0 must not mask the rank1-owned valid loc.
logical_pages = torch.tensor([[1, 33, 42]], dtype=torch.int64)
remote_current_loc = 42 * page_size + 1
local_current_locs = torch.tensor(
[33 * page_size + 0, 33 * page_size + 1], dtype=torch.int64
)
logical_locs = torch.tensor(
[1 * page_size + 0, 33 * page_size + 1, remote_current_loc],
dtype=torch.int64,
)
layout = CpSharedKVLayout(page_size=page_size, cp_size=cp_size, cp_rank=0)
physical_capacity = 8
kv_cache = torch.zeros((physical_capacity * page_size, 2), dtype=torch.float32)
kv_cache[layout.logical_locs_to_physical(torch.tensor([page_size], dtype=torch.int64))[0]] = torch.tensor([7.0, 8.0])
slot_remap = runtime.build_shared_token_kv_slot_remap(
kv_cache,
logical_locs,
logical_pages,
layout,
page_size,
)
current_kv_cache = torch.tensor([[10.0, 11.0], [12.0, 13.0]])
with patch.object(runtime, "_all_reduce_materialized_buffer_range", _identity_all_reduce):
_mixed_cache, mixed_locs = runtime.materialize_prefix_and_reuse_current_kv_page_slots(
kv_cache=kv_cache,
logical_locs=logical_locs,
current_kv_cache=current_kv_cache,
current_locs=local_current_locs,
slot_remap=slot_remap,
layout=layout,
page_size=page_size,
prefix_pages=1,
current_slot_spans=[(1, 3)],
layer_id=0,
)
self.assertGreaterEqual(
int(mixed_locs[-1].item()),
0,
"remote-rank valid current loc must remain visible after current slot all-reduce",
)
@unittest.skipIf(not torch.cuda.is_available(), "CUDA required")
def test_tai_batched_index_mqa_prepare_matches_getk_gets_reference_gsm8k_bs5(self):
from types import SimpleNamespace
from sglang.srt.layers.attention.nsa import (
cp_shared_kv_runtime as runtime,
index_buf_accessor,
)
page_size = 64
index_head_dim = 128
scale_bytes = 4
row_bytes = index_head_dim + scale_bytes
page_bytes = page_size * row_bytes
batch_size = 5
seq_lens = [704, 704, 704, 704, 704]
q_starts = [640, 640, 640, 640, 640]
q_lens = [64, 64, 64, 64, 64]
pages_per_request = 11
num_pages = batch_size * pages_per_request + 1
index_buffer = torch.empty(
(num_pages, page_bytes), dtype=torch.uint8, device="cuda"
)
token_offsets = torch.arange(
page_size, dtype=torch.float32, device="cuda"
).view(page_size, 1)
byte_offsets = torch.arange(
index_head_dim, dtype=torch.int64, device="cuda"
).view(1, index_head_dim)
for page_id in range(num_pages):
page_view = index_buffer[page_id].view(page_size, row_bytes)
k_bytes = (
byte_offsets
+ page_id * 17
+ torch.arange(page_size, dtype=torch.int64, device="cuda").view(
page_size, 1
)
).remainder_(251).to(torch.uint8)
scales = token_offsets + float(page_id * 1000) + 0.25
page_view[:, :index_head_dim] = k_bytes
page_view[:, index_head_dim : index_head_dim + scale_bytes] = (
scales.contiguous().view(torch.uint8).view(page_size, scale_bytes)
)
block_tables = torch.tensor(
[
[1 + req_id * pages_per_request + page for page in range(pages_per_request)]
for req_id in range(batch_size)
],
dtype=torch.int64,
device="cuda",
)
batch_indices = torch.arange(batch_size, dtype=torch.int32, device="cuda")
kv_lens = torch.tensor(seq_lens, dtype=torch.int32, device="cuda")
q_starts_tensor = torch.tensor(q_starts, dtype=torch.int32, device="cuda")
q_lens_tensor = torch.tensor(q_lens, dtype=torch.int32, device="cuda")
k_bases = torch.tensor(
[sum(seq_lens[:i]) for i in range(batch_size)],
dtype=torch.int32,
device="cuda",
)
q_bases = torch.tensor(
[sum(q_lens[:i]) for i in range(batch_size)],
dtype=torch.int32,
device="cuda",
)
with patch.object(
runtime, "cp_shared_kv_tai_index_mqa_prepare_enabled", return_value=True
):
prepared = runtime.try_tai_prepare_cp_mqa_index_batch(
index_buffer=index_buffer,
block_tables=block_tables,
batch_indices=batch_indices,
kv_lens=kv_lens,
q_starts=q_starts_tensor,
q_lens=q_lens_tensor,
k_bases=k_bases,
q_bases=q_bases,
total_kv_len=sum(seq_lens),
total_q_count=sum(q_lens),
max_kv_len=max(seq_lens),
max_q_len=max(q_lens),
page_size=page_size,
index_head_dim=index_head_dim,
)
if prepared is None:
self.skipTest("TAI batched index MQA prepare kernel unavailable")
k_fp8_u8, k_scale, ks, ke_offset = prepared
fake_pool = SimpleNamespace(
page_size=page_size,
index_head_dim=index_head_dim,
quant_block_size=index_head_dim,
device="cuda",
)
ref_k = []
ref_scale = []
ref_ks = []
ref_ke_offset = []
for req_id, (seq_len, q_start, q_len) in enumerate(
zip(seq_lens, q_starts, q_lens)
):
ref_k.append(
index_buf_accessor.GetK.execute(
fake_pool,
index_buffer,
seq_len=seq_len,
page_indices=block_tables[req_id],
)
)
ref_scale.append(
index_buf_accessor.GetS.execute(
fake_pool,
index_buffer,
seq_len=seq_len,
page_indices=block_tables[req_id],
)
.view(torch.float32)
.squeeze(-1)
)
ref_ks.append(
torch.full(
(q_len,),
int(k_bases[req_id].item()),
dtype=torch.int32,
device="cuda",
)
)
ref_ke_offset.append(
torch.arange(
q_start + 1,
q_start + q_len + 1,
dtype=torch.int32,
device="cuda",
)
)
ref_k = torch.cat(ref_k, dim=0)
ref_scale = torch.cat(ref_scale, dim=0)
ref_ks = torch.cat(ref_ks, dim=0)
ref_ke_offset = torch.cat(ref_ke_offset, dim=0)
torch.testing.assert_close(k_fp8_u8, ref_k, atol=0, rtol=0)
torch.testing.assert_close(k_scale, ref_scale, atol=0, rtol=0)
torch.testing.assert_close(ks, ref_ks, atol=0, rtol=0)
torch.testing.assert_close(ke_offset, ref_ke_offset, atol=0, rtol=0)
def test_token_range_materialize_uses_tai_kernel_when_enabled(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
@@ -1,8 +1,19 @@
import unittest
from types import SimpleNamespace
from unittest.mock import patch
import torch
from sglang.srt.layers.attention.nsa.utils import (
NSAContextParallelMetadata,
build_batch_page_aligned_in_seq_split_plan,
get_cp_shared_kv_local_out_cache_loc,
select_cp_local_valid_rows_for_cache_write,
)
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
build_in_seq_page_compute_owners,
)
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool
from sglang.srt.mem_cache.memory_pool_host import (
ALLOC_MEMORY_FUNCS,
@@ -133,6 +144,815 @@ class TestNSAHiCacheTransfer(CustomTestCase):
].cpu()
self.assertTrue(torch.equal(got_kv, expected_kv))
@staticmethod
def _build_owned_logical_out_locs(
*,
extend_lens,
prefix_lens,
page_size: int,
cp_size: int,
dst_page_stride: int = 0,
) -> tuple[torch.Tensor, torch.Tensor]:
valid_chunks = []
padded_chunks = []
logical_page_cursor = 0
for req_id, (extend_len, prefix_len) in enumerate(
zip(extend_lens, prefix_lens)
):
owners = build_in_seq_page_compute_owners(
extend_len=int(extend_len),
extend_prefix_len=int(prefix_len),
page_size=page_size,
cp_size=cp_size,
)
if owners is None:
raise AssertionError("test shape must support owner planning")
remaining = int(extend_len)
req_valid_chunks = []
req_padded_chunks = []
for page_idx, owner in enumerate(owners):
logical_page = (
int(owner)
+ 1
+ cp_size
* (
dst_page_stride
+ req_id * (len(owners) + 4)
+ logical_page_cursor
)
)
page_locs = torch.arange(
logical_page * page_size,
(logical_page + 1) * page_size,
dtype=torch.int64,
)
valid_rows = min(page_size, max(remaining, 0))
if valid_rows > 0:
req_valid_chunks.append(page_locs[:valid_rows])
remaining -= valid_rows
req_padded_chunks.append(page_locs)
logical_page_cursor += 1
if remaining != 0:
raise AssertionError(
f"failed to cover request extend_len={extend_len}"
)
valid_chunks.append(torch.cat(req_valid_chunks, dim=0))
padded_chunks.append(torch.cat(req_padded_chunks, dim=0))
return torch.cat(valid_chunks, dim=0), torch.cat(padded_chunks, dim=0)
@staticmethod
def _index_row_bytes(row_values: torch.Tensor, page_size: int) -> torch.Tensor:
# index_k_with_scale_buffer stores one page as
# [token0(index+scale), token1(index+scale), ...].
return row_values.view(row_values.shape[0], page_size, -1)
def test_fp8_page_first_direct_bs5_valid_rows_survive_l2_roundtrip(self):
page_size = 64
cp_size = 8
layer_num = 2
# GSM8K-like second-pass shape: same long page-aligned prompt prefix,
# short per-request suffixes, and compute padding on every request.
prefix_lens = [640, 640, 640, 640, 640]
extend_lens = [95, 81, 140, 66, 86]
src_valid_locs, src_padded_locs = self._build_owned_logical_out_locs(
extend_lens=extend_lens,
prefix_lens=prefix_lens,
page_size=page_size,
cp_size=cp_size,
)
dst_valid_locs, dst_padded_locs = self._build_owned_logical_out_locs(
extend_lens=extend_lens,
prefix_lens=prefix_lens,
page_size=page_size,
cp_size=cp_size,
dst_page_stride=128,
)
max_physical_page = (
int(torch.div(dst_padded_locs.max(), page_size, rounding_mode="floor"))
// cp_size
+ 4
)
size = max_physical_page * page_size
device_pool = NSATokenToKVPool(
size=size,
page_size=page_size,
kv_lora_rank=512,
dtype=torch.float8_e4m3fn,
qk_rope_head_dim=64,
layer_num=layer_num,
device="cuda",
enable_memory_saver=False,
kv_cache_dim=656,
index_head_dim=128,
)
host_pool = NSATokenToKVPoolHost(
device_pool=device_pool,
host_to_device_ratio=2.0,
host_size=0,
page_size=page_size,
layout="page_first_direct",
pin_memory=True,
device="cpu",
)
for cp_rank in range(cp_size):
layout = CpSharedKVLayout(
page_size=page_size,
cp_size=cp_size,
cp_rank=cp_rank,
)
plan = build_batch_page_aligned_in_seq_split_plan(
extend_lens=extend_lens,
prefix_lens=prefix_lens,
page_size=page_size,
cp_size=cp_size,
cp_rank=cp_rank,
)
forward_batch = SimpleNamespace(
uses_cp_shared_kv=True,
cp_shared_kv_layout=layout,
nsa_cp_metadata=NSAContextParallelMetadata(
batch_size=len(extend_lens),
batch_plan=plan,
page_aligned=True,
page_size=page_size,
extend_prefix_len=prefix_lens[0],
),
out_cache_loc=src_valid_locs.clone(),
token_to_kv_pool=SimpleNamespace(page_size=page_size),
)
with patch(
"sglang.srt.layers.attention.nsa.utils.is_nsa_prefill_cp_round_robin_split",
return_value=False,
):
local_valid_logical_locs = get_cp_shared_kv_local_out_cache_loc(
forward_batch
)
local_valid_rows = int(local_valid_logical_locs.numel())
if local_valid_rows == 0:
continue
local_compute_rows = sum(
int(x) for x in plan.request_compute_rank_local_tokens
)
kv_rows = (
torch.arange(
local_compute_rows * 656,
device="cuda",
dtype=torch.int64,
)
.view(local_compute_rows, 1, 656)
.add_(1000 * cp_rank)
.remainder_(251)
.to(torch.uint8)
)
index_row_stride = device_pool.index_k_with_scale_buffer[0].shape[1] // page_size
index_rows = (
torch.arange(
local_compute_rows * index_row_stride,
device="cuda",
dtype=torch.int64,
)
.view(local_compute_rows, index_row_stride)
.add_(2000 * cp_rank)
.remainder_(253)
.to(torch.uint8)
)
valid_kv_rows = select_cp_local_valid_rows_for_cache_write(
forward_batch, kv_rows
)
valid_index_rows = select_cp_local_valid_rows_for_cache_write(
forward_batch, index_rows
)
self.assertEqual(int(valid_kv_rows.shape[0]), local_valid_rows)
self.assertEqual(int(valid_index_rows.shape[0]), local_valid_rows)
local_valid_physical_locs = layout.logical_locs_to_physical(
local_valid_logical_locs
).to(device="cuda")
src_owned_padded_locs = src_padded_locs[
layout.owned_by_this_rank(src_padded_locs)
]
dst_owned_padded_locs = dst_padded_locs[
layout.owned_by_this_rank(dst_padded_locs)
]
src_physical_padded_locs = layout.logical_locs_to_physical(
src_owned_padded_locs
)
dst_physical_padded_locs = layout.logical_locs_to_physical(
dst_owned_padded_locs
)
self.assertEqual(
int(src_physical_padded_locs.numel()),
int(dst_physical_padded_locs.numel()),
)
host_indices = torch.arange(
cp_rank * 4096,
cp_rank * 4096 + int(src_physical_padded_locs.numel()),
dtype=torch.int64,
)
# The real reservation backs up full owned pages. Poison both src
# and dst first so the assertion proves valid rows were written and
# copied, not accidentally preserved.
for layer_id in range(layer_num):
device_pool.kv_buffer[layer_id].fill_(17 + cp_rank)
device_pool.index_k_with_scale_buffer[layer_id].fill_(23 + cp_rank)
device_pool.kv_buffer[layer_id][
local_valid_physical_locs
] = valid_kv_rows
page_ids = torch.div(
local_valid_physical_locs, page_size, rounding_mode="floor"
).to(torch.long)
page_offsets = torch.remainder(
local_valid_physical_locs, page_size
).to(torch.long)
index_view = self._index_row_bytes(
device_pool.index_k_with_scale_buffer[layer_id], page_size
)
index_view[page_ids, page_offsets] = valid_index_rows
expected_kv = [
device_pool.kv_buffer[layer_id][local_valid_physical_locs]
.detach()
.clone()
for layer_id in range(layer_num)
]
expected_index = []
page_ids = torch.div(
local_valid_physical_locs, page_size, rounding_mode="floor"
).to(torch.long)
page_offsets = torch.remainder(local_valid_physical_locs, page_size).to(
torch.long
)
for layer_id in range(layer_num):
expected_index.append(
self._index_row_bytes(
device_pool.index_k_with_scale_buffer[layer_id], page_size
)[page_ids, page_offsets]
.detach()
.clone()
)
host_pool.backup_from_device_all_layer(
device_pool,
host_indices,
src_physical_padded_locs,
io_backend="direct",
)
torch.cuda.synchronize()
for layer_id in range(layer_num):
device_pool.kv_buffer[layer_id][
dst_physical_padded_locs.to(device="cuda")
].fill_(91 + cp_rank)
dst_pages = torch.div(
dst_physical_padded_locs, page_size, rounding_mode="floor"
).to(device="cuda", dtype=torch.long)
device_pool.index_k_with_scale_buffer[layer_id][dst_pages].fill_(
97 + cp_rank
)
host_pool.begin_load_to_device_op(
host_indices,
dst_physical_padded_locs,
io_backend="direct",
)
try:
for layer_id in range(layer_num):
host_pool.load_to_device_per_layer(
device_pool,
host_indices,
dst_physical_padded_locs,
layer_id=layer_id,
io_backend="direct",
)
finally:
host_pool.end_load_to_device_op()
torch.cuda.synchronize()
dst_valid_logical_locs_rank = dst_valid_locs[
layout.owned_by_this_rank(dst_valid_locs)
]
dst_valid_physical_locs = layout.logical_locs_to_physical(
dst_valid_logical_locs_rank
).to(device="cuda")
self.assertEqual(
int(dst_valid_physical_locs.numel()), int(local_valid_rows)
)
for layer_id in range(layer_num):
got_kv = device_pool.kv_buffer[layer_id][dst_valid_physical_locs]
self.assertTrue(
torch.equal(got_kv, expected_kv[layer_id]),
f"KV valid-row L2 roundtrip mismatch cp_rank={cp_rank} layer={layer_id}",
)
got_pages = torch.div(
dst_valid_physical_locs, page_size, rounding_mode="floor"
).to(torch.long)
got_offsets = torch.remainder(dst_valid_physical_locs, page_size).to(
torch.long
)
got_index = self._index_row_bytes(
device_pool.index_k_with_scale_buffer[layer_id], page_size
)[got_pages, got_offsets]
self.assertTrue(
torch.equal(got_index, expected_index[layer_id]),
f"index valid-row L2 roundtrip mismatch cp_rank={cp_rank} layer={layer_id}",
)
def test_fp8_l2_loaded_bs5_suffix_materializes_from_ragged_logical_locs(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
page_size = 64
cp_size = 8
layer_num = 1
prefix_lens = [640, 640, 640, 640, 640]
extend_lens = [95, 81, 140, 66, 86]
src_valid_locs, src_padded_locs = self._build_owned_logical_out_locs(
extend_lens=extend_lens,
prefix_lens=prefix_lens,
page_size=page_size,
cp_size=cp_size,
)
dst_valid_locs, dst_padded_locs = self._build_owned_logical_out_locs(
extend_lens=extend_lens,
prefix_lens=prefix_lens,
page_size=page_size,
cp_size=cp_size,
dst_page_stride=128,
)
max_physical_page = (
int(torch.div(dst_padded_locs.max(), page_size, rounding_mode="floor"))
// cp_size
+ 4
)
size = max_physical_page * page_size
expected_rows_by_dst_loc: dict[int, torch.Tensor] = {}
dense_outputs = []
dense_locs_reference = None
device_pool = NSATokenToKVPool(
size=size,
page_size=page_size,
kv_lora_rank=512,
dtype=torch.float8_e4m3fn,
qk_rope_head_dim=64,
layer_num=layer_num,
device="cuda",
enable_memory_saver=False,
kv_cache_dim=656,
index_head_dim=128,
)
host_pool = NSATokenToKVPoolHost(
device_pool=device_pool,
host_to_device_ratio=2.0,
host_size=0,
page_size=page_size,
layout="page_first_direct",
pin_memory=True,
device="cpu",
)
for cp_rank in range(cp_size):
layout = CpSharedKVLayout(
page_size=page_size,
cp_size=cp_size,
cp_rank=cp_rank,
)
plan = build_batch_page_aligned_in_seq_split_plan(
extend_lens=extend_lens,
prefix_lens=prefix_lens,
page_size=page_size,
cp_size=cp_size,
cp_rank=cp_rank,
)
forward_batch = SimpleNamespace(
uses_cp_shared_kv=True,
cp_shared_kv_layout=layout,
nsa_cp_metadata=NSAContextParallelMetadata(
batch_size=len(extend_lens),
batch_plan=plan,
page_aligned=True,
page_size=page_size,
extend_prefix_len=prefix_lens[0],
),
out_cache_loc=src_valid_locs.clone(),
token_to_kv_pool=SimpleNamespace(page_size=page_size),
)
with patch(
"sglang.srt.layers.attention.nsa.utils.is_nsa_prefill_cp_round_robin_split",
return_value=False,
):
local_src_valid_locs = get_cp_shared_kv_local_out_cache_loc(
forward_batch
)
local_rows = int(local_src_valid_locs.numel())
local_dst_valid_locs = dst_valid_locs[
layout.owned_by_this_rank(dst_valid_locs)
]
self.assertEqual(int(local_dst_valid_locs.numel()), local_rows)
src_owned_padded_locs = src_padded_locs[
layout.owned_by_this_rank(src_padded_locs)
]
dst_owned_padded_locs = dst_padded_locs[
layout.owned_by_this_rank(dst_padded_locs)
]
src_physical_padded_locs = layout.logical_locs_to_physical(
src_owned_padded_locs
)
dst_physical_padded_locs = layout.logical_locs_to_physical(
dst_owned_padded_locs
)
host_indices = torch.arange(
cp_rank * 4096,
cp_rank * 4096 + int(src_physical_padded_locs.numel()),
dtype=torch.int64,
)
for layer_id in range(layer_num):
device_pool.kv_buffer[layer_id].fill_(11 + cp_rank)
if local_rows > 0:
local_src_physical_locs = layout.logical_locs_to_physical(
local_src_valid_locs
).to(device="cuda")
row_bytes = (
torch.arange(
local_rows * 656,
device="cuda",
dtype=torch.int64,
)
.view(local_rows, 1, 656)
.add_(cp_rank * 37)
.remainder_(251)
.to(torch.uint8)
)
device_pool.kv_buffer[0][local_src_physical_locs] = row_bytes
for dst_loc, row in zip(
local_dst_valid_locs.tolist(),
row_bytes.detach().cpu(),
strict=True,
):
expected_rows_by_dst_loc[int(dst_loc)] = row
host_pool.backup_from_device_all_layer(
device_pool,
host_indices,
src_physical_padded_locs,
io_backend="direct",
)
torch.cuda.synchronize()
for layer_id in range(layer_num):
device_pool.kv_buffer[layer_id][
dst_physical_padded_locs.to(device="cuda")
].fill_(99 + cp_rank)
host_pool.begin_load_to_device_op(
host_indices,
dst_physical_padded_locs,
io_backend="direct",
)
try:
host_pool.load_to_device_per_layer(
device_pool,
host_indices,
dst_physical_padded_locs,
layer_id=0,
io_backend="direct",
)
finally:
host_pool.end_load_to_device_op()
torch.cuda.synchronize()
with patch.object(
runtime,
"_all_reduce_materialized_buffer",
side_effect=lambda buffer, *args, **kwargs: buffer,
):
dense_kv, dense_locs = runtime.materialize_shared_token_kv_buffer(
kv_cache=device_pool.kv_buffer[0],
logical_locs=dst_valid_locs.to(device="cuda"),
layout=layout,
page_size=page_size,
nvtx_source="test.ragged_l2_loaded",
nvtx_layer_id=0,
)
dense_outputs.append(dense_kv.detach().cpu().to(torch.int16))
dense_locs = dense_locs.detach().cpu()
if dense_locs_reference is None:
dense_locs_reference = dense_locs
else:
self.assertTrue(torch.equal(dense_locs_reference, dense_locs))
self.assertEqual(set(expected_rows_by_dst_loc), set(dst_valid_locs.tolist()))
assert dense_locs_reference is not None
merged_dense = torch.stack(dense_outputs, dim=0).sum(dim=0).to(torch.uint8)
got = merged_dense[dense_locs_reference]
expected = torch.stack(
[expected_rows_by_dst_loc[int(loc)] for loc in dst_valid_locs.tolist()],
dim=0,
)
self.assertTrue(
torch.equal(got, expected),
"L2-loaded FP8 suffix pages must remain readable through RAGGED logical-loc materialize",
)
def test_fp8_l2_loaded_index_pages_materialize_after_fused_store_bs5(self):
from sglang.jit_kernel.fused_store_index_cache import (
can_use_nsa_fused_store,
fused_store_index_k_cache,
)
from sglang.srt.layers.attention.nsa import (
cp_shared_kv_runtime as runtime,
index_buf_accessor,
)
from sglang.srt.layers.attention.nsa.triton_kernel import act_quant
if not hasattr(torch, "float8_e4m3fn"):
self.skipTest("torch.float8_e4m3fn unavailable")
if not can_use_nsa_fused_store(torch.bfloat16, torch.int64, 64):
self.skipTest("NSA fused index store JIT unavailable")
page_size = 64
cp_size = 8
layer_num = 1
prefix_lens = [640, 640, 640, 640, 640]
extend_lens = [95, 81, 140, 66, 86]
src_valid_locs, src_padded_locs = self._build_owned_logical_out_locs(
extend_lens=extend_lens,
prefix_lens=prefix_lens,
page_size=page_size,
cp_size=cp_size,
)
dst_valid_locs, dst_padded_locs = self._build_owned_logical_out_locs(
extend_lens=extend_lens,
prefix_lens=prefix_lens,
page_size=page_size,
cp_size=cp_size,
dst_page_stride=128,
)
max_physical_page = (
int(torch.div(dst_padded_locs.max(), page_size, rounding_mode="floor"))
// cp_size
+ 4
)
size = max_physical_page * page_size
device_pool = NSATokenToKVPool(
size=size,
page_size=page_size,
kv_lora_rank=512,
dtype=torch.float8_e4m3fn,
qk_rope_head_dim=64,
layer_num=layer_num,
device="cuda",
enable_memory_saver=False,
kv_cache_dim=656,
index_head_dim=128,
)
host_pool = NSATokenToKVPoolHost(
device_pool=device_pool,
host_to_device_ratio=2.0,
host_size=0,
page_size=page_size,
layout="page_first_direct",
pin_memory=True,
device="cpu",
)
dst_pages_by_req = []
cursor = 0
for extend_len in extend_lens:
num_pages = (int(extend_len) + page_size - 1) // page_size
req_padded = dst_padded_locs[cursor : cursor + num_pages * page_size]
dst_pages_by_req.append(
torch.div(
req_padded[::page_size], page_size, rounding_mode="floor"
).tolist()
)
cursor += num_pages * page_size
max_pages = max(len(pages) for pages in dst_pages_by_req)
logical_pages = torch.zeros(
(len(extend_lens), max_pages), dtype=torch.int64, device="cuda"
)
for req_id, pages in enumerate(dst_pages_by_req):
logical_pages[req_id, : len(pages)] = torch.tensor(
pages, dtype=torch.int64, device="cuda"
)
expected_rows_by_dst_loc: dict[int, torch.Tensor] = {}
dense_outputs = []
dense_pages_reference = None
for cp_rank in range(cp_size):
layout = CpSharedKVLayout(
page_size=page_size,
cp_size=cp_size,
cp_rank=cp_rank,
)
plan = build_batch_page_aligned_in_seq_split_plan(
extend_lens=extend_lens,
prefix_lens=prefix_lens,
page_size=page_size,
cp_size=cp_size,
cp_rank=cp_rank,
)
forward_batch = SimpleNamespace(
uses_cp_shared_kv=True,
cp_shared_kv_layout=layout,
nsa_cp_metadata=NSAContextParallelMetadata(
batch_size=len(extend_lens),
batch_plan=plan,
page_aligned=True,
page_size=page_size,
extend_prefix_len=prefix_lens[0],
),
out_cache_loc=src_valid_locs.clone(),
token_to_kv_pool=SimpleNamespace(page_size=page_size),
)
with patch(
"sglang.srt.layers.attention.nsa.utils.is_nsa_prefill_cp_round_robin_split",
return_value=False,
):
local_src_valid_locs = get_cp_shared_kv_local_out_cache_loc(
forward_batch
)
local_rows = int(local_src_valid_locs.numel())
local_dst_valid_locs = dst_valid_locs[
layout.owned_by_this_rank(dst_valid_locs)
]
self.assertEqual(int(local_dst_valid_locs.numel()), local_rows)
src_owned_padded_locs = src_padded_locs[
layout.owned_by_this_rank(src_padded_locs)
]
dst_owned_padded_locs = dst_padded_locs[
layout.owned_by_this_rank(dst_padded_locs)
]
src_physical_padded_locs = layout.logical_locs_to_physical(
src_owned_padded_locs
)
dst_physical_padded_locs = layout.logical_locs_to_physical(
dst_owned_padded_locs
)
host_indices = torch.arange(
cp_rank * 4096,
cp_rank * 4096 + int(src_physical_padded_locs.numel()),
dtype=torch.int64,
)
device_pool.index_k_with_scale_buffer[0].fill_(0)
if local_rows > 0:
local_compute_rows = sum(
int(x) for x in plan.request_compute_rank_local_tokens
)
torch.manual_seed(20260607 + cp_rank)
local_key = (
torch.randn(
(local_compute_rows, 128),
device="cuda",
dtype=torch.bfloat16,
)
* 2.0
+ 0.125
).contiguous()
valid_key = select_cp_local_valid_rows_for_cache_write(
forward_batch,
local_key,
).contiguous()
self.assertEqual(int(valid_key.shape[0]), local_rows)
src_physical_valid_locs = layout.logical_locs_to_physical(
local_src_valid_locs
).to(device="cuda")
fused_store_index_k_cache(
valid_key,
device_pool.index_k_with_scale_buffer[0],
src_physical_valid_locs.contiguous(),
page_size=page_size,
)
for dst_loc, row in zip(
local_dst_valid_locs.tolist(),
valid_key.detach().cpu(),
strict=True,
):
expected_rows_by_dst_loc[int(dst_loc)] = row
host_pool.backup_from_device_all_layer(
device_pool,
host_indices,
src_physical_padded_locs,
io_backend="direct",
)
torch.cuda.synchronize()
device_pool.index_k_with_scale_buffer[0][
torch.div(
dst_physical_padded_locs.to(device="cuda"),
page_size,
rounding_mode="floor",
).to(torch.long)
].fill_(77 + cp_rank)
host_pool.begin_load_to_device_op(
host_indices,
dst_physical_padded_locs,
io_backend="direct",
)
try:
host_pool.load_to_device_per_layer(
device_pool,
host_indices,
dst_physical_padded_locs,
layer_id=0,
io_backend="direct",
)
finally:
host_pool.end_load_to_device_op()
torch.cuda.synchronize()
slot_remap = runtime.build_shared_paged_buffer_slot_remap(
device_pool.index_k_with_scale_buffer[0],
logical_pages,
layout,
)
with patch.object(
runtime,
"_all_reduce_materialized_buffer",
side_effect=lambda buffer, *args, **kwargs: buffer,
), patch.object(
runtime, "_try_tai_materialize_shared_pages", return_value=None
), patch.object(
runtime,
"_try_tai_ipc_materialize_paged_buffer_page_slots_into",
return_value=False,
):
dense_index, dense_pages = runtime.materialize_shared_paged_buffer(
device_pool.index_k_with_scale_buffer[0],
logical_pages,
layout,
slot_remap=slot_remap,
nvtx_source="test.index_l2_loaded",
nvtx_layer_id=0,
)
dense_outputs.append(dense_index.detach().cpu().to(torch.int16))
dense_pages = dense_pages.detach().cpu()
if dense_pages_reference is None:
dense_pages_reference = dense_pages
else:
self.assertTrue(torch.equal(dense_pages_reference, dense_pages))
self.assertEqual(set(expected_rows_by_dst_loc), set(dst_valid_locs.tolist()))
assert dense_pages_reference is not None
merged_dense = torch.stack(dense_outputs, dim=0).sum(dim=0).to(torch.uint8)
class FakePool:
page_size = 64
index_head_dim = 128
quant_block_size = 128
device = "cuda"
for req_id, (extend_len, req_locs) in enumerate(
zip(extend_lens, torch.split(dst_valid_locs, extend_lens))
):
page_count = len(dst_pages_by_req[req_id])
page_indices = dense_pages_reference[req_id, :page_count].to(
device="cuda"
)
dense_gpu = merged_dense.to(device="cuda")
got_k = index_buf_accessor.GetK.execute(
FakePool,
dense_gpu,
seq_len=int(extend_len),
page_indices=page_indices.contiguous(),
)
got_s = index_buf_accessor.GetS.execute(
FakePool,
dense_gpu,
seq_len=int(extend_len),
page_indices=page_indices.contiguous(),
)
got_deq = (
got_k.contiguous().view(torch.float8_e4m3fn).float()
* got_s.contiguous().view(torch.float32).view(-1, 1)
)
expected_key = torch.stack(
[expected_rows_by_dst_loc[int(loc)] for loc in req_locs.tolist()],
dim=0,
).to(device="cuda")
ref_fp8, ref_scale = act_quant(
expected_key.contiguous(),
block_size=128,
scale_fmt="ue8m0",
)
ref_deq = ref_fp8.float() * ref_scale.float().view(-1, 1)
torch.testing.assert_close(
got_deq,
expected_key.float(),
rtol=0.20,
atol=0.65,
)
torch.testing.assert_close(
got_deq,
ref_deq,
rtol=0.35,
atol=0.90,
)
def test_device_to_host_indexer_kernel_layer_first(self):
self._run_device_to_host_indexer_copy(
io_backend="kernel", layout="layer_first"
@@ -148,6 +968,127 @@ class TestNSAHiCacheTransfer(CustomTestCase):
io_backend="direct", layout="page_first_direct"
)
def test_fp8_page_first_direct_roundtrip_preserves_kv_and_indexer_pages(self):
page_size = 64
layer_num = 3
size = page_size * 20
device_pool = NSATokenToKVPool(
size=size,
page_size=page_size,
kv_lora_rank=512,
dtype=torch.float8_e4m3fn,
qk_rope_head_dim=64,
layer_num=layer_num,
device="cuda",
enable_memory_saver=False,
kv_cache_dim=656,
index_head_dim=128,
)
host_pool = NSATokenToKVPoolHost(
device_pool=device_pool,
host_to_device_ratio=2.0,
host_size=0,
page_size=page_size,
layout="page_first_direct",
pin_memory=True,
device="cpu",
)
for layer_id in range(layer_num):
kv_buf = device_pool.kv_buffer[layer_id]
kv_data = torch.arange(
kv_buf.numel(), device="cuda", dtype=torch.int64
).view_as(kv_buf)
kv_buf.copy_(((kv_data + 17 * layer_id) % 251).to(torch.uint8))
index_buf = device_pool.index_k_with_scale_buffer[layer_id]
index_data = torch.arange(
index_buf.numel(), device="cuda", dtype=torch.int64
).view_as(index_buf)
index_buf.copy_(((index_data + 29 * layer_id) % 253).to(torch.uint8))
src_pages = torch.tensor([1, 2, 5, 8], dtype=torch.int64)
host_pages = torch.tensor([3, 4, 7, 12], dtype=torch.int64)
dst_pages = torch.tensor([10, 11, 13, 15], dtype=torch.int64)
device_indices = self._token_indices_for_pages(
src_pages, page_size, device="cpu"
)
host_indices = self._token_indices_for_pages(
host_pages, page_size, device="cpu"
)
load_device_indices = self._token_indices_for_pages(
dst_pages, page_size, device="cpu"
)
expected_kv = []
expected_index = []
for layer_id in range(layer_num):
expected_kv.append(
[
device_pool.kv_buffer[layer_id][
int(page) * page_size : (int(page) + 1) * page_size
]
.detach()
.clone()
for page in src_pages.tolist()
]
)
expected_index.append(
[
device_pool.index_k_with_scale_buffer[layer_id][int(page)]
.detach()
.clone()
for page in src_pages.tolist()
]
)
host_pool.backup_from_device_all_layer(
device_pool, host_indices, device_indices, io_backend="direct"
)
torch.cuda.synchronize()
for layer_id in range(layer_num):
for page in dst_pages.tolist():
start = int(page) * page_size
device_pool.kv_buffer[layer_id][start : start + page_size].fill_(0)
device_pool.index_k_with_scale_buffer[layer_id][int(page)].fill_(0)
host_pool.begin_load_to_device_op(
host_indices, load_device_indices, io_backend="direct"
)
try:
for layer_id in range(layer_num):
host_pool.load_to_device_per_layer(
device_pool,
host_indices,
load_device_indices,
layer_id=layer_id,
io_backend="direct",
)
finally:
host_pool.end_load_to_device_op()
torch.cuda.synchronize()
for layer_id in range(layer_num):
for page_idx, dst_page in enumerate(dst_pages.tolist()):
dst_start = int(dst_page) * page_size
got_kv = device_pool.kv_buffer[layer_id][
dst_start : dst_start + page_size
]
self.assertTrue(
torch.equal(got_kv, expected_kv[layer_id][page_idx]),
f"KV roundtrip mismatch layer={layer_id} dst_page={dst_page}",
)
got_index = device_pool.index_k_with_scale_buffer[layer_id][
int(dst_page)
]
self.assertTrue(
torch.equal(got_index, expected_index[layer_id][page_idx]),
f"index roundtrip mismatch layer={layer_id} dst_page={dst_page}",
)
class TestPageFirstDirectAllLayerBackupRoute(CustomTestCase):
def test_mla_page_first_direct_all_layer_backup_uses_tai_per_layer_route(self):