From 1a178f7c7c21ce35d2250a78a60e05e3d973adca Mon Sep 17 00:00:00 2001 From: Yuang Chen <77919385+cccccya@users.noreply.github.com> Date: Fri, 14 Aug 2026 22:01:12 +0800 Subject: [PATCH] [EPD] Batch embedding cache host-device range copies (#31574) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: 晟海 Co-authored-by: liusy58 Co-authored-by: Xiaoyu Zhang <1182563586@qq.com> --- .../kernels/aot/csrc/common_extension.cc | 4 + .../kernels/aot/csrc/kvcacheio/transfer.cu | 152 ++++++++++++++++++ .../kernels/aot/include/sgl_kernel_ops.h | 7 + .../aot/python/sgl_kernel/kvcacheio.py | 13 ++ .../kernels/aot/tests/test_kvcacheio.py | 60 +++++++ .../mem_cache/embedding_cache_controller.py | 123 ++++++++++---- .../test_embedding_cache_controller.py | 39 ++++- 7 files changed, 364 insertions(+), 34 deletions(-) diff --git a/python/sglang/kernels/aot/csrc/common_extension.cc b/python/sglang/kernels/aot/csrc/common_extension.cc index 4f6665cf0..39d9a6a2c 100644 --- a/python/sglang/kernels/aot/csrc/common_extension.cc +++ b/python/sglang/kernels/aot/csrc/common_extension.cc @@ -325,6 +325,10 @@ TORCH_LIBRARY_FRAGMENT(sgl_kernel, m) { "transfer_kv_direct(Tensor[] src_layers, Tensor[] dst_layers, Tensor src_indices, Tensor dst_indices, int " "page_size) -> ()"); m.impl("transfer_kv_direct", torch::kCUDA, &transfer_kv_direct); + m.def( + "transfer_embedding_ranges_direct(Tensor src, Tensor! dst, int[] src_starts, int[] dst_starts, int[] " + "lengths) -> ()"); + m.impl("transfer_embedding_ranges_direct", torch::kCUDA, &transfer_embedding_ranges_direct); m.def( "transfer_kv_per_layer_direct_pf_lf(Tensor[] src_ptrs, Tensor[] dst_ptrs, Tensor src_indices, " "Tensor dst_indices, int layer_id, int page_size)->() "); diff --git a/python/sglang/kernels/aot/csrc/kvcacheio/transfer.cu b/python/sglang/kernels/aot/csrc/kvcacheio/transfer.cu index b4372a956..248571089 100644 --- a/python/sglang/kernels/aot/csrc/kvcacheio/transfer.cu +++ b/python/sglang/kernels/aot/csrc/kvcacheio/transfer.cu @@ -1,5 +1,6 @@ #include #include +#include #include #include @@ -732,6 +733,157 @@ void transfer_kv_direct( } } +void transfer_embedding_ranges_direct( + const at::Tensor& src, + at::Tensor& dst, + const std::vector& src_starts, + const std::vector& dst_starts, + const std::vector& lengths) { + TORCH_CHECK(src.dim() == 2, "Source embedding tensor must be 2D"); + TORCH_CHECK(dst.dim() == 2, "Destination embedding tensor must be 2D"); + TORCH_CHECK(src.scalar_type() == dst.scalar_type(), "Source and destination dtypes must match"); + TORCH_CHECK(src.size(1) == dst.size(1), "Source and destination embedding dims must match"); + TORCH_CHECK(src.is_contiguous() && dst.is_contiguous(), "Embedding tensors must be contiguous"); + TORCH_CHECK(src.is_cuda() != dst.is_cuda(), "Exactly one embedding tensor must be on CUDA"); + TORCH_CHECK(src_starts.size() == dst_starts.size(), "src_starts and dst_starts must have the same length"); + TORCH_CHECK(src_starts.size() == lengths.size(), "src_starts and lengths must have the same length"); + + const auto num_ranges = lengths.size(); + if (num_ranges == 0) { + return; + } + + const auto copy_device = src.is_cuda() ? src.device() : dst.device(); + const at::cuda::OptionalCUDAGuard device_guard(copy_device); + const cudaStream_t stream = at::cuda::getCurrentCUDAStream(); + + const size_t row_bytes = static_cast(src.size(1)) * src.element_size(); + const char* src_base = static_cast(src.data_ptr()); + char* dst_base = static_cast(dst.data_ptr()); + + thread_local std::vector batch_srcs; + thread_local std::vector batch_dsts; + thread_local std::vector batch_sizes; + batch_srcs.clear(); + batch_dsts.clear(); + batch_sizes.clear(); + batch_srcs.reserve(num_ranges); + batch_dsts.reserve(num_ranges); + batch_sizes.reserve(num_ranges); + + // Validate the complete plan before submitting any asynchronous copy so a + // bad later range cannot leave the destination partially updated. + for (size_t i = 0; i < num_ranges; ++i) { + const int64_t src_start = src_starts[i]; + const int64_t dst_start = dst_starts[i]; + const int64_t length = lengths[i]; + + TORCH_CHECK(length >= 0, "Range length must be non-negative"); + if (length == 0) { + continue; + } + TORCH_CHECK(src_start >= 0, "Source range start must be non-negative"); + TORCH_CHECK(dst_start >= 0, "Destination range start must be non-negative"); + TORCH_CHECK(length <= src.size(0) - src_start, "Source range is out of bounds"); + TORCH_CHECK(length <= dst.size(0) - dst_start, "Destination range is out of bounds"); + + batch_srcs.push_back(const_cast(src_base + static_cast(src_start) * row_bytes)); + batch_dsts.push_back(dst_base + static_cast(dst_start) * row_bytes); + batch_sizes.push_back(static_cast(length) * row_bytes); + } + + const auto fallback_to_async_copies = [&]() { + for (size_t i = 0; i < batch_sizes.size(); ++i) { + C10_CUDA_CHECK(cudaMemcpyAsync(batch_dsts[i], batch_srcs[i], batch_sizes[i], cudaMemcpyDefault, stream)); + } + }; + + if (batch_sizes.empty()) { + return; + } + +#if defined(USE_ROCM) || defined(USE_MUSA) || !defined(CUDA_VERSION) || CUDA_VERSION < 12080 + fallback_to_async_copies(); + return; +#else + // cudaMemcpyBatchAsync rejects the legacy NULL stream. + if (stream == nullptr) { + fallback_to_async_copies(); + return; + } + + int driver_version = 0; + const cudaError_t driver_version_err = cudaDriverGetVersion(&driver_version); + if (driver_version_err != cudaSuccess || driver_version < 12080) { + fallback_to_async_copies(); + return; + } + + static void* cuda_memcpy_batch_async_sym = dlsym(RTLD_DEFAULT, "cudaMemcpyBatchAsync"); + if (cuda_memcpy_batch_async_sym == nullptr) { + fallback_to_async_copies(); + return; + } + + static int runtime_version = 0; + static const cudaError_t runtime_version_err = cudaRuntimeGetVersion(&runtime_version); + if (runtime_version_err != cudaSuccess) { + fallback_to_async_copies(); + return; + } + static const bool use_v13_signature = runtime_version >= 13000; + + const int device_id = copy_device.index(); + std::vector attrs_idxs(1, 0); + cudaMemcpyAttributes attrs{}; + attrs.srcAccessOrder = cudaMemcpySrcAccessOrderStream; + attrs.srcLocHint.type = src.is_cuda() ? cudaMemLocationTypeDevice : cudaMemLocationTypeHost; + attrs.srcLocHint.id = src.is_cuda() ? device_id : 0; + attrs.dstLocHint.type = dst.is_cuda() ? cudaMemLocationTypeDevice : cudaMemLocationTypeHost; + attrs.dstLocHint.id = dst.is_cuda() ? device_id : 0; + attrs.flags = 0; + + cudaError_t err; + size_t fail_idx = std::numeric_limits::max(); + if (use_v13_signature) { + using FnV13 = cudaError_t (*)( + void* const*, const void* const*, const size_t*, size_t, cudaMemcpyAttributes*, size_t*, size_t, cudaStream_t); + auto fn = reinterpret_cast(cuda_memcpy_batch_async_sym); + err = + fn(batch_dsts.data(), + batch_srcs.data(), + batch_sizes.data(), + batch_sizes.size(), + &attrs, + attrs_idxs.data(), + 1, + stream); + } else { + using FnV12 = + cudaError_t (*)(void**, void**, size_t*, size_t, cudaMemcpyAttributes*, size_t*, size_t, size_t*, cudaStream_t); + auto fn = reinterpret_cast(cuda_memcpy_batch_async_sym); + err = + fn(batch_dsts.data(), + batch_srcs.data(), + batch_sizes.data(), + batch_sizes.size(), + &attrs, + attrs_idxs.data(), + 1, + &fail_idx, + stream); + } + + if (err == cudaErrorNotSupported || err == cudaErrorCallRequiresNewerDriver) { + (void)cudaGetLastError(); + fallback_to_async_copies(); + return; + } + TORCH_CHECK( + err == cudaSuccess, "cudaMemcpyBatchAsync failed. failIdx=", fail_idx, " error=", cudaGetErrorString(err)); +#endif +} + template inline void transfer_kv_page_first_direct_impl( const std::vector& src_ptrs, diff --git a/python/sglang/kernels/aot/include/sgl_kernel_ops.h b/python/sglang/kernels/aot/include/sgl_kernel_ops.h index 69dd943ea..4ee916228 100644 --- a/python/sglang/kernels/aot/include/sgl_kernel_ops.h +++ b/python/sglang/kernels/aot/include/sgl_kernel_ops.h @@ -607,6 +607,13 @@ void transfer_kv_direct( const at::Tensor dst_indices, int64_t page_size); +void transfer_embedding_ranges_direct( + const at::Tensor& src, + at::Tensor& dst, + const std::vector& src_starts, + const std::vector& dst_starts, + const std::vector& lengths); + void transfer_kv_per_layer_direct_pf_lf( const std::vector& src_ptrs, std::vector dst_ptrs, diff --git a/python/sglang/kernels/aot/python/sgl_kernel/kvcacheio.py b/python/sglang/kernels/aot/python/sgl_kernel/kvcacheio.py index 090142a98..0601c8fe6 100644 --- a/python/sglang/kernels/aot/python/sgl_kernel/kvcacheio.py +++ b/python/sglang/kernels/aot/python/sgl_kernel/kvcacheio.py @@ -192,6 +192,19 @@ def transfer_kv_direct( ) +def transfer_embedding_ranges_direct( + src: torch.Tensor, + dst: torch.Tensor, + src_starts: List[int], + dst_starts: List[int], + lengths: List[int], +) -> None: + """Copy embedding ranges between host and CUDA tensors.""" + torch.ops.sgl_kernel.transfer_embedding_ranges_direct.default( + src, dst, src_starts, dst_starts, lengths + ) + + def transfer_kv_per_layer_direct_pf_lf( src_ptrs: List[torch.Tensor], dst_ptrs: List[torch.Tensor], diff --git a/python/sglang/kernels/aot/tests/test_kvcacheio.py b/python/sglang/kernels/aot/tests/test_kvcacheio.py index f28d0fc99..c88412f90 100644 --- a/python/sglang/kernels/aot/tests/test_kvcacheio.py +++ b/python/sglang/kernels/aot/tests/test_kvcacheio.py @@ -3,6 +3,7 @@ import sys import pytest import torch from sgl_kernel.kvcacheio import ( + transfer_embedding_ranges_direct, transfer_kv_all_layer, transfer_kv_all_layer_direct_lf_pf, transfer_kv_all_layer_lf_ph, @@ -68,6 +69,65 @@ def ref_copy_with_indices_page_head( ][head_id][src_indices[i] % page_size][layer_id].to(dst_pool.device) +def ref_copy_embedding_ranges(src, dst, src_starts, dst_starts, lengths): + for src_start, dst_start, length in zip(src_starts, dst_starts, lengths): + dst[dst_start : dst_start + length].copy_( + src[src_start : src_start + length], non_blocking=True + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +@pytest.mark.skipif(is_hip(), reason="This test covers the CUDA batch-copy op") +@pytest.mark.parametrize("direction", ["h2d", "d2h"]) +def test_transfer_embedding_ranges_direct(direction: str): + dtype = torch.bfloat16 + embedding_dim = 37 + page_size = 4 + fragmented_starts = [1, 11, 23] + contiguous_starts = [2, 6, 10] + lengths = [page_size, page_size, 2] + host_rows = 28 + device_rows = 16 + + host_values = torch.arange(host_rows * embedding_dim, dtype=torch.float32).reshape( + host_rows, embedding_dim + ) + device_values = torch.arange( + device_rows * embedding_dim, dtype=torch.float32 + ).reshape(device_rows, embedding_dim) + + if direction == "h2d": + src = host_values.to(dtype).pin_memory() + direct_dst = torch.full( + (device_rows, embedding_dim), -1, dtype=dtype, device="cuda" + ) + reference_dst = torch.full_like(direct_dst, -1) + src_starts, dst_starts = fragmented_starts, contiguous_starts + else: + src = device_values.to(dtype).to("cuda") + direct_dst = torch.full( + (host_rows, embedding_dim), -1, dtype=dtype, pin_memory=True + ) + reference_dst = torch.full( + (host_rows, embedding_dim), -1, dtype=dtype, pin_memory=True + ) + src_starts, dst_starts = contiguous_starts, fragmented_starts + + torch.cuda.synchronize() + copy_stream = torch.cuda.Stream() + assert copy_stream.cuda_stream != torch.cuda.default_stream().cuda_stream + with torch.cuda.stream(copy_stream): + ref_copy_embedding_ranges(src, reference_dst, src_starts, dst_starts, lengths) + transfer_embedding_ranges_direct( + src, direct_dst, src_starts, dst_starts, lengths + ) + completion_event = torch.cuda.Event() + completion_event.record(copy_stream) + + completion_event.synchronize() + torch.testing.assert_close(direct_dst, reference_dst) + + @pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) @pytest.mark.parametrize("num_items_to_transfer", [1, 128, 1024]) @pytest.mark.parametrize("page_size", [1, 16, 64]) diff --git a/python/sglang/srt/mem_cache/embedding_cache_controller.py b/python/sglang/srt/mem_cache/embedding_cache_controller.py index 88abf1437..be912813b 100644 --- a/python/sglang/srt/mem_cache/embedding_cache_controller.py +++ b/python/sglang/srt/mem_cache/embedding_cache_controller.py @@ -11,11 +11,21 @@ from typing import List, Optional, Tuple import torch +try: + from sgl_kernel.kvcacheio import transfer_embedding_ranges_direct +except ImportError: + transfer_embedding_ranges_direct = None + from sglang.srt.managers.schedule_batch import Modality from sglang.srt.mem_cache.embedding_store import EmbeddingStore logger = logging.getLogger(__name__) +if transfer_embedding_ranges_direct is not None and not hasattr( + torch.ops.sgl_kernel, "transfer_embedding_ranges_direct" +): + transfer_embedding_ranges_direct = None + TARGET_PAGE_BYTES = 256 * 1024 VISION_POOL_RATIO = 0.8 @@ -226,10 +236,10 @@ class EmbeddingCacheEntry: return self.state == EntryState.READY and self.ref_count == 0 -def build_transfer_buffers( +def _build_storage_transfer_buffers( entry: EmbeddingCacheEntry, pool: EmbeddingPool ) -> Tuple[List[int], List[int]]: - """Build one pointer/size pair per physical page run.""" + """Build host page-run buffers for storage GET and PUT operations.""" if not entry.page_runs: return [], [] @@ -250,6 +260,36 @@ def build_transfer_buffers( return ptrs, sizes +def _build_host_device_transfer_plan( + entry: EmbeddingCacheEntry, + pool: EmbeddingPool, + src_is_pool: bool, + dst_token_offset: int = 0, +) -> Tuple[List[int], List[int], List[int]]: + """Build token ranges for transferring between the host pool and device.""" + src_starts: List[int] = [] + dst_starts: List[int] = [] + lengths: List[int] = [] + copied = 0 + + for run in entry.page_runs: + valid_tokens = min(pool.page_size * run.length, entry.num_tokens - copied) + if valid_tokens <= 0: + break + + pool_start = run.start * pool.page_size + if src_is_pool: + src_starts.append(pool_start) + dst_starts.append(dst_token_offset + copied) + else: + src_starts.append(copied) + dst_starts.append(pool_start) + lengths.append(valid_tokens) + copied += valid_tokens + + return src_starts, dst_starts, lengths + + @dataclass class AsyncCopyHandle: event: object @@ -565,7 +605,7 @@ class EmbeddingCacheController: ) self.entries[mm_hash] = entry keys.append(mm_hash) - entry_ptrs, entry_sizes = build_transfer_buffers(entry, pool) + entry_ptrs, entry_sizes = _build_storage_transfer_buffers(entry, pool) all_ptrs.append(entry_ptrs) all_sizes.append(entry_sizes) @@ -609,7 +649,7 @@ class EmbeddingCacheController: self._pin_read(entry) keys.append(mm_hash) - entry_ptrs, entry_sizes = build_transfer_buffers(entry, pool) + entry_ptrs, entry_sizes = _build_storage_transfer_buffers(entry, pool) all_ptrs.append(entry_ptrs) all_sizes.append(entry_sizes) @@ -749,21 +789,15 @@ class EmbeddingCacheController: device = dst_tensor.device copy_stream = self._get_copy_stream(device) event = torch.cuda.Event() - copied = 0 with torch.cuda.stream(copy_stream): - for run in entry.page_runs: - valid_tokens = min( - pool.page_size * run.length, entry.num_tokens - copied - ) - if valid_tokens <= 0: - break - src_start = run.start * pool.page_size - dst_start = dst_token_offset + copied - dst_tensor[dst_start : dst_start + valid_tokens].copy_( - pool.tensor[src_start : src_start + valid_tokens], - non_blocking=True, - ) - copied += valid_tokens + self._copy_embedding_page_runs( + src=pool.tensor, + dst=dst_tensor, + entry=entry, + pool=pool, + src_is_pool=True, + dst_token_offset=dst_token_offset, + ) event.record(copy_stream) return AsyncCopyHandle(event, mm_hash, device=torch.device(device)) except Exception: @@ -828,6 +862,39 @@ class EmbeddingCacheController: self._copy_streams[key] = stream return stream + def _copy_embedding_page_runs( + self, + src: torch.Tensor, + dst: torch.Tensor, + entry: EmbeddingCacheEntry, + pool: EmbeddingPool, + src_is_pool: bool, + dst_token_offset: int = 0, + ) -> None: + """Copy one embedding entry between its host pool and a CUDA tensor.""" + src_starts, dst_starts, lengths = _build_host_device_transfer_plan( + entry, pool, src_is_pool, dst_token_offset + ) + if not lengths: + return + + has_cuda_side = src.device.type == "cuda" or dst.device.type == "cuda" + if has_cuda_side and transfer_embedding_ranges_direct is not None: + transfer_embedding_ranges_direct( + src, + dst, + src_starts, + dst_starts, + lengths, + ) + return + + for src_start, dst_start, valid_tokens in zip(src_starts, dst_starts, lengths): + dst[dst_start : dst_start + valid_tokens].copy_( + src[src_start : src_start + valid_tokens], + non_blocking=True, + ) + def has_local_embedding(self, mm_hash: str) -> bool: with self.lock: entry = self.entries.get(mm_hash) @@ -936,20 +1003,14 @@ class EmbeddingCacheController: copy_stream.wait_stream(producer_stream) src.record_stream(copy_stream) event = torch.cuda.Event() - copied = 0 with torch.cuda.stream(copy_stream): - for run in entry.page_runs: - valid_tokens = min( - pool.page_size * run.length, entry.num_tokens - copied - ) - if valid_tokens <= 0: - break - start = run.start * pool.page_size - pool.tensor[start : start + valid_tokens].copy_( - src[copied : copied + valid_tokens], - non_blocking=True, - ) - copied += valid_tokens + self._copy_embedding_page_runs( + src=src, + dst=pool.tensor, + entry=entry, + pool=pool, + src_is_pool=False, + ) event.record(copy_stream) return AsyncCopyHandle( event=event, diff --git a/test/registered/unit/mem_cache/test_embedding_cache_controller.py b/test/registered/unit/mem_cache/test_embedding_cache_controller.py index 557c31af1..1fb4efd6e 100644 --- a/test/registered/unit/mem_cache/test_embedding_cache_controller.py +++ b/test/registered/unit/mem_cache/test_embedding_cache_controller.py @@ -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):