[HiCache] Optimize HiCache hash generation with bulk token byte conversion (#28287)

This commit is contained in:
huangtingwei
2026-07-01 15:44:23 +08:00
committed by GitHub
parent df0dfbaa45
commit 5e1ccd9320
7 changed files with 675 additions and 180 deletions
+5 -11
View File
@@ -1002,18 +1002,12 @@ class HiCacheController:
storage_query_count = 0
hash_value = []
page_hashes = self.get_hash_str(
tokens_to_fetch, last_hash, page_size=self.page_size
)
for start in range(
0, len(tokens_to_fetch), self.page_size * STORAGE_BATCH_SIZE
):
end = min(start + self.page_size * STORAGE_BATCH_SIZE, len(tokens_to_fetch))
batch_tokens = tokens_to_fetch[start:end]
batch_hashes = []
for i in range(0, len(batch_tokens), self.page_size):
last_hash = self.get_hash_str(
batch_tokens[i : i + self.page_size], last_hash
)
batch_hashes.append(last_hash)
for start in range(0, len(page_hashes), STORAGE_BATCH_SIZE):
batch_hashes = page_hashes[start : start + STORAGE_BATCH_SIZE]
extra_info = HiCacheStorageExtraInfo(prefix_keys=prefix_keys)
hit_page_num = self.storage_backend.batch_exists(batch_hashes, extra_info)
hash_value.extend(batch_hashes[:hit_page_num])
@@ -0,0 +1,351 @@
#include <Python.h>
#if defined(__AVX2__)
#include <immintrin.h>
#endif
#include <openssl/sha.h>
#include <pybind11/pybind11.h>
#include <algorithm>
#include <array>
#include <cstdint>
#include <cstring>
#include <limits>
#include <stdexcept>
#include <string>
#include <type_traits>
#include <vector>
namespace py = pybind11;
namespace {
constexpr std::size_t kDigestLen = SHA256_DIGEST_LENGTH;
constexpr std::size_t kHexLen = SHA256_DIGEST_LENGTH * 2;
inline std::uint32_t checked_u32(std::uint64_t value) {
if (value > UINT32_MAX) {
throw std::out_of_range("token id does not fit in uint32");
}
return static_cast<std::uint32_t>(value);
}
inline void digest_to_hex_chars(const unsigned char *digest, char *out) {
static constexpr char kHex[] = "0123456789abcdef";
for (std::size_t i = 0; i < kDigestLen; ++i) {
const unsigned char byte = digest[i];
out[i * 2] = kHex[byte >> 4];
out[i * 2 + 1] = kHex[byte & 0x0f];
}
}
std::string digest_to_hex_string(const unsigned char *digest) {
std::string out(kHexLen, '\0');
digest_to_hex_chars(digest, out.data());
return out;
}
std::array<unsigned char, kDigestLen>
parse_prior_digest(py::object prior_digest_obj, bool *has_prior_digest) {
std::array<unsigned char, kDigestLen> prior_digest{};
*has_prior_digest = false;
if (!prior_digest_obj.is_none()) {
std::string prior = prior_digest_obj.cast<std::string>();
if (prior.size() != kDigestLen) {
throw std::invalid_argument("prior_digest must be exactly 32 bytes");
}
std::copy(prior.begin(), prior.end(), prior_digest.begin());
*has_prior_digest = true;
}
return prior_digest;
}
inline void hash_page(const unsigned char *data, std::size_t len,
bool &has_prior_digest,
std::array<unsigned char, kDigestLen> &prior_digest) {
SHA256_CTX ctx;
SHA256_Init(&ctx);
if (has_prior_digest) {
SHA256_Update(&ctx, prior_digest.data(), prior_digest.size());
}
if (len > 0) {
SHA256_Update(&ctx, data, len);
}
SHA256_Final(prior_digest.data(), &ctx);
has_prior_digest = true;
}
template <typename RawToken>
inline void fill_regular_page(const RawToken *raw, std::size_t start,
std::size_t count, std::uint32_t *out) {
for (std::size_t i = 0; i < count; ++i) {
out[i] = checked_u32(raw[start + i]);
}
}
template <>
inline void
fill_regular_page<std::uint32_t>(const std::uint32_t *raw, std::size_t start,
std::size_t count, std::uint32_t *out) {
std::copy(raw + start, raw + start + count, out);
}
template <>
inline void
fill_regular_page<std::uint64_t>(const std::uint64_t *raw, std::size_t start,
std::size_t count, std::uint32_t *out) {
#if defined(__AVX2__)
std::size_t i = 0;
for (; i + 4 <= count; i += 4) {
const __m256i v =
_mm256_loadu_si256(reinterpret_cast<const __m256i *>(raw + start + i));
const __m256i high = _mm256_srli_epi64(v, 32);
if (!_mm256_testz_si256(high, high)) {
throw std::out_of_range("token id does not fit in uint32");
}
const __m256i low_pairs = _mm256_shuffle_epi32(v, _MM_SHUFFLE(2, 0, 2, 0));
const __m128i lane0 = _mm256_castsi256_si128(low_pairs);
const __m128i lane1 = _mm256_extracti128_si256(low_pairs, 1);
_mm_storel_epi64(reinterpret_cast<__m128i *>(out + i), lane0);
_mm_storel_epi64(reinterpret_cast<__m128i *>(out + i + 2), lane1);
}
for (; i < count; ++i) {
out[i] = checked_u32(raw[start + i]);
}
#else
for (std::size_t i = 0; i < count; ++i) {
out[i] = checked_u32(raw[start + i]);
}
#endif
}
template <typename RawToken>
inline void fill_bigram_page(const RawToken *raw, std::size_t start,
std::size_t count, std::uint32_t *out) {
std::uint32_t prev = checked_u32(raw[start]);
for (std::size_t i = 0; i < count; ++i) {
const std::uint32_t next = checked_u32(raw[start + i + 1]);
out[i * 2] = prev;
out[i * 2 + 1] = next;
prev = next;
}
}
template <>
inline void
fill_bigram_page<std::uint32_t>(const std::uint32_t *raw, std::size_t start,
std::size_t count, std::uint32_t *out) {
std::uint32_t prev = raw[start];
for (std::size_t i = 0; i < count; ++i) {
const std::uint32_t next = raw[start + i + 1];
out[i * 2] = prev;
out[i * 2 + 1] = next;
prev = next;
}
}
template <>
inline void
fill_bigram_page<std::uint64_t>(const std::uint64_t *raw, std::size_t start,
std::size_t count, std::uint32_t *out) {
std::uint32_t prev = checked_u32(raw[start]);
for (std::size_t i = 0; i < count; ++i) {
const std::uint32_t next = checked_u32(raw[start + i + 1]);
out[i * 2] = prev;
out[i * 2 + 1] = next;
prev = next;
}
}
struct RawTokenBuffer {
py::buffer_info info;
std::size_t logical_len;
bool is_bigram;
};
RawTokenBuffer get_raw_token_buffer(const py::buffer &raw_tokens,
std::size_t logical_len,
std::size_t unit_width, bool is_bigram) {
py::buffer_info info = raw_tokens.request();
if (info.ndim != 1) {
throw std::invalid_argument("raw_tokens must be a one-dimensional buffer");
}
if (info.itemsize != 4 && info.itemsize != 8) {
throw std::invalid_argument("raw_tokens itemsize must be 4 or 8 bytes");
}
const std::size_t need_raw_tokens =
is_bigram && logical_len > 0 ? logical_len + 1 : logical_len * unit_width;
if (static_cast<std::size_t>(info.size) < need_raw_tokens) {
throw std::invalid_argument("raw_tokens is shorter than logical_len");
}
return RawTokenBuffer{std::move(info), logical_len, is_bigram};
}
template <typename RawToken>
void hash_pages_to_hex_blob(const RawToken *raw, std::size_t logical_len,
std::size_t page_size, std::size_t unit_width,
bool is_bigram, bool has_prior_digest,
std::array<unsigned char, kDigestLen> prior_digest,
std::string &hex_blob) {
if (page_size == 0) {
throw std::invalid_argument("page_size must be positive");
}
if (is_bigram) {
unit_width = 2;
}
const bool can_hash_raw_bytes =
std::is_same_v<RawToken, std::uint32_t> && !is_bigram;
std::vector<std::uint32_t> page_words;
if (!can_hash_raw_bytes) {
page_words.resize(page_size * unit_width);
}
for (std::size_t start = 0, page_idx = 0; start < logical_len;
start += page_size, ++page_idx) {
const std::size_t page_units = std::min(page_size, logical_len - start);
const std::size_t page_bytes =
page_units * unit_width * sizeof(std::uint32_t);
const unsigned char *bytes = nullptr;
if (can_hash_raw_bytes) {
bytes = reinterpret_cast<const unsigned char *>(raw + start * unit_width);
} else {
if (is_bigram) {
fill_bigram_page(raw, start, page_units, page_words.data());
} else {
fill_regular_page(raw, start * unit_width, page_units * unit_width,
page_words.data());
}
bytes = reinterpret_cast<const unsigned char *>(page_words.data());
}
hash_page(bytes, page_bytes, has_prior_digest, prior_digest);
digest_to_hex_chars(prior_digest.data(),
hex_blob.data() + page_idx * kHexLen);
}
}
template <typename RawToken>
std::string hash_all(const RawToken *raw, std::size_t logical_len,
std::size_t unit_width, bool is_bigram,
bool has_prior_digest,
std::array<unsigned char, kDigestLen> prior_digest) {
if (is_bigram) {
unit_width = 2;
}
const bool can_hash_raw_bytes =
std::is_same_v<RawToken, std::uint32_t> && !is_bigram;
const unsigned char *bytes = nullptr;
std::size_t num_bytes = logical_len * unit_width * sizeof(std::uint32_t);
std::vector<std::uint32_t> words;
if (can_hash_raw_bytes) {
bytes = reinterpret_cast<const unsigned char *>(raw);
} else {
words.resize(logical_len * unit_width);
if (logical_len > 0) {
if (is_bigram) {
fill_bigram_page(raw, 0, logical_len, words.data());
} else {
fill_regular_page(raw, 0, logical_len * unit_width, words.data());
}
}
bytes = reinterpret_cast<const unsigned char *>(words.data());
}
hash_page(bytes, num_bytes, has_prior_digest, prior_digest);
return digest_to_hex_string(prior_digest.data());
}
py::object hex_blob_to_pylist(const std::string &hex_blob) {
const std::size_t num_pages = hex_blob.size() / kHexLen;
if (num_pages >
static_cast<std::size_t>(std::numeric_limits<Py_ssize_t>::max())) {
throw std::overflow_error("too many hash pages");
}
PyObject *raw_list = PyList_New(static_cast<Py_ssize_t>(num_pages));
if (raw_list == nullptr) {
throw py::error_already_set();
}
py::object list = py::reinterpret_steal<py::object>(raw_list);
const char *data = hex_blob.data();
for (std::size_t i = 0; i < num_pages; ++i) {
PyObject *item = PyUnicode_FromStringAndSize(
data + i * kHexLen, static_cast<Py_ssize_t>(kHexLen));
if (item == nullptr) {
throw py::error_already_set();
}
PyList_SET_ITEM(raw_list, static_cast<Py_ssize_t>(i), item);
}
return list;
}
std::string hash_str(const py::buffer &raw_tokens, std::size_t logical_len,
std::size_t unit_width, bool is_bigram,
py::object prior_digest_obj) {
RawTokenBuffer buffer =
get_raw_token_buffer(raw_tokens, logical_len, unit_width, is_bigram);
bool has_prior_digest = false;
auto prior_digest = parse_prior_digest(prior_digest_obj, &has_prior_digest);
py::gil_scoped_release release;
if (buffer.info.itemsize == 4) {
return hash_all(static_cast<const std::uint32_t *>(buffer.info.ptr),
logical_len, unit_width, is_bigram, has_prior_digest,
prior_digest);
}
return hash_all(static_cast<const std::uint64_t *>(buffer.info.ptr),
logical_len, unit_width, is_bigram, has_prior_digest,
prior_digest);
}
py::object pages_hashes(const py::buffer &raw_tokens, std::size_t logical_len,
std::size_t page_size, std::size_t unit_width,
bool is_bigram, py::object prior_digest_obj) {
RawTokenBuffer buffer =
get_raw_token_buffer(raw_tokens, logical_len, unit_width, is_bigram);
const std::size_t num_pages =
page_size == 0 ? 0 : (logical_len + page_size - 1) / page_size;
bool has_prior_digest = false;
auto prior_digest = parse_prior_digest(prior_digest_obj, &has_prior_digest);
std::string hex_blob(num_pages * kHexLen, '\0');
{
py::gil_scoped_release release;
if (buffer.info.itemsize == 4) {
hash_pages_to_hex_blob(
static_cast<const std::uint32_t *>(buffer.info.ptr), logical_len,
page_size, unit_width, is_bigram, has_prior_digest, prior_digest,
hex_blob);
} else {
hash_pages_to_hex_blob(
static_cast<const std::uint64_t *>(buffer.info.ptr), logical_len,
page_size, unit_width, is_bigram, has_prior_digest, prior_digest,
hex_blob);
}
}
return hex_blob_to_pylist(hex_blob);
}
py::object get_hash(const py::buffer &raw_tokens, std::size_t logical_len,
std::size_t unit_width, bool is_bigram,
py::object prior_digest_obj, py::object page_size_obj) {
if (page_size_obj.is_none()) {
return py::cast(hash_str(raw_tokens, logical_len, unit_width, is_bigram,
prior_digest_obj));
}
return pages_hashes(raw_tokens, logical_len,
page_size_obj.cast<std::size_t>(), unit_width, is_bigram,
prior_digest_obj);
}
} // namespace
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("get_hash", &get_hash, py::arg("raw_tokens"), py::arg("logical_len"),
py::arg("unit_width"), py::arg("is_bigram"), py::arg("prior_digest"),
py::arg("page_size") = py::none());
}
@@ -0,0 +1,86 @@
import os
import platform
import sys
from array import array
from functools import lru_cache
from typing import Any, Optional
def _cpu_supports_avx2() -> bool:
if platform.machine().lower() not in ("x86_64", "amd64"):
return False
try:
with open("/proc/cpuinfo", "r", encoding="utf-8", errors="ignore") as f:
return "avx2" in f.read().lower()
except OSError:
return False
@lru_cache(maxsize=1)
def _load_native_hash_module() -> Any:
if sys.byteorder != "little" or not sys.platform.startswith("linux"):
raise RuntimeError(
"HiCache native hash is only supported on little-endian Linux"
)
try:
from torch.utils.cpp_extension import load
abs_path = os.path.dirname(os.path.abspath(__file__))
extra_cflags = ["-O3", "-std=c++17", "-DNDEBUG"]
if _cpu_supports_avx2():
extra_cflags.append("-mavx2")
return load(
name="hicache_hash_cpp",
sources=[f"{abs_path}/hash_binding.cpp"],
extra_cflags=extra_cflags,
extra_ldflags=["-lcrypto"],
with_cuda=False,
verbose=False,
)
except Exception as exc:
raise RuntimeError("Failed to load HiCache native hash extension") from exc
def _native_hash_input(token_ids: Any) -> tuple[array, int, int, bool]:
raw_token_ids = getattr(token_ids, "raw_token_ids", None)
raw = (
raw_token_ids()
if raw_token_ids is not None
else getattr(token_ids, "token_ids", token_ids)
)
logical_len = len(token_ids)
is_bigram = getattr(token_ids, "is_bigram", False)
if isinstance(raw, array) and raw.typecode in ("I", "q", "Q", "L"):
if is_bigram and logical_len > 0 and len(raw) < logical_len + 1:
raise ValueError("bigram token buffer is shorter than logical length")
return raw, logical_len, 2 if is_bigram else 1, is_bigram
if is_bigram:
return array("I", raw[: logical_len + 1]), logical_len, 2, is_bigram
if logical_len == 0:
return array("I"), logical_len, 1, is_bigram
first_token = raw[0]
if isinstance(first_token, tuple):
unit_width = len(first_token)
return (
array("I", (elem for token in raw[:logical_len] for elem in token)),
logical_len,
unit_width,
is_bigram,
)
return array("I", raw[:logical_len]), logical_len, 1, is_bigram
def get_native_hash(
token_ids: Any, prior_digest: Optional[bytes], page_size: Optional[int] = None
) -> str | list[str]:
raw, logical_len, unit_width, is_bigram = _native_hash_input(token_ids)
return _load_native_hash_module().get_hash(
raw, logical_len, unit_width, is_bigram, prior_digest, page_size
)
@@ -585,13 +585,9 @@ class HybridCacheController(BaseHiCacheController):
return operation.id
def _storage_hit_query(self, operation) -> tuple[list[str], int]:
last_hash = operation.last_hash
hash_value = []
for start in range(0, len(operation.token_ids), self.page_size):
last_hash = self.get_hash_str(
operation.token_ids[start : start + self.page_size], last_hash
)
hash_value.append(last_hash)
hash_value = self.get_hash_str(
operation.token_ids, operation.last_hash, page_size=self.page_size
)
extra_info = HiCacheStorageExtraInfo(
prefix_keys=operation.prefix_keys.copy() if operation.prefix_keys else None
+8 -14
View File
@@ -21,7 +21,6 @@ limitations under the License.
The radix tree data structure for managing the KV cache.
"""
import hashlib
import heapq
import logging
import sys
@@ -48,7 +47,11 @@ from sglang.srt.mem_cache.base_prefix_cache import (
)
from sglang.srt.mem_cache.events import KVCacheEventMixin
from sglang.srt.mem_cache.session_radix_cache import SessionRadixCacheMixin
from sglang.srt.mem_cache.utils import get_eviction_strategy, split_node_hash_value
from sglang.srt.mem_cache.utils import (
get_eviction_strategy,
get_hash_str,
split_node_hash_value,
)
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
@@ -206,18 +209,9 @@ class RadixKey:
def hash_page(self, start: int, end: int, prior_hash: Optional[str] = None) -> str:
"""SHA256 for logical units [start, end); bigram mode feeds overlapping (t_i, t_{i+1}) byte pairs."""
hasher = hashlib.sha256()
if prior_hash:
hasher.update(bytes.fromhex(prior_hash))
t = self.token_ids
if self.is_bigram:
for j in range(start, end):
hasher.update(t[j].to_bytes(4, byteorder="little", signed=False))
hasher.update(t[j + 1].to_bytes(4, byteorder="little", signed=False))
else:
for j in range(start, end):
hasher.update(t[j].to_bytes(4, byteorder="little", signed=False))
return hasher.hexdigest()
hash_value = get_hash_str(self[start:end], prior_hash)
assert isinstance(hash_value, str)
return hash_value
class TreeNode:
+10 -27
View File
@@ -13,10 +13,10 @@
# ==============================================================================
"""Common utilities."""
import hashlib
from typing import Any, Callable, List, Optional, Tuple
from sglang.srt.environ import envs
from sglang.srt.mem_cache.cpp_utils.native_hash import get_native_hash
from sglang.srt.mem_cache.evict_policy import (
EvictionStrategy,
FIFOStrategy,
@@ -103,22 +103,13 @@ def maybe_init_custom_mem_pool(
return False, None, None
def get_hash_str(token_ids: List[int], prior_hash: Optional[str] = None) -> str:
hasher = hashlib.sha256()
if prior_hash:
hasher.update(bytes.fromhex(prior_hash))
for t in token_ids:
if isinstance(t, tuple):
# EAGLE bigram mode: hash both elements to uniquely identify the bigram
for elem in t:
hasher.update(elem.to_bytes(4, byteorder="little", signed=False))
else:
# Regular mode: single integer token
hasher.update(t.to_bytes(4, byteorder="little", signed=False))
return hasher.hexdigest()
def get_hash_str(
token_ids: List[int],
prior_hash: Optional[str] = None,
page_size: Optional[int] = None,
) -> str | List[str]:
prior_digest = bytes.fromhex(prior_hash) if prior_hash else None
return get_native_hash(token_ids, prior_digest, page_size)
def hash_str_to_int64(hash_str: str) -> int:
@@ -134,21 +125,13 @@ def hash_str_to_int64(hash_str: str) -> int:
def compute_node_hash_values(node: Any, page_size: int) -> List[str]:
"""Compute SHA256-based hash values for position-aware KV block IDs."""
hash_values = []
parent_hash = None
if node.parent is not None and node.parent.hash_value is not None:
if len(node.parent.key) > 0 and len(node.parent.hash_value) > 0:
parent_hash = node.parent.hash_value[-1]
logical_len = len(node.key)
for start in range(0, logical_len, page_size):
end = min(start + page_size, logical_len)
if end <= start:
continue
hash_val = node.key.hash_page(start, end, parent_hash)
hash_values.append(hash_val)
parent_hash = hash_val
hash_values = get_hash_str(node.key, parent_hash, page_size=page_size)
assert isinstance(hash_values, list)
return hash_values
@@ -1,7 +1,10 @@
"""Unit tests for mem_cache/utils.py — no server, no model loading."""
import hashlib
import sys
import types
import unittest
from array import array
from unittest.mock import MagicMock, patch
from sglang.srt.mem_cache.evict_policy import (
@@ -22,12 +25,117 @@ from sglang.srt.mem_cache.utils import (
split_node_hash_value,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=8, suite="base-a-test-cpu")
class TestGetEvictionStrategy(CustomTestCase):
def _legacy_get_hash_str(token_ids, prior_hash=None):
hasher = hashlib.sha256()
if prior_hash:
hasher.update(bytes.fromhex(prior_hash))
for t in token_ids:
if isinstance(t, tuple):
for elem in t:
hasher.update(elem.to_bytes(4, byteorder="little", signed=False))
else:
hasher.update(t.to_bytes(4, byteorder="little", signed=False))
return hasher.hexdigest()
def _legacy_page_hashes(key, page_size, prior_hash=None):
hashes = []
running_hash = prior_hash
for start in range(0, len(key), page_size):
running_hash = _legacy_get_hash_str(
key[start : start + page_size], running_hash
)
hashes.append(running_hash)
return hashes
class _HashKey:
def __init__(self, token_ids, is_bigram=False):
self.token_ids = token_ids
self.is_bigram = is_bigram
def __len__(self):
if self.is_bigram:
return max(0, len(self.token_ids) - 1)
return len(self.token_ids)
def __getitem__(self, index):
if isinstance(index, slice):
start = index.start or 0
stop = index.stop if index.stop is not None else len(self)
if self.is_bigram:
return _HashKey(self.token_ids[start : stop + 1], is_bigram=True)
return _HashKey(self.token_ids[start:stop])
if self.is_bigram:
return (self.token_ids[index], self.token_ids[index + 1])
return self.token_ids[index]
def raw_token_ids(self):
return self.token_ids
def hash_page(self, start, end, prior_hash=None):
return _legacy_get_hash_str(self[start:end], prior_hash)
def _single_hash_compatibility_cases():
prior_hash = _legacy_get_hash_str([7, 8, 9])
return [
("empty_list", [], None),
("plain_list", [1, 2, 3, 4, 5], None),
("array_q", _HashKey(array("q", range(1, 258))), None),
("array_i_with_prior", _HashKey(array("I", range(1, 258))), prior_hash),
(
"tuple_bigram_with_prior",
[(10, 20), (20, 30), (30, 40), (40, 50)],
prior_hash,
),
(
"eagle_bigram_with_prior",
_HashKey(
array("q", ((i * 2654435761) & 0x00FFFFFF for i in range(258))),
is_bigram=True,
),
prior_hash,
),
("empty_bigram", _HashKey(array("q"), is_bigram=True), None),
]
def _page_hash_compatibility_cases():
prior_hash = _legacy_get_hash_str([7, 8, 9])
return [
("empty_list", [], 8, None),
("empty_bigram", _HashKey(array("q"), is_bigram=True), 8, None),
("array_q_page_64", _HashKey(array("q", range(1, 258))), 64, None),
(
"array_i_page_64_with_prior",
_HashKey(array("I", range(1, 258))),
64,
prior_hash,
),
(
"eagle_bigram_page_64_with_prior",
_HashKey(
array("q", ((i * 2654435761) & 0x00FFFFFF for i in range(258))),
is_bigram=True,
),
64,
prior_hash,
),
(
"eagle_bigram_page_1_with_prior",
_HashKey(array("q", [11, 22, 33, 44, 55]), is_bigram=True),
1,
prior_hash,
),
]
class TestGetEvictionStrategy(unittest.TestCase):
def test_lru(self):
self.assertIsInstance(get_eviction_strategy("lru"), LRUStrategy)
@@ -73,7 +181,7 @@ class TestGetEvictionStrategy(CustomTestCase):
self.assertIsNot(s1, s2)
class TestMaybeInitCustomMemPool(CustomTestCase):
class TestMaybeInitCustomMemPool(unittest.TestCase):
@patch("sglang.srt.mem_cache.utils.envs.SGLANG_MOONCAKE_CUSTOM_MEM_POOL.get")
def test_disabled_by_default(self, mock_env_get):
mock_env_get.return_value = None
@@ -83,69 +191,76 @@ class TestMaybeInitCustomMemPool(CustomTestCase):
self.assertIsNone(pool_type)
@patch("sglang.srt.mem_cache.utils.envs.SGLANG_MOONCAKE_CUSTOM_MEM_POOL.get")
@patch("sglang.srt.disaggregation.mooncake.utils.init_mooncake_custom_mem_pool")
def test_enabled_via_env(self, mock_init, mock_env_get):
def test_enabled_via_env(self, mock_env_get):
mock_env_get.return_value = "enabled"
mock_init = MagicMock()
mock_init.return_value = (True, "mock_pool_instance", "mooncake")
enabled, pool, pool_type = maybe_init_custom_mem_pool("cuda:0")
mooncake_pkg = types.ModuleType("sglang.srt.disaggregation.mooncake")
mooncake_utils = types.ModuleType("sglang.srt.disaggregation.mooncake.utils")
mooncake_utils.init_mooncake_custom_mem_pool = mock_init
with patch.dict(
sys.modules,
{
"sglang.srt.disaggregation.mooncake": mooncake_pkg,
"sglang.srt.disaggregation.mooncake.utils": mooncake_utils,
},
):
enabled, pool, pool_type = maybe_init_custom_mem_pool("cuda:0")
self.assertTrue(enabled)
self.assertEqual(pool, "mock_pool_instance")
self.assertEqual(pool_type, "mooncake")
mock_init.assert_called_once_with("cuda:0")
class TestGetHashStr(CustomTestCase):
def test_empty_list(self):
result = get_hash_str([])
self.assertIsInstance(result, str)
self.assertEqual(len(result), 64)
expected = hashlib.sha256().hexdigest()
self.assertEqual(result, expected)
class TestGetHashStr(unittest.TestCase):
def test_hash_str_matches_pre_optimization_per_token_loop(self):
for name, tokens, prior_hash in _single_hash_compatibility_cases():
with self.subTest(name=name):
self.assertEqual(
get_hash_str(tokens, prior_hash),
_legacy_get_hash_str(tokens, prior_hash),
)
def test_page_hashes_match_pre_optimization_per_token_loop(self):
for name, tokens, page_size, prior_hash in _page_hash_compatibility_cases():
with self.subTest(name=name):
self.assertEqual(
get_hash_str(tokens, prior_hash, page_size=page_size),
_legacy_page_hashes(tokens, page_size, prior_hash),
)
def test_hash_properties(self):
self.assertEqual(get_hash_str([]), hashlib.sha256().hexdigest())
def test_different_sequence_different_hash(self):
h1 = get_hash_str([1, 2, 3])
h2 = get_hash_str([3, 2, 1])
self.assertNotEqual(h1, h2)
def test_different_values_different_hash(self):
h1 = get_hash_str([100])
h2 = get_hash_str([200])
self.assertNotEqual(h1, h2)
self.assertNotEqual(get_hash_str([100]), get_hash_str([200]))
self.assertEqual(get_hash_str([1, 2]), get_hash_str([(1, 2)]))
self.assertNotEqual(get_hash_str([(1, 2)]), get_hash_str([(2, 1)]))
def test_bigram_vs_flat_token_equivalent(self):
h_flat = get_hash_str([1, 2])
h_bigram = get_hash_str([(1, 2)])
self.assertEqual(h_flat, h_bigram)
def test_bigram_order_matters(self):
h1 = get_hash_str([(1, 2)])
h2 = get_hash_str([(2, 1)])
self.assertNotEqual(h1, h2)
def test_prior_hash_chaining(self):
chained = get_hash_str([3, 4], prior_hash=get_hash_str([1, 2]))
# prior_hash must fold into the digest, so chaining differs from
# hashing [3, 4] alone...
self.assertNotEqual(chained, get_hash_str([3, 4]))
# ...and a different prior_hash must yield a different chained digest.
self.assertNotEqual(
chained, get_hash_str([3, 4], prior_hash=get_hash_str([9, 9]))
)
def test_prior_hash_single_step(self):
step1 = get_hash_str([1])
step2 = get_hash_str([2], prior_hash=step1)
direct = get_hash_str([1, 2])
self.assertNotEqual(step2, direct)
def test_returns_64_char_hex(self):
for tokens in [[], [1], [1, 2, 3], [(1, 2)], [1, 2, 3, 4, 5]]:
result = get_hash_str(tokens)
self.assertRegex(result, r"^[0-9a-f]{64}$")
with self.subTest(tokens=tokens):
self.assertRegex(get_hash_str(tokens), r"^[0-9a-f]{64}$")
def test_hash_key_hash_page_matches_get_hash_str(self):
key = _HashKey(array("q", [1, 2, 3, 4, 5, 6]), is_bigram=True)
prior_hash = get_hash_str([(9, 10)])
self.assertEqual(
key.hash_page(1, 4, prior_hash),
get_hash_str(key[1:4], prior_hash),
)
class TestHashStrToInt64(CustomTestCase):
class TestHashStrToInt64(unittest.TestCase):
def test_zero_hash(self):
result = hash_str_to_int64("0" * 64)
self.assertEqual(result, 0)
@@ -178,96 +293,72 @@ class TestHashStrToInt64(CustomTestCase):
self.assertTrue(-(2**63) <= int64_val < 2**63)
class TestComputeNodeHashValues(CustomTestCase):
def setUp(self):
def mock_hash_page(start, end, parent_hash):
parts = [f"p{start}-{end}"]
if parent_hash is not None:
parts.append(parent_hash)
return "-".join(parts)
self.mock_hash_page = mock_hash_page
def _make_node(self, key_len, parent=None, parent_hash_values=None):
class TestComputeNodeHashValues(unittest.TestCase):
def _make_node(self, key, parent=None, parent_hash_values=None):
node = MagicMock()
node.key.__len__.return_value = key_len
node.key.hash_page = self.mock_hash_page
node.key = key
node.parent = parent
if parent is not None:
parent.hash_value = parent_hash_values
return node
def test_single_page_root(self):
node = self._make_node(key_len=3)
result = compute_node_hash_values(node, page_size=16)
self.assertEqual(len(result), 1)
self.assertIn("p0-3", result[0])
def test_root_node_hashes_match_legacy_page_hashes(self):
cases = [
("single_page", _HashKey(array("q", [1, 2, 3])), 16),
("multiple_pages", _HashKey(array("q", range(1, 31))), 16),
("page_aligned_boundary", _HashKey(array("q", range(1, 33))), 8),
("shorter_than_page", _HashKey(array("q", [1, 2, 3, 4, 5])), 16),
]
def test_multiple_pages(self):
node = self._make_node(key_len=30)
result = compute_node_hash_values(node, page_size=16)
self.assertEqual(len(result), 2)
self.assertIn("p0-16", result[0])
self.assertIn("p16-30", result[1])
for name, key, page_size in cases:
with self.subTest(name=name):
node = self._make_node(key)
self.assertEqual(
compute_node_hash_values(node, page_size=page_size),
_legacy_page_hashes(key, page_size=page_size),
)
def test_page_aligned_boundary(self):
node = self._make_node(key_len=32)
result = compute_node_hash_values(node, page_size=8)
self.assertEqual(len(result), 4)
self.assertIn("p24-32", result[3])
def test_key_shorter_than_page_size(self):
node = self._make_node(key_len=5)
result = compute_node_hash_values(node, page_size=16)
self.assertEqual(result, ["p0-5"])
def test_chained_parent_hash(self):
def test_parent_hash_is_used_only_when_parent_has_nonempty_key_and_hash(self):
parent = MagicMock()
parent.key.__len__.return_value = 8
parent.hash_value = ["parent_hash_0", "parent_hash_1"]
parent.key.hash_page = self.mock_hash_page
parent.key = _HashKey(array("q", range(1, 17)))
parent.hash_value = _legacy_page_hashes(parent.key, page_size=8)
child_key = _HashKey(array("q", range(101, 109)))
cases = [
("valid_parent_hash", parent, parent.hash_value, parent.hash_value[-1]),
(
"empty_parent_key",
self._make_node(_HashKey(array("q"))),
[get_hash_str([1, 2, 3])],
None,
),
(
"empty_parent_hash",
self._make_node(_HashKey(array("q", range(1, 9)))),
[],
None,
),
(
"none_parent_hash",
self._make_node(_HashKey(array("q", range(1, 9)))),
None,
None,
),
]
child = self._make_node(
key_len=16, parent=parent, parent_hash_values=parent.hash_value
)
result = compute_node_hash_values(child, page_size=8)
self.assertEqual(result, ["p0-8-parent_hash_1", "p8-16-p0-8-parent_hash_1"])
def test_parent_with_empty_key(self):
parent = MagicMock()
parent.key.__len__.return_value = 0
parent.hash_value = ["some_hash"]
parent.key.hash_page = self.mock_hash_page
child = self._make_node(
key_len=8, parent=parent, parent_hash_values=parent.hash_value
)
result = compute_node_hash_values(child, page_size=8)
self.assertEqual(len(result), 1)
self.assertNotIn("some_hash", result[0])
def test_parent_without_hash_value(self):
parent = MagicMock()
parent.key.__len__.return_value = 8
parent.hash_value = []
parent.key.hash_page = self.mock_hash_page
child = self._make_node(key_len=8, parent=parent, parent_hash_values=[])
result = compute_node_hash_values(child, page_size=8)
self.assertEqual(result, ["p0-8"])
def test_parent_with_none_hash_value(self):
parent = MagicMock()
parent.key.__len__.return_value = 8
parent.hash_value = None
parent.key.hash_page = self.mock_hash_page
child = self._make_node(key_len=8, parent=parent, parent_hash_values=None)
result = compute_node_hash_values(child, page_size=8)
self.assertEqual(result, ["p0-8"])
for name, parent, parent_hash_values, expected_prior in cases:
with self.subTest(name=name):
child = self._make_node(
child_key, parent=parent, parent_hash_values=parent_hash_values
)
self.assertEqual(
compute_node_hash_values(child, page_size=8),
_legacy_page_hashes(
child_key, page_size=8, prior_hash=expected_prior
),
)
class TestSplitNodeHashValue(CustomTestCase):
class TestSplitNodeHashValue(unittest.TestCase):
def test_none_input_returns_none_tuple(self):
result = split_node_hash_value(None, 10, 4)
self.assertEqual(result, (None, None))