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:
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user