[HiCache] Align chunked CUDA host registrations (#36798)
Co-authored-by: Zhangheng <hzh0425@apache.org>
This commit is contained in:
@@ -696,6 +696,8 @@ class Envs:
|
|||||||
# ===================================================================
|
# ===================================================================
|
||||||
# HiCache storage backends and mmap allocation
|
# HiCache storage backends and mmap allocation
|
||||||
# ===================================================================
|
# ===================================================================
|
||||||
|
# Per-call cudaHostRegister limit in GB.
|
||||||
|
SGLANG_HICACHE_HOST_REGISTER_CHUNK_GB = EnvInt(256)
|
||||||
SGLANG_HICACHE_HF3FS_CONFIG_PATH = EnvStr(None)
|
SGLANG_HICACHE_HF3FS_CONFIG_PATH = EnvStr(None)
|
||||||
SGLANG_HICACHE_DECODE_OFFLOAD_STRIDE = EnvInt(None)
|
SGLANG_HICACHE_DECODE_OFFLOAD_STRIDE = EnvInt(None)
|
||||||
SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR = EnvStr(None)
|
SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR = EnvStr(None)
|
||||||
|
|||||||
@@ -236,6 +236,7 @@ class DeepSeekV4PagedHostPool(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
pin_memory=self.pin_memory,
|
pin_memory=self.pin_memory,
|
||||||
allocator=self.allocator,
|
allocator=self.allocator,
|
||||||
|
registration_granularity_bytes=self.layer_num * self.item_bytes,
|
||||||
)
|
)
|
||||||
elif self.layout == "page_first_direct":
|
elif self.layout == "page_first_direct":
|
||||||
self.kv_buffer = alloc_func(
|
self.kv_buffer = alloc_func(
|
||||||
@@ -244,6 +245,7 @@ class DeepSeekV4PagedHostPool(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
pin_memory=self.pin_memory,
|
pin_memory=self.pin_memory,
|
||||||
allocator=self.allocator,
|
allocator=self.allocator,
|
||||||
|
registration_granularity_bytes=self.layer_num * self.item_bytes,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported layout: {self.layout}")
|
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||||
@@ -639,6 +641,7 @@ class DeepSeekV4StateHostPool(HostKVCache):
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
pin_memory=self.pin_memory,
|
pin_memory=self.pin_memory,
|
||||||
allocator=self.allocator,
|
allocator=self.allocator,
|
||||||
|
registration_granularity_bytes=(self.layer_num * self.state_page_bytes),
|
||||||
)
|
)
|
||||||
elif self.layout == "page_first_direct":
|
elif self.layout == "page_first_direct":
|
||||||
self.kv_buffer = alloc_func(
|
self.kv_buffer = alloc_func(
|
||||||
@@ -647,6 +650,7 @@ class DeepSeekV4StateHostPool(HostKVCache):
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
pin_memory=self.pin_memory,
|
pin_memory=self.pin_memory,
|
||||||
allocator=self.allocator,
|
allocator=self.allocator,
|
||||||
|
registration_granularity_bytes=(self.layer_num * self.state_page_bytes),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported layout: {self.layout}")
|
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||||
|
|||||||
@@ -7,10 +7,13 @@ from collections import defaultdict
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.mem_cache.storage.mmap import alloc_mmap
|
from sglang.srt.mem_cache.storage.mmap import alloc_mmap
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_CUDA_HOST_REGISTERED_RANGES_ATTR = "_sglang_cuda_host_registered_ranges"
|
||||||
|
|
||||||
|
|
||||||
class HostTensorAllocator:
|
class HostTensorAllocator:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
@@ -118,30 +121,101 @@ def get_allocator_type() -> str:
|
|||||||
return backend or "default"
|
return backend or "default"
|
||||||
|
|
||||||
|
|
||||||
def _cuda_host_register(buffer: torch.Tensor) -> None:
|
def _cuda_host_register(
|
||||||
|
buffer: torch.Tensor, registration_granularity_bytes: int | None = None
|
||||||
|
) -> None:
|
||||||
|
# Avoid oversized cudaHostRegister calls on large host pools.
|
||||||
cudart = torch.cuda.cudart()
|
cudart = torch.cuda.cudart()
|
||||||
n_bytes = buffer.numel() * buffer.element_size()
|
base = buffer.data_ptr()
|
||||||
rc = cudart.cudaHostRegister(buffer.data_ptr(), n_bytes, 0)
|
total = buffer.numel() * buffer.element_size()
|
||||||
if int(rc) != 0:
|
chunk_limit_bytes = (
|
||||||
raise RuntimeError(
|
max(envs.SGLANG_HICACHE_HOST_REGISTER_CHUNK_GB.get(), 1) * 1024**3
|
||||||
f"cudaHostRegister failed (rc={int(rc)}, "
|
)
|
||||||
f"{cudart.cudaGetErrorString(rc)}) for ptr={buffer.data_ptr():#x} "
|
# Preserve the legacy single-call behavior unless the caller provides a
|
||||||
f"size={n_bytes}; host buffer is not pinned and device transfers "
|
# copy granularity. Splitting an unknown page-first layout at an arbitrary
|
||||||
f"may silently return stale data."
|
# byte offset can make one cudaMemcpyBatchAsync span two registrations.
|
||||||
|
chunk_bytes = total
|
||||||
|
if registration_granularity_bytes is not None:
|
||||||
|
if registration_granularity_bytes <= 0:
|
||||||
|
raise ValueError(
|
||||||
|
"registration_granularity_bytes must be positive, got "
|
||||||
|
f"{registration_granularity_bytes}"
|
||||||
|
)
|
||||||
|
if registration_granularity_bytes > chunk_limit_bytes:
|
||||||
|
raise ValueError(
|
||||||
|
"Host registration granularity exceeds the configured chunk limit: "
|
||||||
|
f"granularity={registration_granularity_bytes}, "
|
||||||
|
f"chunk_limit={chunk_limit_bytes}"
|
||||||
|
)
|
||||||
|
chunk_bytes = (
|
||||||
|
chunk_limit_bytes // registration_granularity_bytes
|
||||||
|
) * registration_granularity_bytes
|
||||||
|
registered_ranges: list[tuple[int, int]] = []
|
||||||
|
try:
|
||||||
|
offset = 0
|
||||||
|
while offset < total:
|
||||||
|
size = min(chunk_bytes, total - offset)
|
||||||
|
ptr = base + offset
|
||||||
|
rc = int(cudart.cudaHostRegister(ptr, size, 0))
|
||||||
|
if rc != 0:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"cudaHostRegister failed (rc={rc}, "
|
||||||
|
f"{cudart.cudaGetErrorString(rc)}) at offset={offset} size={size} "
|
||||||
|
f"(total={total}, chunk_limit={chunk_bytes}); host buffer is not "
|
||||||
|
f"pinned and device transfers may silently return stale data."
|
||||||
|
)
|
||||||
|
registered_ranges.append((ptr, size))
|
||||||
|
offset += size
|
||||||
|
|
||||||
|
# Keep the exact registration bases alive with the tensor. CUDA requires
|
||||||
|
# cudaHostUnregister to receive each base pointer, not just the tensor's
|
||||||
|
# original base once after several independent registrations.
|
||||||
|
setattr(buffer, _CUDA_HOST_REGISTERED_RANGES_ATTR, registered_ranges)
|
||||||
|
except Exception:
|
||||||
|
remaining_ranges = _cuda_host_unregister_ranges(
|
||||||
|
cudart, registered_ranges, operation="registration rollback"
|
||||||
)
|
)
|
||||||
|
if remaining_ranges:
|
||||||
|
setattr(buffer, _CUDA_HOST_REGISTERED_RANGES_ATTR, remaining_ranges)
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
def _cuda_host_unregister_ranges(
|
||||||
|
cudart, registered_ranges: list[tuple[int, int]], *, operation: str
|
||||||
|
) -> list[tuple[int, int]]:
|
||||||
|
failed_ranges = []
|
||||||
|
for ptr, size in reversed(registered_ranges):
|
||||||
|
rc = int(cudart.cudaHostUnregister(ptr))
|
||||||
|
if rc != 0:
|
||||||
|
failed_ranges.append((ptr, size))
|
||||||
|
logger.warning(
|
||||||
|
"cudaHostUnregister failed during %s (rc=%d, %s) "
|
||||||
|
"for ptr=%#x size=%d",
|
||||||
|
operation,
|
||||||
|
rc,
|
||||||
|
cudart.cudaGetErrorString(rc),
|
||||||
|
ptr,
|
||||||
|
size,
|
||||||
|
)
|
||||||
|
failed_ranges.reverse()
|
||||||
|
return failed_ranges
|
||||||
|
|
||||||
|
|
||||||
def _cuda_host_unregister(buffer: torch.Tensor) -> None:
|
def _cuda_host_unregister(buffer: torch.Tensor) -> None:
|
||||||
cudart = torch.cuda.cudart()
|
cudart = torch.cuda.cudart()
|
||||||
rc = cudart.cudaHostUnregister(buffer.data_ptr())
|
registered_ranges = getattr(buffer, _CUDA_HOST_REGISTERED_RANGES_ATTR, None)
|
||||||
if int(rc) != 0:
|
if registered_ranges is None:
|
||||||
# Best-effort on shutdown: warn, don't raise -- a leak is reclaimed at exit.
|
# Compatibility for buffers registered before range metadata was added.
|
||||||
logger.warning(
|
registered_ranges = [
|
||||||
"cudaHostUnregister failed (rc=%d, %s) for ptr=%#x",
|
(buffer.data_ptr(), buffer.numel() * buffer.element_size())
|
||||||
int(rc),
|
]
|
||||||
cudart.cudaGetErrorString(rc),
|
if not registered_ranges:
|
||||||
buffer.data_ptr(),
|
return
|
||||||
)
|
|
||||||
|
remaining_ranges = _cuda_host_unregister_ranges(
|
||||||
|
cudart, registered_ranges, operation="host-pool destroy"
|
||||||
|
)
|
||||||
|
setattr(buffer, _CUDA_HOST_REGISTERED_RANGES_ATTR, remaining_ranges)
|
||||||
|
|
||||||
|
|
||||||
def alloc_with_host_register(
|
def alloc_with_host_register(
|
||||||
@@ -150,6 +224,7 @@ def alloc_with_host_register(
|
|||||||
device: str,
|
device: str,
|
||||||
pin_memory: bool,
|
pin_memory: bool,
|
||||||
allocator: HostTensorAllocator,
|
allocator: HostTensorAllocator,
|
||||||
|
registration_granularity_bytes: int | None = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
Allocate tensor and register host memory with cudaHostRegister.
|
Allocate tensor and register host memory with cudaHostRegister.
|
||||||
@@ -157,7 +232,7 @@ def alloc_with_host_register(
|
|||||||
"""
|
"""
|
||||||
buffer = allocator.allocate(dims, dtype=dtype, device=device)
|
buffer = allocator.allocate(dims, dtype=dtype, device=device)
|
||||||
if pin_memory:
|
if pin_memory:
|
||||||
_cuda_host_register(buffer)
|
_cuda_host_register(buffer, registration_granularity_bytes)
|
||||||
return buffer
|
return buffer
|
||||||
|
|
||||||
|
|
||||||
@@ -167,6 +242,7 @@ def alloc_with_pin_memory(
|
|||||||
device: str,
|
device: str,
|
||||||
pin_memory: bool,
|
pin_memory: bool,
|
||||||
allocator: None,
|
allocator: None,
|
||||||
|
registration_granularity_bytes: int | None = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
Allocate tensor using PyTorch's built-in pin_memory flag.
|
Allocate tensor using PyTorch's built-in pin_memory flag.
|
||||||
|
|||||||
@@ -173,6 +173,7 @@ class DSAIndexerPoolHost(HostKVCache):
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
pin_memory=self.pin_memory,
|
pin_memory=self.pin_memory,
|
||||||
allocator=self.allocator,
|
allocator=self.allocator,
|
||||||
|
registration_granularity_bytes=self.indexer_layout_dim,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported layout: {self.layout}")
|
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||||
|
|||||||
@@ -146,6 +146,9 @@ class MambaPoolHost(HostKVCache):
|
|||||||
device=device,
|
device=device,
|
||||||
pin_memory=pin_memory,
|
pin_memory=pin_memory,
|
||||||
allocator=allocator,
|
allocator=allocator,
|
||||||
|
registration_granularity_bytes=(
|
||||||
|
int(np.prod(dims[1:])) * dtype.itemsize
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.layout in ["page_first", "page_first_direct"]:
|
if self.layout in ["page_first", "page_first_direct"]:
|
||||||
|
|||||||
@@ -194,6 +194,11 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
pin_memory=self.pin_memory,
|
pin_memory=self.pin_memory,
|
||||||
allocator=self.allocator,
|
allocator=self.allocator,
|
||||||
|
registration_granularity_bytes=(
|
||||||
|
self.page_size * self.layout_dim
|
||||||
|
if self.layout in ("page_first", "page_first_direct")
|
||||||
|
else None
|
||||||
|
),
|
||||||
)
|
)
|
||||||
return buffer
|
return buffer
|
||||||
|
|
||||||
@@ -794,6 +799,11 @@ class MHATokenToKOnlyPoolHost(HostKVCache):
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
pin_memory=self.pin_memory,
|
pin_memory=self.pin_memory,
|
||||||
allocator=self.allocator,
|
allocator=self.allocator,
|
||||||
|
registration_granularity_bytes=(
|
||||||
|
self.page_size * self.layout_dim
|
||||||
|
if self.layout in ("page_first", "page_first_direct")
|
||||||
|
else None
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_hybrid_pool_buffer(self):
|
def get_hybrid_pool_buffer(self):
|
||||||
@@ -1117,6 +1127,7 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
pin_memory=self.pin_memory,
|
pin_memory=self.pin_memory,
|
||||||
allocator=self.allocator,
|
allocator=self.allocator,
|
||||||
|
registration_granularity_bytes=self.page_size * self._k_layout_dim(),
|
||||||
)
|
)
|
||||||
v_buffer = alloc_func(
|
v_buffer = alloc_func(
|
||||||
v_dims,
|
v_dims,
|
||||||
@@ -1124,6 +1135,7 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
pin_memory=self.pin_memory,
|
pin_memory=self.pin_memory,
|
||||||
allocator=self.allocator,
|
allocator=self.allocator,
|
||||||
|
registration_granularity_bytes=self.page_size * self._v_layout_dim(),
|
||||||
)
|
)
|
||||||
return (k_buffer, v_buffer)
|
return (k_buffer, v_buffer)
|
||||||
|
|
||||||
|
|||||||
@@ -206,6 +206,11 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
pin_memory=self.pin_memory,
|
pin_memory=self.pin_memory,
|
||||||
allocator=self.allocator,
|
allocator=self.allocator,
|
||||||
|
registration_granularity_bytes=(
|
||||||
|
self.page_size * self.layout_dim
|
||||||
|
if self.layout in ("page_first", "page_first_direct")
|
||||||
|
else None
|
||||||
|
),
|
||||||
)
|
)
|
||||||
return buffer
|
return buffer
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,412 @@
|
|||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest import mock
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.mem_cache import memory_pool_host
|
||||||
|
from sglang.srt.mem_cache.memory_pool_host import (
|
||||||
|
DeepSeekV4PagedHostPool,
|
||||||
|
DeepSeekV4StateHostPool,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.pool_host import mha as mha_pool_host
|
||||||
|
from sglang.srt.mem_cache.pool_host import mla as mla_pool_host
|
||||||
|
from sglang.srt.mem_cache.pool_host.common import (
|
||||||
|
ALLOC_MEMORY_FUNCS,
|
||||||
|
_cuda_host_register,
|
||||||
|
_cuda_host_unregister,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.pool_host.dsa import DSAIndexerPoolHost
|
||||||
|
from sglang.srt.mem_cache.pool_host.mamba import MambaPoolHost
|
||||||
|
from sglang.srt.mem_cache.pool_host.mha import (
|
||||||
|
AsymmetricMHATokenToKVPoolHost,
|
||||||
|
MHATokenToKOnlyPoolHost,
|
||||||
|
MHATokenToKVPoolHost,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeBuffer:
|
||||||
|
def __init__(self, base: int, size: int):
|
||||||
|
self._base = base
|
||||||
|
self._size = size
|
||||||
|
|
||||||
|
def data_ptr(self) -> int:
|
||||||
|
return self._base
|
||||||
|
|
||||||
|
def numel(self) -> int:
|
||||||
|
return self._size
|
||||||
|
|
||||||
|
def element_size(self) -> int:
|
||||||
|
return 1
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeCudart:
|
||||||
|
def __init__(self, fail_on_registration: int | None = None):
|
||||||
|
self.registrations = []
|
||||||
|
self.unregistrations = []
|
||||||
|
self.fail_on_registration = fail_on_registration
|
||||||
|
|
||||||
|
def cudaHostRegister(self, ptr: int, size: int, flags: int) -> int:
|
||||||
|
self.registrations.append((ptr, size, flags))
|
||||||
|
if len(self.registrations) == self.fail_on_registration:
|
||||||
|
return 1
|
||||||
|
return 0
|
||||||
|
|
||||||
|
def cudaHostUnregister(self, ptr: int) -> int:
|
||||||
|
self.unregistrations.append(ptr)
|
||||||
|
return 0
|
||||||
|
|
||||||
|
def cudaGetErrorString(self, rc: int) -> str:
|
||||||
|
return "injected error"
|
||||||
|
|
||||||
|
|
||||||
|
class TestHiCacheHostRegister(unittest.TestCase):
|
||||||
|
def test_dsa_page_layouts_with_draft_use_page_registration_granularity(self):
|
||||||
|
target_buffers = [torch.empty(1, dtype=torch.uint8) for _ in range(3)]
|
||||||
|
draft_buffer = torch.empty(1, dtype=torch.uint8)
|
||||||
|
|
||||||
|
for layout in ("page_first", "page_first_direct"):
|
||||||
|
with self.subTest(layout=layout):
|
||||||
|
host = DSAIndexerPoolHost.__new__(DSAIndexerPoolHost)
|
||||||
|
host.device_pool = SimpleNamespace(
|
||||||
|
device="cpu", index_k_with_scale_buffer=target_buffers
|
||||||
|
)
|
||||||
|
host.mtp_draft_device_pools = [
|
||||||
|
SimpleNamespace(index_k_with_scale_buffer=[draft_buffer])
|
||||||
|
]
|
||||||
|
host.layout = layout
|
||||||
|
host.layer_num = 4
|
||||||
|
host.indexer_page_num = 3
|
||||||
|
host.indexer_page_stride_size = 512
|
||||||
|
host.indexer_layout_dim = host.layer_num * host.indexer_page_stride_size
|
||||||
|
host.indexer_dtype = torch.uint8
|
||||||
|
host.device = "cpu"
|
||||||
|
host.pin_memory = True
|
||||||
|
host.allocator = mock.sentinel.allocator
|
||||||
|
alloc = mock.Mock(return_value=torch.empty(1, dtype=torch.uint8))
|
||||||
|
|
||||||
|
with mock.patch.dict(ALLOC_MEMORY_FUNCS, {"cpu": alloc}):
|
||||||
|
host.init_kv_buffer()
|
||||||
|
|
||||||
|
self.assertEqual(len(host.packed_device_index_buffers), 4)
|
||||||
|
self.assertIs(host.packed_device_index_buffers[-1], draft_buffer)
|
||||||
|
self.assertEqual(
|
||||||
|
alloc.call_args.kwargs["registration_granularity_bytes"],
|
||||||
|
host.indexer_layout_dim,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_page_first_direct_mla_uses_page_registration_granularity(self):
|
||||||
|
pool = MLATokenToKVPoolHost.__new__(MLATokenToKVPoolHost)
|
||||||
|
pool.layout = "page_first_direct"
|
||||||
|
pool.page_num = 4
|
||||||
|
pool.layer_num = 3
|
||||||
|
pool.page_size = 2
|
||||||
|
pool.kv_cache_dim = 5
|
||||||
|
pool.dtype = torch.float16
|
||||||
|
pool.device_pool = SimpleNamespace(device="cuda")
|
||||||
|
pool.device = "cpu"
|
||||||
|
pool.pin_memory = True
|
||||||
|
pool.allocator = object()
|
||||||
|
alloc = mock.Mock(return_value=object())
|
||||||
|
|
||||||
|
with mock.patch.dict(mla_pool_host.ALLOC_MEMORY_FUNCS, {"cuda": alloc}):
|
||||||
|
pool.init_kv_buffer()
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
alloc.call_args.kwargs["registration_granularity_bytes"],
|
||||||
|
pool.page_size * pool.layer_num * pool.kv_cache_dim * pool.dtype.itemsize,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_page_first_direct_mha_uses_page_registration_granularity(self):
|
||||||
|
pool = MHATokenToKVPoolHost.__new__(MHATokenToKVPoolHost)
|
||||||
|
pool.layout = "page_first_direct"
|
||||||
|
pool.page_num = 4
|
||||||
|
pool.layer_num = 3
|
||||||
|
pool.page_size = 2
|
||||||
|
pool.head_num = 2
|
||||||
|
pool.head_dim = 4
|
||||||
|
pool.dtype = torch.float16
|
||||||
|
pool.device_pool = SimpleNamespace(device="cuda")
|
||||||
|
pool.device = "cpu"
|
||||||
|
pool.pin_memory = True
|
||||||
|
pool.allocator = object()
|
||||||
|
alloc = mock.Mock(return_value=object())
|
||||||
|
|
||||||
|
with mock.patch.dict(mha_pool_host.ALLOC_MEMORY_FUNCS, {"cuda": alloc}):
|
||||||
|
pool.init_kv_buffer()
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
alloc.call_args.kwargs["registration_granularity_bytes"],
|
||||||
|
pool.page_size
|
||||||
|
* pool.layer_num
|
||||||
|
* pool.head_num
|
||||||
|
* pool.head_dim
|
||||||
|
* pool.dtype.itemsize,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_mamba_page_layouts_use_per_buffer_page_granularity(self):
|
||||||
|
for layout in ("page_first", "page_first_direct"):
|
||||||
|
with self.subTest(layout=layout):
|
||||||
|
pool = MambaPoolHost.__new__(MambaPoolHost)
|
||||||
|
pool.layout = layout
|
||||||
|
pool.size = 4
|
||||||
|
pool.num_mamba_layers = 3
|
||||||
|
pool.temporal_state_shape = (2, 5)
|
||||||
|
pool.conv_state_shapes = [(7,), (2, 2)]
|
||||||
|
pool.temporal_dtype = torch.float16
|
||||||
|
pool.conv_dtype = torch.float32
|
||||||
|
pool.device_pool = SimpleNamespace(device="cuda")
|
||||||
|
pool.device = "cpu"
|
||||||
|
pool.pin_memory = True
|
||||||
|
pool.allocator = object()
|
||||||
|
alloc = mock.Mock(
|
||||||
|
side_effect=lambda *args, **kwargs: torch.empty(
|
||||||
|
1, dtype=torch.uint8
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
with mock.patch.dict(ALLOC_MEMORY_FUNCS, {"cuda": alloc}):
|
||||||
|
pool.init_kv_buffer()
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
[
|
||||||
|
call.kwargs["registration_granularity_bytes"]
|
||||||
|
for call in alloc.call_args_list
|
||||||
|
],
|
||||||
|
[
|
||||||
|
3 * 2 * 5 * torch.float16.itemsize,
|
||||||
|
3 * 7 * torch.float32.itemsize,
|
||||||
|
3 * 2 * 2 * torch.float32.itemsize,
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_deepseek_v4_page_layouts_use_page_registration_granularity(self):
|
||||||
|
for layout in ("page_first", "page_first_direct"):
|
||||||
|
with self.subTest(pool="paged", layout=layout):
|
||||||
|
alloc = mock.Mock(return_value=torch.empty(1, dtype=torch.uint8))
|
||||||
|
device_buffers = [torch.empty(1, dtype=torch.uint8) for _ in range(3)]
|
||||||
|
with (
|
||||||
|
mock.patch.object(
|
||||||
|
memory_pool_host,
|
||||||
|
"host_memory_budget_bytes",
|
||||||
|
return_value=1024**3,
|
||||||
|
),
|
||||||
|
mock.patch.dict(ALLOC_MEMORY_FUNCS, {torch.device("cpu"): alloc}),
|
||||||
|
):
|
||||||
|
DeepSeekV4PagedHostPool(
|
||||||
|
pool_name="test",
|
||||||
|
device_buffers=device_buffers,
|
||||||
|
item_bytes=11,
|
||||||
|
num_host_pages=4,
|
||||||
|
slot_page_size=2,
|
||||||
|
layout=layout,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
alloc.call_args.kwargs["registration_granularity_bytes"],
|
||||||
|
3 * 11,
|
||||||
|
)
|
||||||
|
|
||||||
|
with self.subTest(pool="state", layout=layout):
|
||||||
|
alloc = mock.Mock(return_value=torch.empty(1, dtype=torch.uint8))
|
||||||
|
state_pools = [
|
||||||
|
SimpleNamespace(
|
||||||
|
ring_size=2,
|
||||||
|
kv_score_buffer=SimpleNamespace(
|
||||||
|
kv_score=torch.empty((4, 3), dtype=torch.uint8)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
for _ in range(2)
|
||||||
|
]
|
||||||
|
with (
|
||||||
|
mock.patch.object(
|
||||||
|
memory_pool_host,
|
||||||
|
"host_memory_budget_bytes",
|
||||||
|
return_value=1024**3,
|
||||||
|
),
|
||||||
|
mock.patch.dict(ALLOC_MEMORY_FUNCS, {torch.device("cpu"): alloc}),
|
||||||
|
):
|
||||||
|
DeepSeekV4StateHostPool(
|
||||||
|
pool_name="test",
|
||||||
|
state_pools=state_pools,
|
||||||
|
num_host_pages=4,
|
||||||
|
swa_page_size=2,
|
||||||
|
layout=layout,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
alloc.call_args.kwargs["registration_granularity_bytes"],
|
||||||
|
2 * 2 * 3,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_k_only_mha_page_layouts_use_page_registration_granularity(self):
|
||||||
|
for layout in ("page_first", "page_first_direct"):
|
||||||
|
with self.subTest(layout=layout):
|
||||||
|
pool = MHATokenToKOnlyPoolHost.__new__(MHATokenToKOnlyPoolHost)
|
||||||
|
pool.layout = layout
|
||||||
|
pool.size = 8
|
||||||
|
pool.page_num = 4
|
||||||
|
pool.page_size = 2
|
||||||
|
pool.layer_num = 3
|
||||||
|
pool.head_num = 2
|
||||||
|
pool.head_dim = 5
|
||||||
|
pool.dtype = torch.float16
|
||||||
|
pool.layout_dim = (
|
||||||
|
pool.layer_num * pool.head_num * pool.head_dim * pool.dtype.itemsize
|
||||||
|
)
|
||||||
|
pool.device_pool = SimpleNamespace(device="cuda")
|
||||||
|
pool.device = "cpu"
|
||||||
|
pool.pin_memory = True
|
||||||
|
pool.allocator = object()
|
||||||
|
alloc = mock.Mock(return_value=object())
|
||||||
|
|
||||||
|
with mock.patch.dict(ALLOC_MEMORY_FUNCS, {"cuda": alloc}):
|
||||||
|
pool.init_kv_buffer()
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
alloc.call_args.kwargs["registration_granularity_bytes"],
|
||||||
|
pool.page_size * pool.layout_dim,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_asymmetric_mha_page_layouts_use_native_page_granularities(self):
|
||||||
|
for layout in ("page_first", "page_first_direct"):
|
||||||
|
with self.subTest(layout=layout):
|
||||||
|
pool = AsymmetricMHATokenToKVPoolHost.__new__(
|
||||||
|
AsymmetricMHATokenToKVPoolHost
|
||||||
|
)
|
||||||
|
pool.layout = layout
|
||||||
|
pool.size = 8
|
||||||
|
pool.page_num = 4
|
||||||
|
pool.page_size = 2
|
||||||
|
pool.layer_num = 3
|
||||||
|
pool.head_num = 2
|
||||||
|
pool.head_dim = 5
|
||||||
|
pool.v_head_dim = 7
|
||||||
|
pool.dtype = torch.float16
|
||||||
|
pool.device_pool = SimpleNamespace(device="cuda")
|
||||||
|
pool.device = "cpu"
|
||||||
|
pool.pin_memory = True
|
||||||
|
pool.allocator = object()
|
||||||
|
alloc = mock.Mock(side_effect=[object(), object()])
|
||||||
|
|
||||||
|
with mock.patch.dict(ALLOC_MEMORY_FUNCS, {"cuda": alloc}):
|
||||||
|
pool.init_kv_buffer()
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
[
|
||||||
|
call.kwargs["registration_granularity_bytes"]
|
||||||
|
for call in alloc.call_args_list
|
||||||
|
],
|
||||||
|
[
|
||||||
|
pool.page_size * pool._k_layout_dim(),
|
||||||
|
pool.page_size * pool._v_layout_dim(),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_unregister_releases_every_registered_chunk_once(self):
|
||||||
|
gib = 1024**3
|
||||||
|
base = 0x10000000
|
||||||
|
buffer = _FakeBuffer(base, 2 * gib + 17)
|
||||||
|
cudart = _FakeCudart()
|
||||||
|
|
||||||
|
with (
|
||||||
|
mock.patch.object(
|
||||||
|
envs.SGLANG_HICACHE_HOST_REGISTER_CHUNK_GB,
|
||||||
|
"get",
|
||||||
|
return_value=1,
|
||||||
|
),
|
||||||
|
mock.patch.object(torch.cuda, "cudart", return_value=cudart),
|
||||||
|
):
|
||||||
|
_cuda_host_register(buffer, registration_granularity_bytes=gib)
|
||||||
|
_cuda_host_unregister(buffer)
|
||||||
|
_cuda_host_unregister(buffer)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
cudart.unregistrations,
|
||||||
|
[base + 2 * gib, base + gib, base],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_registration_failure_rolls_back_prior_chunks(self):
|
||||||
|
gib = 1024**3
|
||||||
|
base = 0x10000000
|
||||||
|
buffer = _FakeBuffer(base, 2 * gib + 17)
|
||||||
|
cudart = _FakeCudart(fail_on_registration=2)
|
||||||
|
|
||||||
|
with (
|
||||||
|
mock.patch.object(
|
||||||
|
envs.SGLANG_HICACHE_HOST_REGISTER_CHUNK_GB,
|
||||||
|
"get",
|
||||||
|
return_value=1,
|
||||||
|
),
|
||||||
|
mock.patch.object(torch.cuda, "cudart", return_value=cudart),
|
||||||
|
self.assertRaisesRegex(RuntimeError, "offset=1073741824"),
|
||||||
|
):
|
||||||
|
_cuda_host_register(buffer, registration_granularity_bytes=gib)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
cudart.registrations,
|
||||||
|
[(base, gib, 0), (base + gib, gib, 0)],
|
||||||
|
)
|
||||||
|
self.assertEqual(cudart.unregistrations, [base])
|
||||||
|
|
||||||
|
def test_missing_copy_granularity_preserves_single_registration(self):
|
||||||
|
gib = 1024**3
|
||||||
|
base = 0x10000000
|
||||||
|
total = 2 * gib + 17
|
||||||
|
buffer = _FakeBuffer(base, total)
|
||||||
|
cudart = _FakeCudart()
|
||||||
|
|
||||||
|
with (
|
||||||
|
mock.patch.object(
|
||||||
|
envs.SGLANG_HICACHE_HOST_REGISTER_CHUNK_GB,
|
||||||
|
"get",
|
||||||
|
return_value=1,
|
||||||
|
),
|
||||||
|
mock.patch.object(torch.cuda, "cudart", return_value=cudart),
|
||||||
|
):
|
||||||
|
_cuda_host_register(buffer)
|
||||||
|
|
||||||
|
self.assertEqual(cudart.registrations, [(base, total, 0)])
|
||||||
|
|
||||||
|
def test_registration_boundaries_honor_page_copy_granularity(self):
|
||||||
|
mib = 1024**2
|
||||||
|
gib = 1024**3
|
||||||
|
base = 0x10000000
|
||||||
|
total = 2500 * mib
|
||||||
|
page_copy_bytes = 300 * mib
|
||||||
|
cudart = _FakeCudart()
|
||||||
|
|
||||||
|
with (
|
||||||
|
mock.patch.object(
|
||||||
|
envs.SGLANG_HICACHE_HOST_REGISTER_CHUNK_GB,
|
||||||
|
"get",
|
||||||
|
return_value=1,
|
||||||
|
),
|
||||||
|
mock.patch.object(torch.cuda, "cudart", return_value=cudart),
|
||||||
|
):
|
||||||
|
_cuda_host_register(
|
||||||
|
_FakeBuffer(base, total),
|
||||||
|
registration_granularity_bytes=page_copy_bytes,
|
||||||
|
)
|
||||||
|
|
||||||
|
aligned_chunk = 900 * mib
|
||||||
|
self.assertLessEqual(aligned_chunk, gib)
|
||||||
|
self.assertEqual(
|
||||||
|
cudart.registrations,
|
||||||
|
[
|
||||||
|
(base, aligned_chunk, 0),
|
||||||
|
(base + aligned_chunk, aligned_chunk, 0),
|
||||||
|
(base + 2 * aligned_chunk, 700 * mib, 0),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
for ptr, _, _ in cudart.registrations:
|
||||||
|
self.assertEqual((ptr - base) % page_copy_bytes, 0)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user