Lock fp8 CP cache-hit roundtrips across batch requests

The GSM8K cache-hit regression needs executable coverage for the fp8 paths that combine bs>1 page planning, valid-row writes, prefix materialization, and HiCache L2 reload.  These tests construct bs=5 page-aligned layouts and assert valid rows survive the persistent-page store/materialize and host roundtrip paths.

Constraint: Production runs use fp8_e4m3 CP shared KV with page_first_direct HiCache and bs>1 prefill.

Rejected: Validate with bf16-only tests | the observed regressions are fp8/cache-hit sensitive and bf16 coverage does not exercise scale/index byte layouts.

Confidence: medium

Scope-risk: narrow

Directive: Keep fp8 cache-hit tests aligned with the page-as-minimum-cache-unit contract.

Tested: python -m py_compile on changed runtime files.

Not-tested: CUDA/TAI tests not run locally; local pytest blocked before collection by missing orjson dependency.
(cherry picked from commit 6600f75fc8ee803610df0078bb880df808655635)
This commit is contained in:
laoyao0822
2026-06-07 23:53:42 +08:00
parent 0f9b445131
commit 8f23b06427
2 changed files with 545 additions and 0 deletions

View File

@@ -3402,6 +3402,256 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
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_mla_persistent_pages_survive_bs5_cache_hit_materialize(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
batch_size = 5
prefix_lens = [640, 640, 640, 640, 640]
# GSM8K second-run cache hits usually consume one or two full suffix
# pages from the first run, then recompute only a short tail.
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 prefix_len, extend_len in 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)
req_locs.append(
torch.arange(
logical_page * page_size,
logical_page * page_size + valid_rows,
dtype=torch.int64,
)
)
remaining -= valid_rows
self.assertEqual(remaining, 0)
request_pages.append(req_pages)
request_valid_locs.append(torch.cat(req_locs, dim=0))
all_valid_locs.append(request_valid_locs[-1])
# Second-run radix flooring should only expose complete cached suffix
# pages. Tail partial pages remain current rows and are not part of the
# cache-hit prefix materialized below.
cache_hit_pages: list[list[int]] = []
for extend_len, pages in zip(extend_lens, request_pages):
full_pages = int(extend_len) // page_size
cache_hit_pages.append(pages[:full_pages])
max_cache_hit_pages = max(len(pages) for pages in cache_hit_pages)
self.assertGreater(max_cache_hit_pages, 0)
logical_pages = torch.zeros(
(batch_size, max_cache_hit_pages),
dtype=torch.int64,
device="cuda",
)
for req_id, pages in enumerate(cache_hit_pages):
if 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")
prefix_slot_spans = runtime.build_batch_prefix_slot_spans(
logical_pages=logical_pages,
prefix_lens_cpu=[len(pages) * page_size for pages in cache_hit_pages],
page_size=page_size,
)
max_logical_page = max(max(pages) for pages in request_pages)
physical_token_capacity = (((max_logical_page - 1) // cp_size) + 3) * page_size
expected_by_loc: dict[int, torch.Tensor] = {}
rank_dense_buffers = []
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_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)
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()))
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(
(physical_token_capacity, 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)
if logical_locs.numel() > 0:
physical_locs = layout.logical_locs_to_physical(logical_locs)
expected = torch.zeros_like(pool.kv_buffer[0])
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,
)
torch.testing.assert_close(
pool.kv_buffer[0],
expected,
atol=0,
rtol=0,
)
for loc, row in zip(logical_locs.cpu().tolist(), expected[physical_locs]):
expected_by_loc[int(loc)] = row.detach()
slot_remap = runtime.build_shared_token_kv_slot_remap(
pool.kv_buffer[0],
logical_locs=out_cache_loc,
remap_logical_pages=logical_pages,
layout=layout,
page_size=page_size,
)
with patch.object(
runtime, "_all_reduce_materialized_buffer_range", _identity_all_reduce
), patch.object(
runtime, "_try_tai_ipc_materialize_token_kv_page_slots_into",
return_value=False,
):
dense_kv_cache, _dense_locs = (
runtime.materialize_prefix_and_reuse_current_kv_page_slots(
kv_cache=pool.kv_buffer[0],
logical_locs=out_cache_loc,
current_kv_cache=pool.kv_buffer[0].new_empty((0, 1, 656)),
current_locs=out_cache_loc.new_empty((0,)),
slot_remap=slot_remap,
layout=layout,
page_size=page_size,
prefix_pages=0,
prefix_slot_spans=prefix_slot_spans,
current_slot_spans=[],
layer_id=0,
)
)
rank_dense_buffers.append(dense_kv_cache.to(torch.int16))
merged_dense = torch.stack(rank_dense_buffers, dim=0).sum(dim=0).to(torch.uint8)
dense_pages = runtime.build_slot_page_remap(logical_pages)[1].to("cuda")
for req_id, pages in enumerate(cache_hit_pages):
for page_offset, logical_page in enumerate(pages):
dense_page = int(dense_pages[req_id, page_offset].item())
for token_offset in range(page_size):
logical_loc = int(logical_page) * page_size + token_offset
expected_row = expected_by_loc[logical_loc].to(device="cuda")
actual_row = merged_dense[dense_page * page_size + token_offset]
torch.testing.assert_close(actual_row, expected_row, 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 (

View File

@@ -1,9 +1,12 @@
import unittest
import threading
from types import SimpleNamespace
from unittest.mock import patch
import torch
from sglang.srt.managers.cache_controller import HiCacheController, HiCacheWriteFailure
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
from sglang.srt.layers.attention.nsa.utils import (
NSAContextParallelMetadata,
build_batch_page_aligned_in_seq_split_plan,
@@ -206,6 +209,298 @@ class TestNSAHiCacheTransfer(CustomTestCase):
# [token0(index+scale), token1(index+scale), ...].
return row_values.view(row_values.shape[0], page_size, -1)
def test_cp_hicache_controller_grouped_bs5_l2_roundtrip_preserves_valid_rows(
self,
):
page_size = 64
cp_size = 8
layer_num = 2
prefix_lens = [640, 640, 640, 640, 640]
extend_lens = [95, 81, 140, 66, 86]
physical_size = page_size * 256
logical_size = physical_size * cp_size
original_alloc = ALLOC_MEMORY_FUNCS["cuda"]
ALLOC_MEMORY_FUNCS["cuda"] = alloc_with_pin_memory
try:
for cp_rank in range(cp_size):
device_pool = NSATokenToKVPool(
size=physical_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,
)
allocator = CPSharedPagedTokenToKVPoolAllocator(
logical_size=logical_size,
physical_size=physical_size,
page_size=page_size,
dtype=device_pool.store_dtype,
device="cuda",
kvcache=device_pool,
need_sort=False,
cp_size=cp_size,
cp_rank=cp_rank,
)
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",
)
layout = CpSharedKVLayout(
page_size=page_size,
cp_size=cp_size,
cp_rank=cp_rank,
)
controller = HiCacheController(
token_to_kv_pool_allocator=allocator,
mem_pool_host=host_pool,
page_size=page_size,
tp_group=None,
load_cache_event=threading.Event(),
write_policy="write_through",
io_backend="direct",
cp_shared_kv_layout=layout,
)
page_compute_owners = []
for extend_len, prefix_len in zip(
extend_lens, prefix_lens, strict=True
):
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)
page_compute_owners.extend(owners)
prefix_lens_cpu = torch.tensor(prefix_lens, dtype=torch.int64)
seq_lens_cpu = torch.tensor(
[p + e for p, e in zip(prefix_lens, extend_lens, strict=True)],
dtype=torch.int64,
)
out_cache_loc = allocator.alloc_extend_compute_owner(
prefix_lens=prefix_lens_cpu.to(device="cuda"),
prefix_lens_cpu=prefix_lens_cpu,
seq_lens=seq_lens_cpu.to(device="cuda"),
seq_lens_cpu=seq_lens_cpu,
last_loc=torch.full(
(len(extend_lens),),
-1,
dtype=torch.int64,
device="cuda",
),
extend_num_tokens=sum(extend_lens),
page_compute_owners=page_compute_owners,
)
self.assertIsNotNone(out_cache_loc)
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=out_cache_loc.detach().cpu().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_physical_locs = layout.logical_locs_to_physical(
local_valid_logical_locs
).to(device="cuda")
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_(cp_rank * 31)
.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_(cp_rank * 43)
.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]), int(local_valid_logical_locs.numel())
)
self.assertEqual(
int(valid_index_rows.shape[0]), int(local_valid_logical_locs.numel())
)
for layer_id in range(layer_num):
device_pool.kv_buffer[layer_id].fill_(13 + cp_rank)
device_pool.index_k_with_scale_buffer[layer_id].fill_(17 + cp_rank)
if local_valid_physical_locs.numel() > 0:
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)
self._index_row_bytes(
device_pool.index_k_with_scale_buffer[layer_id], page_size
)[page_ids, page_offsets] = valid_index_rows
reservations = []
nodes = []
expected_by_node = {}
cursor = 0
for req_id, extend_len in enumerate(extend_lens):
req_locs = out_cache_loc[cursor : cursor + extend_len].detach().cpu()
cursor += extend_len
node_id = 2026060700 + cp_rank * 100 + req_id
reservation = controller.reserve_write_cp(req_locs, node_id=node_id)
self.assertNotIsInstance(reservation, HiCacheWriteFailure)
reservations.append(reservation)
nodes.append(
SimpleNamespace(
id=node_id,
host_len=int(reservation.metadata.logical_len),
cp_hicache=reservation.metadata,
)
)
owned_req_locs = req_locs[layout.owned_by_this_rank(req_locs)]
owned_req_physical = layout.logical_locs_to_physical(
owned_req_locs
).to(device="cuda")
expected_by_node[node_id] = {
"logical_len": int(extend_len),
"kv": [
device_pool.kv_buffer[layer_id][owned_req_physical]
.detach()
.clone()
for layer_id in range(layer_num)
],
"index": [
self._index_row_bytes(
device_pool.index_k_with_scale_buffer[layer_id],
page_size,
)[
torch.div(
owned_req_physical,
page_size,
rounding_mode="floor",
).to(torch.long),
torch.remainder(owned_req_physical, page_size).to(
torch.long
),
]
.detach()
.clone()
for layer_id in range(layer_num)
],
}
for reservation in reservations:
controller.submit_write_cp_per_layer(
reservation, catch_up_all_layers=False
)
for layer_id in range(layer_num):
controller.on_layer_end(layer_id, source="target")
torch.cuda.synchronize()
for layer_id in range(layer_num):
device_pool.kv_buffer[layer_id].fill_(89 + cp_rank)
device_pool.index_k_with_scale_buffer[layer_id].fill_(97 + cp_rank)
loaded_visible_by_node = {}
for node in nodes:
visible = controller.load_cp([node], node_id=node.id)
self.assertIsNotNone(visible)
loaded_visible_by_node[node.id] = visible.detach().clone()
producer_id = controller.start_loading()
self.assertGreaterEqual(producer_id, 0)
torch.cuda.synchronize()
for node in nodes:
visible = loaded_visible_by_node[node.id]
self.assertEqual(
int(visible.numel()), expected_by_node[node.id]["logical_len"]
)
owned_visible = visible[layout.owned_by_this_rank(visible)]
loaded_physical = layout.logical_locs_to_physical(
owned_visible
).to(device="cuda")
for layer_id in range(layer_num):
got_kv = device_pool.kv_buffer[layer_id][loaded_physical]
self.assertTrue(
torch.equal(got_kv, expected_by_node[node.id]["kv"][layer_id]),
f"grouped CP L2 KV mismatch cp_rank={cp_rank} node_id={node.id} layer={layer_id}",
)
page_ids = torch.div(
loaded_physical, page_size, rounding_mode="floor"
).to(torch.long)
page_offsets = torch.remainder(loaded_physical, page_size).to(
torch.long
)
got_index = self._index_row_bytes(
device_pool.index_k_with_scale_buffer[layer_id], page_size
)[page_ids, page_offsets]
self.assertTrue(
torch.equal(
got_index, expected_by_node[node.id]["index"][layer_id]
),
f"grouped CP L2 index mismatch cp_rank={cp_rank} node_id={node.id} layer={layer_id}",
)
finally:
ALLOC_MEMORY_FUNCS["cuda"] = original_alloc
def test_fp8_page_first_direct_bs5_valid_rows_survive_l2_roundtrip(self):
page_size = 64
cp_size = 8