Enable CP HiCache direct transfers to use layer-page host layout
CP shared-KV HiCache transfers are per-layer, so the host layout should match the access pattern instead of forcing page-major strides through every layer. This adds a direct-only layer_page_first layout, routes per-layer KV and NSA index backup/load through the TAI LF<->LPF direct kernels, and keeps storage/page-buffer metadata paths fail-fast until their page-level contract is redesigned.\n\nThe direct controller keeps host indices in caller order for both page_first_direct and layer_page_first because the TAI direct path requires CPU index descriptors and owns descriptor coalescing. All-layer backup intentionally loops over per-layer direct kernels rather than using the sgl-kernel all-layer direct ABI.\n\nConstraint: layer_page_first is currently host-only CP HiCache; storage backends assume page-major contiguous page metadata.\nConstraint: TAI LPF direct kernels require CPU int64 page indices and complete page spans.\nRejected: silently fallback to SM copy when TAI LPF kernels are missing | that hides production performance regressions.\nRejected: support storage page metadata in this commit | LPF requires a layer-page-level storage contract, not a one-pointer-per-page contract.\nConfidence: medium\nScope-risk: moderate\nDirective: Do not enable storage or kernel backend for layer_page_first without redesigning page-buffer metadata and adding remote ETE coverage.\nTested: local py_compile for touched runtime files.\nTested: remote py_compile in g0034 container for touched runtime files.\nTested: remote targeted pytest: 5 passed for parser/storage/layout/move_indices smoke coverage.\nNot-tested: full CP HiCache ETE with --hicache-mem-layout layer_page_first after this commit step.\nNot-tested: combined CUDA roundtrip tests in one pytest process; previous independent runs passed but combined run exposed a host-memory registration lifecycle issue.
This commit is contained in:
@@ -105,7 +105,7 @@ for _schema in (
|
||||
if "already" not in str(exc).lower() and "duplicate" not in str(exc).lower():
|
||||
raise
|
||||
|
||||
from sglang.srt.managers.cache_controller import HiCacheController
|
||||
from sglang.srt.managers.cache_controller import CacheOperation, HiCacheController
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
from sglang.srt.mem_cache.hiradix_cache import CpHiCacheNodeMetadata
|
||||
from sglang.srt.mem_cache.memory_pool_host import (
|
||||
@@ -449,6 +449,139 @@ class TestPageFirstPerLayerBackupTaiKernel(CustomTestCase):
|
||||
self.assertEqual(kwargs["layer_id"], 1)
|
||||
self.assertEqual(kwargs["page_size"], 1)
|
||||
|
||||
def test_mla_layer_page_first_per_layer_backup_uses_direct_lf_lpf(self):
|
||||
calls = []
|
||||
|
||||
def fake_direct(src_ptrs, dst_ptrs, src_indices, dst_indices, **kwargs):
|
||||
calls.append(
|
||||
(src_ptrs, dst_ptrs, src_indices.clone(), dst_indices.clone(), kwargs)
|
||||
)
|
||||
|
||||
host_pool = MLATokenToKVPoolHost.__new__(MLATokenToKVPoolHost)
|
||||
host_pool.layout = "layer_page_first"
|
||||
host_pool.page_size = 4
|
||||
host_pool.kv_buffer = torch.empty((3, 8, 4, 1, 16), dtype=torch.uint8)
|
||||
device_pool = type("DevicePool", (), {})()
|
||||
device_pool.kv_buffer = torch.empty((3, 32, 1, 16), dtype=torch.uint8)
|
||||
host_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||
device_indices = torch.tensor([12, 13, 14, 15], dtype=torch.int64)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.mem_cache.memory_pool_host._load_tai_transfer_kv_per_layer_direct_lf_lpf",
|
||||
return_value=fake_direct,
|
||||
):
|
||||
host_pool.backup_from_device_per_layer(
|
||||
device_pool,
|
||||
host_indices,
|
||||
device_indices,
|
||||
layer_id=2,
|
||||
io_backend="direct",
|
||||
)
|
||||
|
||||
self.assertEqual(len(calls), 1)
|
||||
src_ptrs, dst_ptrs, src_indices, dst_indices, kwargs = calls[0]
|
||||
self.assertEqual(len(src_ptrs), 1)
|
||||
self.assertEqual(src_ptrs[0].data_ptr(), device_pool.kv_buffer[2].data_ptr())
|
||||
self.assertEqual(len(dst_ptrs), 1)
|
||||
self.assertEqual(dst_ptrs[0].data_ptr(), host_pool.kv_buffer.data_ptr())
|
||||
self.assertEqual(src_indices.tolist(), [12, 13, 14, 15])
|
||||
self.assertEqual(dst_indices.tolist(), [4, 5, 6, 7])
|
||||
self.assertEqual(kwargs["layer_id"], 2)
|
||||
self.assertEqual(kwargs["page_size"], 4)
|
||||
|
||||
def test_mha_layer_page_first_per_layer_backup_uses_direct_lf_lpf(self):
|
||||
calls = []
|
||||
|
||||
def fake_direct(src_ptrs, dst_ptrs, src_indices, dst_indices, **kwargs):
|
||||
calls.append(
|
||||
(src_ptrs, dst_ptrs, src_indices.clone(), dst_indices.clone(), kwargs)
|
||||
)
|
||||
|
||||
host_pool = MHATokenToKVPoolHost.__new__(MHATokenToKVPoolHost)
|
||||
host_pool.layout = "layer_page_first"
|
||||
host_pool.page_size = 4
|
||||
host_pool.kv_buffer = torch.empty((2, 3, 8, 4, 2, 8), dtype=torch.uint8)
|
||||
device_pool = type("DevicePool", (), {})()
|
||||
device_pool.k_buffer = torch.empty((3, 32, 2, 8), dtype=torch.uint8)
|
||||
device_pool.v_buffer = torch.empty((3, 32, 2, 8), dtype=torch.uint8)
|
||||
host_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||
device_indices = torch.tensor([12, 13, 14, 15], dtype=torch.int64)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.mem_cache.memory_pool_host._load_tai_transfer_kv_per_layer_direct_lf_lpf",
|
||||
return_value=fake_direct,
|
||||
):
|
||||
host_pool.backup_from_device_per_layer(
|
||||
device_pool,
|
||||
host_indices,
|
||||
device_indices,
|
||||
layer_id=2,
|
||||
io_backend="direct",
|
||||
)
|
||||
|
||||
self.assertEqual(len(calls), 1)
|
||||
src_ptrs, dst_ptrs, src_indices, dst_indices, kwargs = calls[0]
|
||||
self.assertEqual(len(src_ptrs), 2)
|
||||
self.assertEqual(src_ptrs[0].data_ptr(), device_pool.k_buffer[2].data_ptr())
|
||||
self.assertEqual(src_ptrs[1].data_ptr(), device_pool.v_buffer[2].data_ptr())
|
||||
self.assertEqual(len(dst_ptrs), 2)
|
||||
self.assertEqual(dst_ptrs[0].data_ptr(), host_pool.k_buffer.data_ptr())
|
||||
self.assertEqual(dst_ptrs[1].data_ptr(), host_pool.v_buffer.data_ptr())
|
||||
self.assertEqual(src_indices.tolist(), [12, 13, 14, 15])
|
||||
self.assertEqual(dst_indices.tolist(), [4, 5, 6, 7])
|
||||
self.assertEqual(kwargs["layer_id"], 2)
|
||||
self.assertEqual(kwargs["page_size"], 4)
|
||||
|
||||
def test_nsa_indexer_layer_page_first_per_layer_backup_uses_direct_lf_lpf(self):
|
||||
calls = []
|
||||
|
||||
def fake_direct(src_ptrs, dst_ptrs, src_indices, dst_indices, **kwargs):
|
||||
calls.append(
|
||||
(src_ptrs, dst_ptrs, src_indices.clone(), dst_indices.clone(), kwargs)
|
||||
)
|
||||
|
||||
host_pool = NSATokenToKVPoolHost.__new__(NSATokenToKVPoolHost)
|
||||
host_pool.layout = "layer_page_first"
|
||||
host_pool.page_size = 4
|
||||
host_pool.index_k_with_scale_buffer = torch.empty(
|
||||
(3, 8, 1, 32), dtype=torch.uint8
|
||||
)
|
||||
device_pool = type("DevicePool", (), {})()
|
||||
device_pool.index_k_with_scale_buffer = torch.empty(
|
||||
(3, 8, 32), dtype=torch.uint8
|
||||
)
|
||||
host_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||
device_indices = torch.tensor([12, 13, 14, 15], dtype=torch.int64)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.mem_cache.memory_pool_host._load_tai_transfer_kv_per_layer_direct_lf_lpf",
|
||||
return_value=fake_direct,
|
||||
):
|
||||
host_pool._backup_indexer_from_device_per_layer(
|
||||
device_pool,
|
||||
host_indices,
|
||||
device_indices,
|
||||
layer_id=1,
|
||||
io_backend="direct",
|
||||
)
|
||||
|
||||
self.assertEqual(len(calls), 1)
|
||||
src_ptrs, dst_ptrs, src_indices, dst_indices, kwargs = calls[0]
|
||||
self.assertEqual(len(src_ptrs), 1)
|
||||
self.assertEqual(
|
||||
src_ptrs[0].data_ptr(),
|
||||
device_pool.index_k_with_scale_buffer[1].data_ptr(),
|
||||
)
|
||||
self.assertEqual(len(dst_ptrs), 1)
|
||||
self.assertEqual(
|
||||
dst_ptrs[0].data_ptr(),
|
||||
host_pool.index_k_with_scale_buffer.data_ptr(),
|
||||
)
|
||||
self.assertEqual(src_indices.tolist(), [3])
|
||||
self.assertEqual(dst_indices.tolist(), [1])
|
||||
self.assertEqual(kwargs["layer_id"], 1)
|
||||
self.assertEqual(kwargs["page_size"], 1)
|
||||
|
||||
def test_mla_page_first_direct_per_layer_load_uses_tai_direct_pf_lf(self):
|
||||
calls = []
|
||||
|
||||
@@ -582,6 +715,139 @@ class TestPageFirstPerLayerBackupTaiKernel(CustomTestCase):
|
||||
self.assertEqual(kwargs["layer_id"], 1)
|
||||
self.assertEqual(kwargs["page_size"], 1)
|
||||
|
||||
def test_mla_layer_page_first_per_layer_load_uses_tai_direct_lpf_lf(self):
|
||||
calls = []
|
||||
|
||||
def fake_direct(src_ptrs, dst_ptrs, src_indices, dst_indices, **kwargs):
|
||||
calls.append(
|
||||
(src_ptrs, dst_ptrs, src_indices.clone(), dst_indices.clone(), kwargs)
|
||||
)
|
||||
|
||||
host_pool = MLATokenToKVPoolHost.__new__(MLATokenToKVPoolHost)
|
||||
host_pool.layout = "layer_page_first"
|
||||
host_pool.page_size = 4
|
||||
host_pool.kv_buffer = torch.empty((3, 8, 4, 1, 16), dtype=torch.uint8)
|
||||
device_pool = type("DevicePool", (), {})()
|
||||
device_pool.kv_buffer = torch.empty((3, 32, 1, 16), dtype=torch.uint8)
|
||||
host_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||
device_indices = torch.tensor([12, 13, 14, 15], dtype=torch.int64)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.mem_cache.memory_pool_host._load_tai_transfer_kv_per_layer_direct_lpf_lf",
|
||||
return_value=fake_direct,
|
||||
):
|
||||
host_pool.load_to_device_per_layer(
|
||||
device_pool,
|
||||
host_indices,
|
||||
device_indices,
|
||||
layer_id=2,
|
||||
io_backend="direct",
|
||||
)
|
||||
|
||||
self.assertEqual(len(calls), 1)
|
||||
src_ptrs, dst_ptrs, src_indices, dst_indices, kwargs = calls[0]
|
||||
self.assertEqual(len(src_ptrs), 1)
|
||||
self.assertEqual(src_ptrs[0].data_ptr(), host_pool.kv_buffer.data_ptr())
|
||||
self.assertEqual(len(dst_ptrs), 1)
|
||||
self.assertEqual(dst_ptrs[0].data_ptr(), device_pool.kv_buffer[2].data_ptr())
|
||||
self.assertEqual(src_indices.tolist(), [4, 5, 6, 7])
|
||||
self.assertEqual(dst_indices.tolist(), [12, 13, 14, 15])
|
||||
self.assertEqual(kwargs["layer_id"], 2)
|
||||
self.assertEqual(kwargs["page_size"], 4)
|
||||
|
||||
def test_mha_layer_page_first_per_layer_load_uses_tai_direct_lpf_lf(self):
|
||||
calls = []
|
||||
|
||||
def fake_direct(src_ptrs, dst_ptrs, src_indices, dst_indices, **kwargs):
|
||||
calls.append(
|
||||
(src_ptrs, dst_ptrs, src_indices.clone(), dst_indices.clone(), kwargs)
|
||||
)
|
||||
|
||||
host_pool = MHATokenToKVPoolHost.__new__(MHATokenToKVPoolHost)
|
||||
host_pool.layout = "layer_page_first"
|
||||
host_pool.page_size = 4
|
||||
host_pool.kv_buffer = torch.empty((2, 3, 8, 4, 2, 8), dtype=torch.uint8)
|
||||
device_pool = type("DevicePool", (), {})()
|
||||
device_pool.k_buffer = torch.empty((3, 32, 2, 8), dtype=torch.uint8)
|
||||
device_pool.v_buffer = torch.empty((3, 32, 2, 8), dtype=torch.uint8)
|
||||
host_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||
device_indices = torch.tensor([12, 13, 14, 15], dtype=torch.int64)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.mem_cache.memory_pool_host._load_tai_transfer_kv_per_layer_direct_lpf_lf",
|
||||
return_value=fake_direct,
|
||||
):
|
||||
host_pool.load_to_device_per_layer(
|
||||
device_pool,
|
||||
host_indices,
|
||||
device_indices,
|
||||
layer_id=2,
|
||||
io_backend="direct",
|
||||
)
|
||||
|
||||
self.assertEqual(len(calls), 1)
|
||||
src_ptrs, dst_ptrs, src_indices, dst_indices, kwargs = calls[0]
|
||||
self.assertEqual(len(src_ptrs), 2)
|
||||
self.assertEqual(src_ptrs[0].data_ptr(), host_pool.k_buffer.data_ptr())
|
||||
self.assertEqual(src_ptrs[1].data_ptr(), host_pool.v_buffer.data_ptr())
|
||||
self.assertEqual(len(dst_ptrs), 2)
|
||||
self.assertEqual(dst_ptrs[0].data_ptr(), device_pool.k_buffer[2].data_ptr())
|
||||
self.assertEqual(dst_ptrs[1].data_ptr(), device_pool.v_buffer[2].data_ptr())
|
||||
self.assertEqual(src_indices.tolist(), [4, 5, 6, 7])
|
||||
self.assertEqual(dst_indices.tolist(), [12, 13, 14, 15])
|
||||
self.assertEqual(kwargs["layer_id"], 2)
|
||||
self.assertEqual(kwargs["page_size"], 4)
|
||||
|
||||
def test_nsa_indexer_layer_page_first_per_layer_load_uses_tai_direct_lpf_lf(self):
|
||||
calls = []
|
||||
|
||||
def fake_direct(src_ptrs, dst_ptrs, src_indices, dst_indices, **kwargs):
|
||||
calls.append(
|
||||
(src_ptrs, dst_ptrs, src_indices.clone(), dst_indices.clone(), kwargs)
|
||||
)
|
||||
|
||||
host_pool = NSATokenToKVPoolHost.__new__(NSATokenToKVPoolHost)
|
||||
host_pool.layout = "layer_page_first"
|
||||
host_pool.page_size = 4
|
||||
host_pool.index_k_with_scale_buffer = torch.empty(
|
||||
(3, 8, 1, 32), dtype=torch.uint8
|
||||
)
|
||||
device_pool = type("DevicePool", (), {})()
|
||||
device_pool.index_k_with_scale_buffer = torch.empty(
|
||||
(3, 8, 32), dtype=torch.uint8
|
||||
)
|
||||
host_indices = torch.tensor([4, 5, 6, 7], dtype=torch.int64)
|
||||
device_indices = torch.tensor([12, 13, 14, 15], dtype=torch.int64)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.mem_cache.memory_pool_host._load_tai_transfer_kv_per_layer_direct_lpf_lf",
|
||||
return_value=fake_direct,
|
||||
):
|
||||
host_pool._load_indexer_to_device_per_layer(
|
||||
device_pool,
|
||||
host_indices,
|
||||
device_indices,
|
||||
layer_id=1,
|
||||
io_backend="direct",
|
||||
)
|
||||
|
||||
self.assertEqual(len(calls), 1)
|
||||
src_ptrs, dst_ptrs, src_indices, dst_indices, kwargs = calls[0]
|
||||
self.assertEqual(len(src_ptrs), 1)
|
||||
self.assertEqual(
|
||||
src_ptrs[0].data_ptr(),
|
||||
host_pool.index_k_with_scale_buffer.data_ptr(),
|
||||
)
|
||||
self.assertEqual(len(dst_ptrs), 1)
|
||||
self.assertEqual(
|
||||
dst_ptrs[0].data_ptr(),
|
||||
device_pool.index_k_with_scale_buffer[1].data_ptr(),
|
||||
)
|
||||
self.assertEqual(src_indices.tolist(), [1])
|
||||
self.assertEqual(dst_indices.tolist(), [3])
|
||||
self.assertEqual(kwargs["layer_id"], 1)
|
||||
self.assertEqual(kwargs["page_size"], 1)
|
||||
|
||||
def test_nsa_indexer_load_reuses_precomputed_page_indices_across_layers(self):
|
||||
calls = []
|
||||
|
||||
@@ -1845,6 +2111,23 @@ class TestHiCacheControllerCPLoad(TestHiCacheControllerCPWrite):
|
||||
self.assertEqual(queued_op.host_indices.tolist(), [100, 101, 102, 103])
|
||||
self.assertEqual(queued_op.device_indices.tolist(), [20, 21, 22, 23])
|
||||
|
||||
def test_direct_layer_page_first_move_indices_keeps_host_order_and_cpu_device_indices(self):
|
||||
host_pool = FakeHostPool(torch.empty((0,), dtype=torch.int64))
|
||||
host_pool.layout = "layer_page_first"
|
||||
controller = self.make_controller(host_pool, cp_rank=1)
|
||||
op = CacheOperation(
|
||||
host_indices=torch.tensor([12, 8, 9, 10], dtype=torch.int64),
|
||||
device_indices=torch.tensor([32, 28, 29, 30], dtype=torch.int64),
|
||||
node_id=7,
|
||||
)
|
||||
|
||||
host_indices, device_indices = controller.move_indices(op, host_pool)
|
||||
|
||||
self.assertEqual(host_indices.tolist(), [12, 8, 9, 10])
|
||||
self.assertEqual(host_indices.device.type, "cpu")
|
||||
self.assertEqual(device_indices.tolist(), [32, 28, 29, 30])
|
||||
self.assertEqual(device_indices.device.type, "cpu")
|
||||
|
||||
def test_cp_load_rejects_non_contiguous_physical_device_page(self):
|
||||
host_pool = FakeHostPool(torch.tensor([100, 101, 102, 103], dtype=torch.int64))
|
||||
allocator = FakeAllocator(
|
||||
|
||||
@@ -17,6 +17,7 @@ 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.hicache_storage import HiCacheStorage
|
||||
from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool
|
||||
from sglang.srt.mem_cache.memory_pool_host import (
|
||||
ALLOC_MEMORY_FUNCS,
|
||||
@@ -32,6 +33,184 @@ from sglang.test.test_utils import CustomTestCase
|
||||
register_cuda_ci(est_time=3, suite="stage-b-test-1-gpu-small")
|
||||
|
||||
|
||||
class _DummyHiCacheStorage(HiCacheStorage):
|
||||
def get(self, key, target_location=None, target_sizes=None):
|
||||
return None
|
||||
|
||||
def batch_get(self, keys, target_locations=None, target_sizes=None):
|
||||
return []
|
||||
|
||||
def set(self, key, value=None, target_location=None, target_sizes=None):
|
||||
return False
|
||||
|
||||
def batch_set(self, keys, values=None, target_locations=None, target_sizes=None):
|
||||
return False
|
||||
|
||||
def exists(self, key):
|
||||
return False
|
||||
|
||||
|
||||
class TestLayerPageFirstDirectHostLayout(CustomTestCase):
|
||||
def test_mha_layer_page_first_direct_host_layout_is_layer_page_major(self):
|
||||
device_pool = SimpleNamespace(
|
||||
store_dtype=torch.float16,
|
||||
size=8,
|
||||
start_layer=0,
|
||||
end_layer=3,
|
||||
device="cpu",
|
||||
head_num=2,
|
||||
head_dim=8,
|
||||
layer_num=3,
|
||||
)
|
||||
|
||||
host_pool = MHATokenToKVPoolHost(
|
||||
device_pool=device_pool,
|
||||
host_to_device_ratio=2.0,
|
||||
host_size=0,
|
||||
page_size=4,
|
||||
layout="layer_page_first",
|
||||
pin_memory=False,
|
||||
device="cpu",
|
||||
host_token_capacity=16,
|
||||
)
|
||||
|
||||
self.assertEqual(tuple(host_pool.kv_buffer.shape), (2, 3, 4, 4, 2, 8))
|
||||
self.assertEqual(tuple(host_pool.k_buffer.shape), (3, 4, 4, 2, 8))
|
||||
self.assertEqual(len(host_pool.k_data_refs), 3)
|
||||
self.assertEqual(tuple(host_pool.k_data_refs[0].shape), (4, 4, 2, 8))
|
||||
|
||||
def test_mla_layer_page_first_direct_host_layout_is_layer_page_major(self):
|
||||
device_pool = SimpleNamespace(
|
||||
store_dtype=torch.float16,
|
||||
size=8,
|
||||
start_layer=0,
|
||||
end_layer=3,
|
||||
device="cpu",
|
||||
kv_lora_rank=16,
|
||||
qk_rope_head_dim=4,
|
||||
layer_num=3,
|
||||
)
|
||||
|
||||
host_pool = MLATokenToKVPoolHost(
|
||||
device_pool=device_pool,
|
||||
host_to_device_ratio=2.0,
|
||||
host_size=0,
|
||||
page_size=4,
|
||||
layout="layer_page_first",
|
||||
pin_memory=False,
|
||||
device="cpu",
|
||||
host_token_capacity=16,
|
||||
)
|
||||
|
||||
self.assertEqual(tuple(host_pool.kv_buffer.shape), (3, 4, 4, 1, 20))
|
||||
self.assertEqual(len(host_pool.data_refs), 3)
|
||||
self.assertEqual(tuple(host_pool.data_refs[0].shape), (4, 4, 1, 20))
|
||||
|
||||
def test_nsa_layer_page_first_direct_indexer_layout_is_layer_page_major(self):
|
||||
indexer_dtype = NSATokenToKVPool.index_k_with_scale_buffer_dtype
|
||||
device_pool = SimpleNamespace(
|
||||
store_dtype=torch.float16,
|
||||
size=8,
|
||||
start_layer=0,
|
||||
end_layer=3,
|
||||
device="cpu",
|
||||
kv_lora_rank=16,
|
||||
qk_rope_head_dim=4,
|
||||
kv_cache_dim=24,
|
||||
layer_num=3,
|
||||
index_head_dim=16,
|
||||
quant_block_size=8,
|
||||
index_k_with_scale_buffer=[
|
||||
torch.empty((5, 96), dtype=indexer_dtype) for _ in range(3)
|
||||
],
|
||||
)
|
||||
|
||||
host_pool = NSATokenToKVPoolHost(
|
||||
device_pool=device_pool,
|
||||
host_to_device_ratio=2.0,
|
||||
host_size=0,
|
||||
page_size=4,
|
||||
layout="layer_page_first",
|
||||
pin_memory=False,
|
||||
device="cpu",
|
||||
host_token_capacity=16,
|
||||
)
|
||||
|
||||
self.assertEqual(tuple(host_pool.kv_buffer.shape), (3, 4, 4, 1, 24))
|
||||
self.assertEqual(
|
||||
tuple(host_pool.index_k_with_scale_buffer.shape),
|
||||
(3, host_pool.indexer_page_num, 1, host_pool.indexer_page_stride_size),
|
||||
)
|
||||
self.assertEqual(len(host_pool.index_k_data_refs), 3)
|
||||
self.assertEqual(
|
||||
tuple(host_pool.index_k_data_refs[0].shape),
|
||||
(host_pool.indexer_page_num, 1, host_pool.indexer_page_stride_size),
|
||||
)
|
||||
|
||||
def test_mha_layer_page_first_page_buffer_meta_fails_fast_for_storage(self):
|
||||
device_pool = SimpleNamespace(
|
||||
store_dtype=torch.float16,
|
||||
size=8,
|
||||
start_layer=0,
|
||||
end_layer=2,
|
||||
device="cpu",
|
||||
head_num=2,
|
||||
head_dim=8,
|
||||
layer_num=2,
|
||||
)
|
||||
host_pool = MHATokenToKVPoolHost(
|
||||
device_pool=device_pool,
|
||||
host_to_device_ratio=2.0,
|
||||
host_size=0,
|
||||
page_size=4,
|
||||
layout="layer_page_first",
|
||||
pin_memory=False,
|
||||
device="cpu",
|
||||
host_token_capacity=16,
|
||||
)
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
RuntimeError, "layer_page_first_page_buffer_meta_unsupported"
|
||||
):
|
||||
host_pool.get_page_buffer_meta(torch.arange(0, 4, dtype=torch.int64))
|
||||
|
||||
def test_mla_layer_page_first_page_buffer_meta_fails_fast_for_storage(self):
|
||||
device_pool = SimpleNamespace(
|
||||
store_dtype=torch.float16,
|
||||
size=8,
|
||||
start_layer=0,
|
||||
end_layer=2,
|
||||
device="cpu",
|
||||
kv_lora_rank=16,
|
||||
qk_rope_head_dim=4,
|
||||
layer_num=2,
|
||||
)
|
||||
host_pool = MLATokenToKVPoolHost(
|
||||
device_pool=device_pool,
|
||||
host_to_device_ratio=2.0,
|
||||
host_size=0,
|
||||
page_size=4,
|
||||
layout="layer_page_first",
|
||||
pin_memory=False,
|
||||
device="cpu",
|
||||
host_token_capacity=16,
|
||||
)
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
RuntimeError, "layer_page_first_page_buffer_meta_unsupported"
|
||||
):
|
||||
host_pool.get_page_buffer_meta(torch.arange(0, 4, dtype=torch.int64))
|
||||
|
||||
def test_storage_registration_fails_fast_for_layer_page_first(self):
|
||||
storage = _DummyHiCacheStorage()
|
||||
mem_pool_host = SimpleNamespace(layout="layer_page_first")
|
||||
|
||||
with self.assertRaisesRegex(
|
||||
RuntimeError, "layer_page_first_storage_unsupported"
|
||||
):
|
||||
storage.register_mem_pool_host(mem_pool_host)
|
||||
|
||||
|
||||
class TestNSAHiCacheTransfer(CustomTestCase):
|
||||
def setUp(self):
|
||||
if not torch.cuda.is_available():
|
||||
@@ -1263,7 +1442,9 @@ class TestNSAHiCacheTransfer(CustomTestCase):
|
||||
io_backend="direct", layout="page_first_direct"
|
||||
)
|
||||
|
||||
def test_fp8_page_first_direct_roundtrip_preserves_kv_and_indexer_pages(self):
|
||||
def _run_direct_roundtrip_preserves_kv_and_indexer_pages(
|
||||
self, *, layout: str, dtype: torch.dtype
|
||||
):
|
||||
page_size = 64
|
||||
layer_num = 3
|
||||
size = page_size * 20
|
||||
@@ -1272,7 +1453,7 @@ class TestNSAHiCacheTransfer(CustomTestCase):
|
||||
size=size,
|
||||
page_size=page_size,
|
||||
kv_lora_rank=512,
|
||||
dtype=torch.float8_e4m3fn,
|
||||
dtype=dtype,
|
||||
qk_rope_head_dim=64,
|
||||
layer_num=layer_num,
|
||||
device="cuda",
|
||||
@@ -1285,7 +1466,7 @@ class TestNSAHiCacheTransfer(CustomTestCase):
|
||||
host_to_device_ratio=2.0,
|
||||
host_size=0,
|
||||
page_size=page_size,
|
||||
layout="page_first_direct",
|
||||
layout=layout,
|
||||
pin_memory=True,
|
||||
device="cpu",
|
||||
)
|
||||
@@ -1373,7 +1554,8 @@ class TestNSAHiCacheTransfer(CustomTestCase):
|
||||
]
|
||||
self.assertTrue(
|
||||
torch.equal(got_kv, expected_kv[layer_id][page_idx]),
|
||||
f"KV roundtrip mismatch layer={layer_id} dst_page={dst_page}",
|
||||
f"KV roundtrip mismatch layout={layout} dtype={dtype} "
|
||||
f"layer={layer_id} dst_page={dst_page}",
|
||||
)
|
||||
|
||||
got_index = device_pool.index_k_with_scale_buffer[layer_id][
|
||||
@@ -1381,9 +1563,24 @@ class TestNSAHiCacheTransfer(CustomTestCase):
|
||||
]
|
||||
self.assertTrue(
|
||||
torch.equal(got_index, expected_index[layer_id][page_idx]),
|
||||
f"index roundtrip mismatch layer={layer_id} dst_page={dst_page}",
|
||||
f"index roundtrip mismatch layout={layout} dtype={dtype} "
|
||||
f"layer={layer_id} dst_page={dst_page}",
|
||||
)
|
||||
|
||||
def test_fp8_page_first_direct_roundtrip_preserves_kv_and_indexer_pages(self):
|
||||
self._run_direct_roundtrip_preserves_kv_and_indexer_pages(
|
||||
layout="page_first_direct", dtype=torch.float8_e4m3fn
|
||||
)
|
||||
|
||||
def test_fp8_layer_page_first_roundtrip_preserves_kv_and_indexer_pages(self):
|
||||
self._run_direct_roundtrip_preserves_kv_and_indexer_pages(
|
||||
layout="layer_page_first", dtype=torch.float8_e4m3fn
|
||||
)
|
||||
|
||||
def test_bf16_layer_page_first_roundtrip_preserves_kv_and_indexer_pages(self):
|
||||
self._run_direct_roundtrip_preserves_kv_and_indexer_pages(
|
||||
layout="layer_page_first", dtype=torch.bfloat16
|
||||
)
|
||||
|
||||
class TestPageFirstDirectAllLayerBackupRoute(CustomTestCase):
|
||||
def test_mla_page_first_direct_all_layer_backup_uses_tai_per_layer_route(self):
|
||||
@@ -1576,6 +1773,126 @@ class TestNSAIndexerPageIndices(CustomTestCase):
|
||||
self.assertEqual(call["dst_indices"].tolist(), [0, 1])
|
||||
self.assertEqual(call["page_size"], 1)
|
||||
|
||||
def test_mla_layer_page_first_all_layer_backup_uses_tai_per_layer_route(self):
|
||||
host_pool = MLATokenToKVPoolHost.__new__(MLATokenToKVPoolHost)
|
||||
host_pool.layout = "layer_page_first"
|
||||
host_pool.page_size = 4
|
||||
host_pool.layer_num = 2
|
||||
host_pool.kv_buffer = "host-mla-layer-page-first"
|
||||
device_pool = type(
|
||||
"FakeDevicePool",
|
||||
(),
|
||||
{"kv_buffer": ["device-mla-layer-0", "device-mla-layer-1"]},
|
||||
)()
|
||||
calls = []
|
||||
|
||||
def fake_tai_transfer(**kwargs):
|
||||
calls.append(kwargs)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.mem_cache.memory_pool_host._load_tai_transfer_kv_per_layer_direct_lf_lpf",
|
||||
return_value=fake_tai_transfer,
|
||||
):
|
||||
host_pool.backup_from_device_all_layer(
|
||||
device_pool,
|
||||
torch.tensor([0, 1, 2, 3], dtype=torch.int64),
|
||||
torch.tensor([8, 9, 10, 11], dtype=torch.int64),
|
||||
"direct",
|
||||
)
|
||||
|
||||
self.assertEqual([call["layer_id"] for call in calls], [0, 1])
|
||||
for layer_id, call in enumerate(calls):
|
||||
self.assertEqual(call["src_ptrs"], [f"device-mla-layer-{layer_id}"])
|
||||
self.assertEqual(call["dst_ptrs"], ["host-mla-layer-page-first"])
|
||||
self.assertEqual(call["src_indices"].tolist(), [8, 9, 10, 11])
|
||||
self.assertEqual(call["dst_indices"].tolist(), [0, 1, 2, 3])
|
||||
self.assertEqual(call["page_size"], 4)
|
||||
|
||||
def test_mha_layer_page_first_all_layer_backup_uses_tai_per_layer_route(self):
|
||||
host_pool = MHATokenToKVPoolHost.__new__(MHATokenToKVPoolHost)
|
||||
host_pool.layout = "layer_page_first"
|
||||
host_pool.page_size = 4
|
||||
host_pool.layer_num = 2
|
||||
host_pool.kv_buffer = ["host-k-layer-page-first", "host-v-layer-page-first"]
|
||||
device_pool = type(
|
||||
"FakeDevicePool",
|
||||
(),
|
||||
{
|
||||
"k_buffer": ["device-k-layer-0", "device-k-layer-1"],
|
||||
"v_buffer": ["device-v-layer-0", "device-v-layer-1"],
|
||||
},
|
||||
)()
|
||||
calls = []
|
||||
|
||||
def fake_tai_transfer(**kwargs):
|
||||
calls.append(kwargs)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.mem_cache.memory_pool_host._load_tai_transfer_kv_per_layer_direct_lf_lpf",
|
||||
return_value=fake_tai_transfer,
|
||||
):
|
||||
host_pool.backup_from_device_all_layer(
|
||||
device_pool,
|
||||
torch.tensor([0, 1, 2, 3], dtype=torch.int64),
|
||||
torch.tensor([8, 9, 10, 11], dtype=torch.int64),
|
||||
"direct",
|
||||
)
|
||||
|
||||
self.assertEqual([call["layer_id"] for call in calls], [0, 1])
|
||||
for layer_id, call in enumerate(calls):
|
||||
self.assertEqual(
|
||||
call["src_ptrs"],
|
||||
[f"device-k-layer-{layer_id}", f"device-v-layer-{layer_id}"],
|
||||
)
|
||||
self.assertEqual(
|
||||
call["dst_ptrs"],
|
||||
["host-k-layer-page-first", "host-v-layer-page-first"],
|
||||
)
|
||||
self.assertEqual(call["src_indices"].tolist(), [8, 9, 10, 11])
|
||||
self.assertEqual(call["dst_indices"].tolist(), [0, 1, 2, 3])
|
||||
self.assertEqual(call["page_size"], 4)
|
||||
|
||||
def test_layer_page_first_all_layer_indexer_backup_uses_tai_per_layer_route(self):
|
||||
host_pool = self.make_host_pool_stub(page_size=4)
|
||||
host_pool.layout = "layer_page_first"
|
||||
host_pool.indexer_page_stride_size = 8
|
||||
host_pool.layer_num = 3
|
||||
host_pool.index_k_with_scale_buffer = "host-layer-page-first-indexer"
|
||||
device_pool = type(
|
||||
"FakeDevicePool",
|
||||
(),
|
||||
{
|
||||
"index_k_with_scale_buffer": [
|
||||
"device-layer-0",
|
||||
"device-layer-1",
|
||||
"device-layer-2",
|
||||
]
|
||||
},
|
||||
)()
|
||||
calls = []
|
||||
|
||||
def fake_tai_transfer(**kwargs):
|
||||
calls.append(kwargs)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.mem_cache.memory_pool_host._load_tai_transfer_kv_per_layer_direct_lf_lpf",
|
||||
return_value=fake_tai_transfer,
|
||||
):
|
||||
host_pool._backup_indexer_from_device_all_layer(
|
||||
device_pool,
|
||||
torch.tensor([0, 1, 2, 3, 4, 5, 6, 7], dtype=torch.int64),
|
||||
torch.tensor([8, 9, 10, 11, 12, 13, 14, 15], dtype=torch.int64),
|
||||
"direct",
|
||||
)
|
||||
|
||||
self.assertEqual([call["layer_id"] for call in calls], [0, 1, 2])
|
||||
for layer_id, call in enumerate(calls):
|
||||
self.assertEqual(call["src_ptrs"], [f"device-layer-{layer_id}"])
|
||||
self.assertEqual(call["dst_ptrs"], ["host-layer-page-first-indexer"])
|
||||
self.assertEqual(call["src_indices"].tolist(), [2, 3])
|
||||
self.assertEqual(call["dst_indices"].tolist(), [0, 1])
|
||||
self.assertEqual(call["page_size"], 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -75,6 +75,23 @@ def test_cp_shared_kv_prefill_bs_gt1_parser_limits():
|
||||
assert args.cp_shared_kv_prefill_max_total_extend_tokens == 8192
|
||||
|
||||
|
||||
def test_hicache_mem_layout_parser_accepts_layer_page_first():
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
ServerArgs.add_cli_args(parser)
|
||||
raw_args = parser.parse_args(
|
||||
[
|
||||
"--model-path",
|
||||
"dummy",
|
||||
"--hicache-mem-layout",
|
||||
"layer_page_first",
|
||||
]
|
||||
)
|
||||
args = ServerArgs.from_cli_args(raw_args)
|
||||
assert args.hicache_mem_layout == "layer_page_first"
|
||||
|
||||
|
||||
class TestLoadBalanceMethod(unittest.TestCase):
|
||||
def test_non_pd_defaults_to_round_robin(self):
|
||||
server_args = ServerArgs(model_path="dummy", disaggregation_mode="null")
|
||||
@@ -556,6 +573,7 @@ class TestHiCacheArgs(CustomTestCase):
|
||||
("kernel", "page_first"),
|
||||
("direct", "layer_first"),
|
||||
("direct", "page_first_direct"),
|
||||
("direct", "layer_page_first"),
|
||||
]
|
||||
|
||||
for io_backend, mem_layout in cases:
|
||||
@@ -598,6 +616,26 @@ class TestHiCacheArgs(CustomTestCase):
|
||||
hicache_mem_layout="page_first_kv_split",
|
||||
)
|
||||
|
||||
def test_cp_hicache_rejects_kernel_layer_page_first_layout(self):
|
||||
with self.assertRaisesRegex(
|
||||
ValueError, "CP shared KV HiCache.*kernel.*layer_page_first"
|
||||
):
|
||||
self._normalize_and_validate_cp_hicache_args(
|
||||
hicache_io_backend="kernel",
|
||||
hicache_mem_layout="layer_page_first",
|
||||
)
|
||||
|
||||
def test_hicache_storage_rejects_layer_page_first_layout(self):
|
||||
args = self._make_args(
|
||||
enable_hierarchical_cache=True,
|
||||
hicache_storage_backend="mooncake",
|
||||
hicache_io_backend="direct",
|
||||
hicache_mem_layout="layer_page_first",
|
||||
)
|
||||
|
||||
with self.assertRaisesRegex(ValueError, "layer_page_first.*storage"):
|
||||
args._handle_hicache()
|
||||
|
||||
def test_cp_hicache_rejects_kernel_ascend_backend(self):
|
||||
with self.assertRaisesRegex(ValueError, "CP shared KV HiCache.*kernel_ascend"):
|
||||
self._normalize_and_validate_cp_hicache_args(
|
||||
|
||||
Reference in New Issue
Block a user