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