From 5e1ccd932053427595f2e9b92dc70e7cb0e99ac9 Mon Sep 17 00:00:00 2001 From: huangtingwei <141888744+huangtingwei9988@users.noreply.github.com> Date: Wed, 1 Jul 2026 15:44:23 +0800 Subject: [PATCH] [HiCache] Optimize HiCache hash generation with bulk token byte conversion (#28287) --- .../sglang/srt/managers/cache_controller.py | 16 +- .../srt/mem_cache/cpp_utils/hash_binding.cpp | 351 ++++++++++++++++++ .../srt/mem_cache/cpp_utils/native_hash.py | 86 +++++ .../hybrid_cache/hybrid_cache_controller.py | 10 +- python/sglang/srt/mem_cache/radix_cache.py | 22 +- python/sglang/srt/mem_cache/utils.py | 37 +- .../unit/mem_cache/test_mem_cache_utils.py | 333 +++++++++++------ 7 files changed, 675 insertions(+), 180 deletions(-) create mode 100644 python/sglang/srt/mem_cache/cpp_utils/hash_binding.cpp create mode 100644 python/sglang/srt/mem_cache/cpp_utils/native_hash.py diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index f6138f6d5..7505f5fe7 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -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]) diff --git a/python/sglang/srt/mem_cache/cpp_utils/hash_binding.cpp b/python/sglang/srt/mem_cache/cpp_utils/hash_binding.cpp new file mode 100644 index 000000000..d415bdf78 --- /dev/null +++ b/python/sglang/srt/mem_cache/cpp_utils/hash_binding.cpp @@ -0,0 +1,351 @@ +#include +#if defined(__AVX2__) +#include +#endif +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +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(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 +parse_prior_digest(py::object prior_digest_obj, bool *has_prior_digest) { + std::array prior_digest{}; + *has_prior_digest = false; + if (!prior_digest_obj.is_none()) { + std::string prior = prior_digest_obj.cast(); + 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 &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 +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(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(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(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 +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(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(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(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 +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 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 && !is_bigram; + std::vector 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(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(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 +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 prior_digest) { + if (is_bigram) { + unit_width = 2; + } + const bool can_hash_raw_bytes = + std::is_same_v && !is_bigram; + const unsigned char *bytes = nullptr; + std::size_t num_bytes = logical_len * unit_width * sizeof(std::uint32_t); + + std::vector words; + if (can_hash_raw_bytes) { + bytes = reinterpret_cast(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(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::numeric_limits::max())) { + throw std::overflow_error("too many hash pages"); + } + + PyObject *raw_list = PyList_New(static_cast(num_pages)); + if (raw_list == nullptr) { + throw py::error_already_set(); + } + py::object list = py::reinterpret_steal(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(kHexLen)); + if (item == nullptr) { + throw py::error_already_set(); + } + PyList_SET_ITEM(raw_list, static_cast(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(buffer.info.ptr), + logical_len, unit_width, is_bigram, has_prior_digest, + prior_digest); + } + return hash_all(static_cast(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(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(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(), 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()); +} diff --git a/python/sglang/srt/mem_cache/cpp_utils/native_hash.py b/python/sglang/srt/mem_cache/cpp_utils/native_hash.py new file mode 100644 index 000000000..9eb8c37d4 --- /dev/null +++ b/python/sglang/srt/mem_cache/cpp_utils/native_hash.py @@ -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 + ) diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py index 011a877e7..2a406596e 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py @@ -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 diff --git a/python/sglang/srt/mem_cache/radix_cache.py b/python/sglang/srt/mem_cache/radix_cache.py index 28ff718fa..6ebf09939 100644 --- a/python/sglang/srt/mem_cache/radix_cache.py +++ b/python/sglang/srt/mem_cache/radix_cache.py @@ -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: diff --git a/python/sglang/srt/mem_cache/utils.py b/python/sglang/srt/mem_cache/utils.py index 9b4aaed2d..85d30fbd7 100644 --- a/python/sglang/srt/mem_cache/utils.py +++ b/python/sglang/srt/mem_cache/utils.py @@ -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 diff --git a/test/registered/unit/mem_cache/test_mem_cache_utils.py b/test/registered/unit/mem_cache/test_mem_cache_utils.py index 73b61b5c4..b2ff08dcf 100644 --- a/test/registered/unit/mem_cache/test_mem_cache_utils.py +++ b/test/registered/unit/mem_cache/test_mem_cache_utils.py @@ -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))