From 8f23b0642714bc8d293a572bde7cccbff3aee49f Mon Sep 17 00:00:00 2001 From: laoyao0822 Date: Sun, 7 Jun 2026 23:53:42 +0800 Subject: [PATCH] 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) --- .../mem_cache/test_cp_shared_kv_runtime.py | 250 +++++++++++++++ .../unit/mem_cache/test_nsa_pool_host_unit.py | 295 ++++++++++++++++++ 2 files changed, 545 insertions(+) diff --git a/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py b/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py index 48d69e6b1..b7b485cdd 100644 --- a/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py +++ b/test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py @@ -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 ( diff --git a/test/registered/unit/mem_cache/test_nsa_pool_host_unit.py b/test/registered/unit/mem_cache/test_nsa_pool_host_unit.py index 405d55c7b..4348f1fd2 100644 --- a/test/registered/unit/mem_cache/test_nsa_pool_host_unit.py +++ b/test/registered/unit/mem_cache/test_nsa_pool_host_unit.py @@ -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