[unified-memory] PD disaggregation for every unified pool shape (#37506)
This commit is contained in:
@@ -18,9 +18,11 @@ import unittest
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.layout.page_major import (
|
||||
build_mha_views,
|
||||
build_mla_views,
|
||||
build_page_major_mamba_views,
|
||||
mamba_entry_bytes,
|
||||
mha_entry_bytes,
|
||||
mla_entry_bytes,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
@@ -74,6 +76,118 @@ class TestMLAEnvelopeTransferAddressing(CustomTestCase):
|
||||
self.assertTrue(torch.equal(got, val), (page, layer, off))
|
||||
|
||||
|
||||
class TestMHAEnvelopeTransferAddressing(CustomTestCase):
|
||||
"""The MHA counterpart of the MLA case above.
|
||||
|
||||
An MHA page envelope holds ``2 * layer_num`` row-blocks (layer l's K at
|
||||
block 2l, its V at 2l+1). PD ships that whole envelope as one item, so a
|
||||
row written through ANY per-layer view must land inside its own page's
|
||||
``page_envelope_bytes`` block -- otherwise the transfer would carry a
|
||||
page's K but another page's V and every kernel would still read fine
|
||||
locally.
|
||||
"""
|
||||
|
||||
def test_page_envelope_matches_per_layer_views(self):
|
||||
layer_num, page_size, head_num, head_dim, num_pages = 3, 4, 2, 8, 6
|
||||
store_dtype = torch.bfloat16
|
||||
entry_bytes = mha_entry_bytes(
|
||||
layer_num=layer_num,
|
||||
head_num=head_num,
|
||||
head_dim=head_dim,
|
||||
v_head_dim=head_dim,
|
||||
itemsize=store_dtype.itemsize,
|
||||
)
|
||||
page_bytes = page_size * entry_bytes
|
||||
row_bytes = head_num * head_dim * store_dtype.itemsize
|
||||
self.assertEqual(page_bytes, page_size * 2 * layer_num * row_bytes)
|
||||
|
||||
# One page envelope of tail pad, as UnifiedKVPool allocates for MHA.
|
||||
raw = torch.zeros((num_pages + 1) * page_bytes, dtype=torch.uint8)
|
||||
k_views, v_views = build_mha_views(
|
||||
raw,
|
||||
layer_num=layer_num,
|
||||
head_num=head_num,
|
||||
head_dim=head_dim,
|
||||
v_head_dim=head_dim,
|
||||
store_dtype=store_dtype,
|
||||
page_size=page_size,
|
||||
num_pages=num_pages,
|
||||
anchor_bytes=0,
|
||||
)
|
||||
|
||||
blocks = 2 * layer_num
|
||||
for page in range(num_pages):
|
||||
for layer in range(layer_num):
|
||||
for is_v, views in ((0, k_views), (1, v_views)):
|
||||
for pos in range(page_size):
|
||||
row = page * blocks * page_size + pos
|
||||
views[layer][row].fill_(1)
|
||||
(nz,) = torch.nonzero(raw, as_tuple=True)
|
||||
lo, hi = int(nz.min()), int(nz.max())
|
||||
self.assertGreaterEqual(
|
||||
lo,
|
||||
page * page_bytes,
|
||||
f"page={page} layer={layer} v={is_v} pos={pos} "
|
||||
"wrote below its page envelope",
|
||||
)
|
||||
self.assertLess(
|
||||
hi,
|
||||
(page + 1) * page_bytes,
|
||||
f"page={page} layer={layer} v={is_v} pos={pos} "
|
||||
"wrote past its page envelope",
|
||||
)
|
||||
views[layer][row].zero_()
|
||||
|
||||
def test_envelope_move_is_a_whole_page_copy(self):
|
||||
"""Relocating a page envelope must move every layer's K and V with it;
|
||||
this is what `UnifiedMHATokenToKVPool.move_kv_cache` relies on and what
|
||||
makes a physical page id a valid PD transfer index after compaction."""
|
||||
layer_num, page_size, head_num, head_dim, num_pages = 2, 2, 1, 4, 4
|
||||
store_dtype = torch.bfloat16
|
||||
entry_bytes = mha_entry_bytes(
|
||||
layer_num=layer_num,
|
||||
head_num=head_num,
|
||||
head_dim=head_dim,
|
||||
v_head_dim=head_dim,
|
||||
itemsize=store_dtype.itemsize,
|
||||
)
|
||||
page_bytes = page_size * entry_bytes
|
||||
raw = torch.zeros((num_pages + 1) * page_bytes, dtype=torch.uint8)
|
||||
k_views, v_views = build_mha_views(
|
||||
raw,
|
||||
layer_num=layer_num,
|
||||
head_num=head_num,
|
||||
head_dim=head_dim,
|
||||
v_head_dim=head_dim,
|
||||
store_dtype=store_dtype,
|
||||
page_size=page_size,
|
||||
num_pages=num_pages,
|
||||
anchor_bytes=0,
|
||||
)
|
||||
blocks = 2 * layer_num
|
||||
# Distinct content in source page 1, every layer, K and V.
|
||||
for layer in range(layer_num):
|
||||
for pos in range(page_size):
|
||||
row = 1 * blocks * page_size + pos
|
||||
k_views[layer][row].fill_(layer + 1)
|
||||
v_views[layer][row].fill_(-(layer + 1))
|
||||
|
||||
env = raw[: num_pages * page_bytes].view(num_pages, page_bytes)
|
||||
env[3] = env[1]
|
||||
|
||||
for layer in range(layer_num):
|
||||
for pos in range(page_size):
|
||||
row = 3 * blocks * page_size + pos
|
||||
self.assertTrue(
|
||||
torch.all(k_views[layer][row] == layer + 1),
|
||||
f"K layer {layer} did not ride the envelope move",
|
||||
)
|
||||
self.assertTrue(
|
||||
torch.all(v_views[layer][row] == -(layer + 1)),
|
||||
f"V layer {layer} did not ride the envelope move",
|
||||
)
|
||||
|
||||
|
||||
class TestMambaEnvelopeTransferAddressing(CustomTestCase):
|
||||
def test_slot_envelope_is_self_contained(self):
|
||||
"""A slot's conv+temporal state for all layers must live exactly in
|
||||
|
||||
@@ -345,11 +345,12 @@ class TestUnifiedMHATokenToKVPool(unittest.TestCase):
|
||||
)
|
||||
|
||||
def test_transfer_entry_points_fail_loud(self):
|
||||
"""PD / CPU-copy entry points assume per-layer buffers indexed by TOKEN
|
||||
id and would silently mis-index the row space, so each must raise."""
|
||||
"""The entry points that assume per-layer buffers indexed by TOKEN id
|
||||
would silently mis-index against the row space (or hit a missing-attr
|
||||
AttributeError), so each must raise. `get_contiguous_buf_infos` is NOT
|
||||
among them: PD addresses this pool as whole page envelopes, pinned by
|
||||
`test_pd_registration_is_one_whole_envelope` below."""
|
||||
_, pool = _make_pool_and_kv(1)
|
||||
with self.assertRaises(NotImplementedError):
|
||||
pool.get_contiguous_buf_infos()
|
||||
with self.assertRaises(NotImplementedError):
|
||||
pool.get_cpu_copy(torch.tensor([1]))
|
||||
with self.assertRaises(NotImplementedError):
|
||||
@@ -357,6 +358,24 @@ class TestUnifiedMHATokenToKVPool(unittest.TestCase):
|
||||
with self.assertRaises(NotImplementedError):
|
||||
pool.set_kv_buffer_prefix_valid()
|
||||
|
||||
def test_pd_registration_is_one_whole_envelope(self):
|
||||
"""PD registers ONE region -- the whole raw buffer -- with the page
|
||||
envelope as the item, so the transfer engine addresses it as
|
||||
`raw_ptr + physical_page * page_envelope_bytes`. Per-layer regions
|
||||
would be wrong here: the per-layer views overlap inside the envelope
|
||||
and index in kernel-facing ids, not token ids."""
|
||||
kv, pool = _make_pool_and_kv(1)
|
||||
ptrs, lens, item_lens = pool.get_contiguous_buf_infos()
|
||||
self.assertEqual(len(ptrs), 1)
|
||||
self.assertEqual(len(lens), 1)
|
||||
self.assertEqual(len(item_lens), 1)
|
||||
self.assertEqual(ptrs[0], kv._raw.data_ptr())
|
||||
self.assertEqual(lens[0], kv._raw.numel())
|
||||
self.assertEqual(item_lens[0], pool._page_bytes)
|
||||
# The whole addressable page range must fit the registered region, or
|
||||
# the last page's write would run off the end of the RDMA mapping.
|
||||
self.assertLessEqual(pool._num_pages * item_lens[0], lens[0])
|
||||
|
||||
def test_hnd_env_cannot_hijack_layout(self):
|
||||
"""SGLANG_USE_HND_KVCACHE must not flip this pool's layout: HND indexes
|
||||
4-D while the per-layer views are 3-D, so the pinned label has to win."""
|
||||
|
||||
Reference in New Issue
Block a user