[EPD] Batch embedding cache host-device range copies (#31574)

Co-authored-by: 晟海 <huangtingwei.htw@antgroup.com>
Co-authored-by: liusy58 <liusy58@linux.alibaba.com>
Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
This commit is contained in:
Yuang Chen
2026-08-14 22:01:12 +08:00
committed by GitHub
co-authored by 晟海 liusy58 Xiaoyu Zhang
parent 5e65dd01a7
commit 1a178f7c7c
7 changed files with 364 additions and 34 deletions
@@ -16,7 +16,8 @@ from sglang.srt.mem_cache.embedding_cache_controller import (
EvictableLRU,
PageRun,
RangePageAllocator,
build_transfer_buffers,
_build_host_device_transfer_plan,
_build_storage_transfer_buffers,
)
from sglang.test.ci.ci_register import register_cpu_ci
@@ -361,7 +362,7 @@ class TestTransferBuffers(unittest.TestCase):
state=EntryState.READY,
)
ptrs, sizes = build_transfer_buffers(entry, pool)
ptrs, sizes = _build_storage_transfer_buffers(entry, pool)
self.assertEqual(ptrs, [pool.tensor[4].data_ptr()])
self.assertEqual(sizes, [5 * 4 * torch.float32.itemsize])
@@ -377,7 +378,7 @@ class TestTransferBuffers(unittest.TestCase):
state=EntryState.READY,
)
ptrs, sizes = build_transfer_buffers(entry, pool)
ptrs, sizes = _build_storage_transfer_buffers(entry, pool)
self.assertEqual(ptrs, [pool.tensor[0].data_ptr(), pool.tensor[6].data_ptr()])
self.assertEqual(
@@ -388,6 +389,38 @@ class TestTransferBuffers(unittest.TestCase):
],
)
def test_build_h2d_copy_plan_for_fragmented_entry(self):
pool = _make_pool(num_pages=10, dim=4, page_size=2)
entry = EmbeddingCacheEntry(
hash="h",
modality=Modality.IMAGE,
num_tokens=5,
dim=4,
page_runs=[PageRun(2, 1), PageRun(7, 2)],
state=EntryState.READY,
)
plan = _build_host_device_transfer_plan(
entry, pool, src_is_pool=True, dst_token_offset=3
)
self.assertEqual(plan, ([4, 14], [3, 5], [2, 3]))
def test_build_d2h_copy_plan_for_fragmented_entry(self):
pool = _make_pool(num_pages=10, dim=4, page_size=2)
entry = EmbeddingCacheEntry(
hash="h",
modality=Modality.IMAGE,
num_tokens=5,
dim=4,
page_runs=[PageRun(2, 1), PageRun(7, 2)],
state=EntryState.READY,
)
plan = _build_host_device_transfer_plan(entry, pool, src_is_pool=False)
self.assertEqual(plan, ([0, 2], [4, 14], [2, 3]))
class TestMooncakeEmbeddingStoreWrappers(unittest.TestCase):
def test_batch_put_multi_buffers_deduplicates_existing_keys(self):