[HiCache][LoRA] Isolate storage pages by extra key (#38577)
Co-authored-by: Shuwen Wang <47200617+alphabetc1@users.noreply.github.com>
This commit is contained in:
co-authored by
Shuwen Wang
parent
b9cb96496d
commit
0b415fa573
@@ -90,6 +90,8 @@ class DecodeHiCachePreallocMixin:
|
|||||||
suffix_tokens,
|
suffix_tokens,
|
||||||
last_hash,
|
last_hash,
|
||||||
prefix_keys,
|
prefix_keys,
|
||||||
|
extra_key=req.extra_key,
|
||||||
|
cache_salt=req.cache_salt,
|
||||||
)
|
)
|
||||||
|
|
||||||
return DecodePrefixMatch(
|
return DecodePrefixMatch(
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ from sglang.srt.mem_cache.memory_pool import (
|
|||||||
MLATokenToKVPool,
|
MLATokenToKVPool,
|
||||||
ReqToTokenPool,
|
ReqToTokenPool,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.mem_cache.utils import storage_namespace_seed
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_memory,
|
get_memory,
|
||||||
get_schedule,
|
get_schedule,
|
||||||
@@ -138,7 +139,9 @@ class DecodeKVCacheOffloadManager:
|
|||||||
state = self.offloaded_state.get(req)
|
state = self.offloaded_state.get(req)
|
||||||
if state is None:
|
if state is None:
|
||||||
prefill_hashes = self._compute_prefix_hash(
|
prefill_hashes = self._compute_prefix_hash(
|
||||||
req.origin_input_ids[:prefill_offloaded_len]
|
req.origin_input_ids[:prefill_offloaded_len],
|
||||||
|
extra_key=req.extra_key,
|
||||||
|
cache_salt=req.cache_salt,
|
||||||
)
|
)
|
||||||
last_prefill_hash = (
|
last_prefill_hash = (
|
||||||
prefill_hashes[-1] if prefill_offloaded_len > 0 else None
|
prefill_hashes[-1] if prefill_offloaded_len > 0 else None
|
||||||
@@ -271,7 +274,12 @@ class DecodeKVCacheOffloadManager:
|
|||||||
self, req, host_indices, incremental_tokens, start_time, prior_hash
|
self, req, host_indices, incremental_tokens, start_time, prior_hash
|
||||||
):
|
):
|
||||||
"""Trigger async backup from host to storage."""
|
"""Trigger async backup from host to storage."""
|
||||||
page_hashes = self._compute_prefix_hash(incremental_tokens, prior_hash)
|
page_hashes = self._compute_prefix_hash(
|
||||||
|
incremental_tokens,
|
||||||
|
prior_hash,
|
||||||
|
extra_key=req.extra_key,
|
||||||
|
cache_salt=req.cache_salt,
|
||||||
|
)
|
||||||
ack_id = self.cache_controller.write_storage(
|
ack_id = self.cache_controller.write_storage(
|
||||||
host_indices,
|
host_indices,
|
||||||
incremental_tokens,
|
incremental_tokens,
|
||||||
@@ -280,9 +288,12 @@ class DecodeKVCacheOffloadManager:
|
|||||||
self.ongoing_backup[ack_id] = (req.rid, host_indices, start_time)
|
self.ongoing_backup[ack_id] = (req.rid, host_indices, start_time)
|
||||||
return page_hashes[-1] if len(page_hashes) > 0 else prior_hash
|
return page_hashes[-1] if len(page_hashes) > 0 else prior_hash
|
||||||
|
|
||||||
def _compute_prefix_hash(self, tokens, prior_hash=""):
|
def _compute_prefix_hash(
|
||||||
|
self, tokens, prior_hash="", extra_key=None, cache_salt=None
|
||||||
|
):
|
||||||
|
"""Match prefill storage hashes."""
|
||||||
page_hashes = []
|
page_hashes = []
|
||||||
last_hash = prior_hash
|
last_hash = prior_hash or storage_namespace_seed(extra_key, cache_salt)
|
||||||
for offset in range(0, len(tokens), self.page_size):
|
for offset in range(0, len(tokens), self.page_size):
|
||||||
page_tokens = tokens[offset : offset + self.page_size]
|
page_tokens = tokens[offset : offset + self.page_size]
|
||||||
last_hash = self.cache_controller.get_hash_str(page_tokens, last_hash)
|
last_hash = self.cache_controller.get_hash_str(page_tokens, last_hash)
|
||||||
|
|||||||
@@ -42,6 +42,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.l2_transfer import L2Transfer, L2TransferEngine
|
from sglang.srt.mem_cache.l2_transfer import L2Transfer, L2TransferEngine
|
||||||
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool
|
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool
|
||||||
|
from sglang.srt.mem_cache.utils import get_storage_hash_str
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import get_device_module
|
from sglang.srt.utils import get_device_module
|
||||||
|
|
||||||
@@ -1197,7 +1198,7 @@ class HiCacheController:
|
|||||||
|
|
||||||
storage_query_count = 0
|
storage_query_count = 0
|
||||||
hash_value = []
|
hash_value = []
|
||||||
page_hashes = self.get_hash_str(
|
page_hashes = get_storage_hash_str(
|
||||||
tokens_to_fetch, last_hash, page_size=self.page_size
|
tokens_to_fetch, last_hash, page_size=self.page_size
|
||||||
)
|
)
|
||||||
operation.all_hash_values = page_hashes
|
operation.all_hash_values = page_hashes
|
||||||
|
|||||||
@@ -89,9 +89,9 @@ class KVCacheEventRecorder:
|
|||||||
"""Hash values to publish for ``node``, computing them if not yet set."""
|
"""Hash values to publish for ``node``, computing them if not yet set."""
|
||||||
if node.hash_value is None:
|
if node.hash_value is None:
|
||||||
node.hash_value = compute_node_hash_values(node, self.page_size)
|
node.hash_value = compute_node_hash_values(node, self.page_size)
|
||||||
if node.key.cache_salt is not None:
|
if node.key.extra_key is None and node.key.cache_salt is None:
|
||||||
return compute_node_event_hash_values(node, self.page_size)
|
return node.hash_value
|
||||||
return node.hash_value
|
return compute_node_event_hash_values(node, self.page_size)
|
||||||
|
|
||||||
def _parent_block_hash(self, node: Any) -> Optional[int]:
|
def _parent_block_hash(self, node: Any) -> Optional[int]:
|
||||||
"""The hash the first page of ``node`` links back to.
|
"""The hash the first page of ``node`` links back to.
|
||||||
@@ -103,7 +103,7 @@ class KVCacheEventRecorder:
|
|||||||
parent = node.parent
|
parent = node.parent
|
||||||
if parent is None or parent.parent is None:
|
if parent is None or parent.parent is None:
|
||||||
return None
|
return None
|
||||||
if node.key.cache_salt is not None:
|
if node.key.extra_key is not None or node.key.cache_salt is not None:
|
||||||
parent_hash_values = parent.event_hash_value
|
parent_hash_values = parent.event_hash_value
|
||||||
assert parent_hash_values is not None
|
assert parent_hash_values is not None
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -1488,15 +1488,17 @@ class HiRadixCache(RadixCache):
|
|||||||
new_input_tokens: List[int],
|
new_input_tokens: List[int],
|
||||||
last_hash: Optional[str] = None,
|
last_hash: Optional[str] = None,
|
||||||
prefix_keys: Optional[List[str]] = None,
|
prefix_keys: Optional[List[str]] = None,
|
||||||
|
extra_key: Optional[str] = None,
|
||||||
|
cache_salt: Optional[str] = None,
|
||||||
) -> int:
|
) -> int:
|
||||||
if not self.enable_storage or self.cache_controller.prefetch_rate_limited():
|
if not self.enable_storage or self.cache_controller.prefetch_rate_limited():
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
prefetch_key = RadixKey(
|
prefetch_key = RadixKey(
|
||||||
new_input_tokens,
|
new_input_tokens,
|
||||||
extra_key=last_host_node.key.extra_key,
|
extra_key=extra_key,
|
||||||
is_bigram=self.is_eagle,
|
is_bigram=self.is_eagle,
|
||||||
cache_salt=last_host_node.key.cache_salt,
|
cache_salt=cache_salt,
|
||||||
).page_aligned(self.page_size)
|
).page_aligned(self.page_size)
|
||||||
if len(prefetch_key) < self.prefetch_threshold:
|
if len(prefetch_key) < self.prefetch_threshold:
|
||||||
return 0
|
return 0
|
||||||
@@ -1773,16 +1775,14 @@ class HiRadixCache(RadixCache):
|
|||||||
prefix_keys: Optional[List[str]] = None,
|
prefix_keys: Optional[List[str]] = None,
|
||||||
# Scheduler-call parity with UnifiedRadixCache; unused in cache mode.
|
# Scheduler-call parity with UnifiedRadixCache; unused in cache mode.
|
||||||
matched_prefix_tokens: Optional[List[int]] = None,
|
matched_prefix_tokens: Optional[List[int]] = None,
|
||||||
# Cache mode write-through keeps the anchor on the request's own path,
|
|
||||||
# so the namespace is already carried by ``last_host_node.key``.
|
|
||||||
extra_key: Optional[str] = None,
|
extra_key: Optional[str] = None,
|
||||||
cache_salt: Optional[str] = None,
|
cache_salt: Optional[str] = None,
|
||||||
):
|
):
|
||||||
prefetch_key = RadixKey(
|
prefetch_key = RadixKey(
|
||||||
new_input_tokens,
|
new_input_tokens,
|
||||||
extra_key=last_host_node.key.extra_key,
|
extra_key=extra_key,
|
||||||
is_bigram=self.is_eagle,
|
is_bigram=self.is_eagle,
|
||||||
cache_salt=last_host_node.key.cache_salt,
|
cache_salt=cache_salt,
|
||||||
)
|
)
|
||||||
# align the number of fetching tokens to the page size
|
# align the number of fetching tokens to the page size
|
||||||
prefetch_key = prefetch_key.page_aligned(self.page_size)
|
prefetch_key = prefetch_key.page_aligned(self.page_size)
|
||||||
|
|||||||
@@ -39,6 +39,8 @@ from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||||
|
|
||||||
|
from sglang.srt.mem_cache.utils import get_storage_hash_str
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@@ -576,7 +578,7 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
return operation.id
|
return operation.id
|
||||||
|
|
||||||
def _storage_hit_query(self, operation) -> tuple[list[str], int]:
|
def _storage_hit_query(self, operation) -> tuple[list[str], int]:
|
||||||
hash_value = self.get_hash_str(
|
hash_value = get_storage_hash_str(
|
||||||
operation.token_ids, operation.last_hash, page_size=self.page_size
|
operation.token_ids, operation.last_hash, page_size=self.page_size
|
||||||
)
|
)
|
||||||
operation.all_hash_values = hash_value
|
operation.all_hash_values = hash_value
|
||||||
|
|||||||
@@ -71,12 +71,9 @@ class RadixKey:
|
|||||||
):
|
):
|
||||||
# token ids sequence (raw ints in both modes)
|
# token ids sequence (raw ints in both modes)
|
||||||
self.token_ids = token_ids
|
self.token_ids = token_ids
|
||||||
# Extra key for caller-defined cache classification.
|
# Namespaces the tree and storage; omitted from KV events.
|
||||||
self.extra_key = extra_key
|
self.extra_key = extra_key
|
||||||
# Cache salt is kept distinct so it cannot collide with extra_key.
|
# Namespaces the tree, storage and KV events.
|
||||||
# It namespaces the in-process radix tree and external KV events;
|
|
||||||
# external L3/remote storage keys remain token-only and are outside
|
|
||||||
# this contract.
|
|
||||||
self.cache_salt = cache_salt or None
|
self.cache_salt = cache_salt or None
|
||||||
# bigram view over token_ids: length = max(0, len(token_ids) - 1)
|
# bigram view over token_ids: length = max(0, len(token_ids) - 1)
|
||||||
self.is_bigram = is_bigram
|
self.is_bigram = is_bigram
|
||||||
|
|||||||
@@ -42,7 +42,7 @@ from sglang.srt.mem_cache.unified_cache.components import (
|
|||||||
LinkerTransferPhase,
|
LinkerTransferPhase,
|
||||||
TreeComponent,
|
TreeComponent,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.utils import get_hash_str
|
from sglang.srt.mem_cache.utils import get_storage_hash_str
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
@@ -279,7 +279,7 @@ class UnifiedCacheLinkerWrapper:
|
|||||||
tail_len = (len(key) - device_hit_len) // page * page
|
tail_len = (len(key) - device_hit_len) // page * page
|
||||||
if tail_len == 0:
|
if tail_len == 0:
|
||||||
return []
|
return []
|
||||||
return get_hash_str(
|
return get_storage_hash_str(
|
||||||
key[device_hit_len : device_hit_len + tail_len],
|
key[device_hit_len : device_hit_len + tail_len],
|
||||||
last_hash,
|
last_hash,
|
||||||
page_size=page,
|
page_size=page,
|
||||||
|
|||||||
@@ -1829,8 +1829,10 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
new_input_tokens: list[int],
|
new_input_tokens: list[int],
|
||||||
last_hash: Optional[str] = None,
|
last_hash: Optional[str] = None,
|
||||||
prefix_keys: Optional[list[str]] = None,
|
prefix_keys: Optional[list[str]] = None,
|
||||||
|
extra_key: Optional[str] = None,
|
||||||
|
cache_salt: Optional[str] = None,
|
||||||
) -> int:
|
) -> int:
|
||||||
"""Synchronously probe L3 storage for the reusable prefix length."""
|
"""Probe L3 with the request namespace."""
|
||||||
if (
|
if (
|
||||||
not self.enable_storage
|
not self.enable_storage
|
||||||
or self.cache_controller is None
|
or self.cache_controller is None
|
||||||
@@ -1838,7 +1840,6 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
):
|
):
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
extra_key, cache_salt = self.tree_core.prefetch_anchor_info(last_host_node_id)
|
|
||||||
prefetch_key = RadixKey(
|
prefetch_key = RadixKey(
|
||||||
new_input_tokens,
|
new_input_tokens,
|
||||||
extra_key=extra_key,
|
extra_key=extra_key,
|
||||||
|
|||||||
@@ -120,6 +120,36 @@ def get_hash_str(
|
|||||||
return get_native_hash(token_ids, prior_digest, page_size)
|
return get_native_hash(token_ids, prior_digest, page_size)
|
||||||
|
|
||||||
|
|
||||||
|
def storage_namespace_seed(
|
||||||
|
extra_key: Optional[str], cache_salt: Optional[str]
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""Seed storage chains; preserve unnamespaced keys and Rust byte parity."""
|
||||||
|
if extra_key is None and cache_salt is None:
|
||||||
|
return None
|
||||||
|
digest = hashlib.sha256(b"sglang-cache-namespace-v1")
|
||||||
|
# Presence and UTF-8 byte length distinguish absent, empty and joined parts.
|
||||||
|
for part in (extra_key, cache_salt):
|
||||||
|
if part is None:
|
||||||
|
digest.update(b"\x00")
|
||||||
|
continue
|
||||||
|
encoded = part.encode("utf-8")
|
||||||
|
digest.update(b"\x01" + len(encoded).to_bytes(8, "little") + encoded)
|
||||||
|
return digest.hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
def get_storage_hash_str(
|
||||||
|
key: Any,
|
||||||
|
prior_hash: Optional[str] = None,
|
||||||
|
page_size: Optional[int] = None,
|
||||||
|
) -> str | List[str]:
|
||||||
|
"""Seed new storage chains with the request namespace."""
|
||||||
|
if prior_hash is None:
|
||||||
|
prior_hash = storage_namespace_seed(
|
||||||
|
getattr(key, "extra_key", None), getattr(key, "cache_salt", None)
|
||||||
|
)
|
||||||
|
return get_hash_str(key, prior_hash, page_size=page_size)
|
||||||
|
|
||||||
|
|
||||||
def hash_str_to_int64(hash_str: str) -> int:
|
def hash_str_to_int64(hash_str: str) -> int:
|
||||||
"""Convert SHA256 hex string to signed 64-bit integer for events.
|
"""Convert SHA256 hex string to signed 64-bit integer for events.
|
||||||
|
|
||||||
@@ -138,15 +168,15 @@ def compute_node_hash_values(node: Any, page_size: int) -> List[str]:
|
|||||||
if len(node.parent.key) > 0 and len(node.parent.hash_value) > 0:
|
if len(node.parent.key) > 0 and len(node.parent.hash_value) > 0:
|
||||||
parent_hash = node.parent.hash_value[-1]
|
parent_hash = node.parent.hash_value[-1]
|
||||||
|
|
||||||
hash_values = get_hash_str(node.key, parent_hash, page_size=page_size)
|
hash_values = get_storage_hash_str(node.key, parent_hash, page_size=page_size)
|
||||||
assert isinstance(hash_values, list)
|
assert isinstance(hash_values, list)
|
||||||
return hash_values
|
return hash_values
|
||||||
|
|
||||||
|
|
||||||
def compute_node_event_hash_values(node: Any, page_size: int) -> List[str]:
|
def compute_node_event_hash_values(node: Any, page_size: int) -> List[str]:
|
||||||
"""Compute and memoize namespace-aware external KV-event hashes."""
|
"""Hash tokens with the legacy salt seed; omit extra_key."""
|
||||||
cache_salt = node.key.cache_salt
|
namespace = (node.key.extra_key, node.key.cache_salt)
|
||||||
if cache_salt is None:
|
if namespace == (None, None):
|
||||||
return compute_node_hash_values(node, page_size)
|
return compute_node_hash_values(node, page_size)
|
||||||
|
|
||||||
if node.event_hash_value is not None:
|
if node.event_hash_value is not None:
|
||||||
@@ -154,31 +184,22 @@ def compute_node_event_hash_values(node: Any, page_size: int) -> List[str]:
|
|||||||
|
|
||||||
missing_nodes = []
|
missing_nodes = []
|
||||||
current = node
|
current = node
|
||||||
while (
|
while current is not None and current.key is not None and len(current.key) > 0:
|
||||||
current is not None
|
if (current.key.extra_key, current.key.cache_salt) != namespace:
|
||||||
and current.key is not None
|
raise ValueError("Radix path contains mismatched cache namespaces")
|
||||||
and len(current.key) > 0
|
if current.event_hash_value is not None:
|
||||||
and current.event_hash_value is None
|
break
|
||||||
):
|
|
||||||
if current.key.cache_salt != cache_salt:
|
|
||||||
raise ValueError("Radix path contains mismatched cache_salt values")
|
|
||||||
missing_nodes.append(current)
|
missing_nodes.append(current)
|
||||||
current = current.parent
|
current = current.parent
|
||||||
|
|
||||||
if (
|
|
||||||
current is not None
|
|
||||||
and current.key is not None
|
|
||||||
and len(current.key) > 0
|
|
||||||
and current.key.cache_salt != cache_salt
|
|
||||||
):
|
|
||||||
raise ValueError("Radix path contains mismatched cache_salt values")
|
|
||||||
|
|
||||||
if current is not None and current.event_hash_value:
|
if current is not None and current.event_hash_value:
|
||||||
parent_hash = current.event_hash_value[-1]
|
parent_hash = current.event_hash_value[-1]
|
||||||
else:
|
elif node.key.cache_salt is not None:
|
||||||
parent_hash = hashlib.sha256(
|
parent_hash = hashlib.sha256(
|
||||||
b"sglang-cache-salt-v1\0" + cache_salt.encode("utf-8")
|
b"sglang-cache-salt-v1\0" + node.key.cache_salt.encode("utf-8")
|
||||||
).hexdigest()
|
).hexdigest()
|
||||||
|
else:
|
||||||
|
parent_hash = None
|
||||||
|
|
||||||
for missing_node in reversed(missing_nodes):
|
for missing_node in reversed(missing_nodes):
|
||||||
hash_values = get_hash_str(missing_node.key, parent_hash, page_size=page_size)
|
hash_values = get_hash_str(missing_node.key, parent_hash, page_size=page_size)
|
||||||
|
|||||||
@@ -1137,7 +1137,7 @@ impl<K: ChildKeyType> NodeArena<K> {
|
|||||||
.filter_map(|(idx, slot)| slot.as_ref().map(|_| NodeIdx_(idx)))
|
.filter_map(|(idx, slot)| slot.as_ref().map(|_| NodeIdx_(idx)))
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Per-page hash values for a node's key, chained from its parent's last hash.
|
/// Chain page hashes from the parent, or seed a new chain with the namespace.
|
||||||
pub fn compute_node_hash_values(&self, node_id: NodeIdx_, page_size: usize) -> Vec<String> {
|
pub fn compute_node_hash_values(&self, node_id: NodeIdx_, page_size: usize) -> Vec<String> {
|
||||||
let node = self.node(node_id);
|
let node = self.node(node_id);
|
||||||
let parent_hash = node.parent.and_then(|parent_id| {
|
let parent_hash = node.parent.and_then(|parent_id| {
|
||||||
@@ -1148,7 +1148,27 @@ impl<K: ChildKeyType> NodeArena<K> {
|
|||||||
None
|
None
|
||||||
}
|
}
|
||||||
});
|
});
|
||||||
crate::node::get_hash_str::<K>(node.key.as_ref(), parent_hash, page_size)
|
let prior = parent_hash.map(str::to_owned).or_else(|| {
|
||||||
|
let namespace = node.namespace.as_ref();
|
||||||
|
if namespace == KeyNamespaceRef::default() {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
// Match Python's storage_namespace_seed byte for byte.
|
||||||
|
let mut hasher = Sha256::new();
|
||||||
|
hasher.update(b"sglang-cache-namespace-v1");
|
||||||
|
for part in [namespace.extra_key, namespace.cache_salt] {
|
||||||
|
match part {
|
||||||
|
None => hasher.update([0u8]),
|
||||||
|
Some(part) => {
|
||||||
|
hasher.update([1u8]);
|
||||||
|
hasher.update((part.len() as u64).to_le_bytes());
|
||||||
|
hasher.update(part.as_bytes());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
Some(digest_to_hex(&hasher.finalize().into()))
|
||||||
|
});
|
||||||
|
crate::node::get_hash_str::<K>(node.key.as_ref(), prior.as_deref(), page_size)
|
||||||
}
|
}
|
||||||
|
|
||||||
/// The ancestor chain's hash values ending at `node_id`, in root-to-node
|
/// The ancestor chain's hash values ending at `node_id`, in root-to-node
|
||||||
|
|||||||
@@ -1945,3 +1945,24 @@ fn iter_yields_all_members() {
|
|||||||
members.sort_unstable();
|
members.sort_unstable();
|
||||||
assert_eq!(members, vec![NodeIdx_(10), NodeIdx_(30)]);
|
assert_eq!(members, vec![NodeIdx_(10), NodeIdx_(30)]);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Pin Python/Rust storage hashes across a parent-child boundary.
|
||||||
|
#[test]
|
||||||
|
fn storage_hashes_match_python() -> Result<(), TreeCoreRuntimeError> {
|
||||||
|
let namespace = KeyNamespaceRef::new(Some("adapter-a"), Some("tenant-a"));
|
||||||
|
let mut arena: NodeArena<Vec<i64>> = NodeArena::new(vec![FULL], 2);
|
||||||
|
let root = arena.root();
|
||||||
|
let parent = arena.alloc_child_in_namespace(root, vec![1, 2], 0, namespace)?;
|
||||||
|
let hashes = arena.compute_node_hash_values(parent, 2);
|
||||||
|
assert_eq!(
|
||||||
|
hashes,
|
||||||
|
vec!["91b8b854063250a84c6f75b3d294bc5d72047c3a15c52b09038a5831f69ecd1a"]
|
||||||
|
);
|
||||||
|
arena.node_mut(parent).hash_value = Some(hashes);
|
||||||
|
let child = arena.alloc_child_in_namespace(parent, vec![3, 4], 0, namespace)?;
|
||||||
|
assert_eq!(
|
||||||
|
arena.compute_node_hash_values(child, 2),
|
||||||
|
vec!["c1ab67afa32b9fdd2ac8429d44d56f207a99c03a30b7a9bb131a32db35354c90"]
|
||||||
|
);
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|||||||
@@ -2754,11 +2754,11 @@ fn insert_coalesces_parent_linked_block_stores() {
|
|||||||
.hash_value,
|
.hash_value,
|
||||||
Some(hashes)
|
Some(hashes)
|
||||||
);
|
);
|
||||||
assert!(tc.salted_event_hashes.is_empty());
|
assert!(tc.namespaced_event_hashes.is_empty());
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn salted_event_hashes_are_sparse_and_removed_with_the_node() {
|
fn namespaced_event_hashes_are_sparse_and_removed_with_the_node() {
|
||||||
let mut tc = events_core(2);
|
let mut tc = events_core(2);
|
||||||
let key = vec![1, 2, 7, 8];
|
let key = vec![1, 2, 7, 8];
|
||||||
tc.insert(&insert_params_in_namespace(
|
tc.insert(&insert_params_in_namespace(
|
||||||
@@ -2773,8 +2773,8 @@ fn salted_event_hashes_are_sparse_and_removed_with_the_node() {
|
|||||||
.match_prefix(&match_params_in_namespace(&key, None, Some("tenant-a")))
|
.match_prefix(&match_params_in_namespace(&key, None, Some("tenant-a")))
|
||||||
.best_match_node_id;
|
.best_match_node_id;
|
||||||
let leaf_idx = tc.arena.resolve(leaf).expect("live test node");
|
let leaf_idx = tc.arena.resolve(leaf).expect("live test node");
|
||||||
assert_eq!(tc.salted_event_hashes[&leaf].len(), 2);
|
assert_eq!(tc.namespaced_event_hashes[&leaf].len(), 2);
|
||||||
assert_eq!(
|
assert_ne!(
|
||||||
tc.arena.node(leaf_idx).hash_value,
|
tc.arena.node(leaf_idx).hash_value,
|
||||||
Some(crate::node::get_hash_str::<Vec<i64>>(&key, None, 2))
|
Some(crate::node::get_hash_str::<Vec<i64>>(&key, None, 2))
|
||||||
);
|
);
|
||||||
@@ -2790,7 +2790,7 @@ fn salted_event_hashes_are_sparse_and_removed_with_the_node() {
|
|||||||
accumulate_step(step, &mut tracker, &mut device_frees, &mut host_frees);
|
accumulate_step(step, &mut tracker, &mut device_frees, &mut host_frees);
|
||||||
tc.evict_device_end(FULL);
|
tc.evict_device_end(FULL);
|
||||||
tc.take_events();
|
tc.take_events();
|
||||||
assert!(tc.salted_event_hashes.is_empty());
|
assert!(tc.namespaced_event_hashes.is_empty());
|
||||||
|
|
||||||
tc.insert(&insert_params_in_namespace(
|
tc.insert(&insert_params_in_namespace(
|
||||||
&key,
|
&key,
|
||||||
@@ -2798,13 +2798,48 @@ fn salted_event_hashes_are_sparse_and_removed_with_the_node() {
|
|||||||
None,
|
None,
|
||||||
Some("tenant-a"),
|
Some("tenant-a"),
|
||||||
));
|
));
|
||||||
assert!(!tc.salted_event_hashes.is_empty());
|
assert!(!tc.namespaced_event_hashes.is_empty());
|
||||||
tc.reset();
|
tc.reset();
|
||||||
assert!(tc.salted_event_hashes.is_empty());
|
assert!(tc.namespaced_event_hashes.is_empty());
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn salted_event_hashes_survive_node_split() {
|
fn extra_key_nodes_publish_token_only_event_hashes() {
|
||||||
|
// Events omit extra_key; storage includes it.
|
||||||
|
let mut tc = events_core(2);
|
||||||
|
let key = vec![1, 2, 7, 8];
|
||||||
|
tc.insert(&insert_params_in_namespace(
|
||||||
|
&key,
|
||||||
|
&[10, 11, 12, 13],
|
||||||
|
Some("lora-a"),
|
||||||
|
None,
|
||||||
|
));
|
||||||
|
let token_only = crate::node::get_hash_str::<Vec<i64>>(&key, None, 2);
|
||||||
|
assert_eq!(
|
||||||
|
tc.take_events(),
|
||||||
|
vec![KvCacheEvent::BlockStored {
|
||||||
|
block_hashes: token_only
|
||||||
|
.iter()
|
||||||
|
.map(|hash| crate::node::hash_str_to_int64(hash))
|
||||||
|
.collect(),
|
||||||
|
parent_block_hash: None,
|
||||||
|
token_ids: key.clone(),
|
||||||
|
block_size: 2,
|
||||||
|
medium: StorageMedium::Gpu,
|
||||||
|
cache_salt: None,
|
||||||
|
}]
|
||||||
|
);
|
||||||
|
|
||||||
|
let leaf = tc
|
||||||
|
.match_prefix(&match_params_in_namespace(&key, Some("lora-a"), None))
|
||||||
|
.best_match_node_id;
|
||||||
|
let leaf_idx = tc.arena.resolve(leaf).expect("live test node");
|
||||||
|
assert_eq!(tc.namespaced_event_hashes[&leaf].len(), 2);
|
||||||
|
assert_ne!(tc.arena.node(leaf_idx).hash_value, Some(token_only));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn namespaced_event_hashes_survive_node_split() {
|
||||||
let mut tc = events_core(2);
|
let mut tc = events_core(2);
|
||||||
let original = vec![1, 2, 3, 4];
|
let original = vec![1, 2, 3, 4];
|
||||||
tc.insert(&insert_params_in_namespace(
|
tc.insert(&insert_params_in_namespace(
|
||||||
@@ -2820,7 +2855,7 @@ fn salted_event_hashes_survive_node_split() {
|
|||||||
Some("tenant-a"),
|
Some("tenant-a"),
|
||||||
))
|
))
|
||||||
.best_match_node_id;
|
.best_match_node_id;
|
||||||
let original_hashes = tc.salted_event_hashes[&original_leaf].clone();
|
let original_hashes = tc.namespaced_event_hashes[&original_leaf].clone();
|
||||||
tc.take_events();
|
tc.take_events();
|
||||||
|
|
||||||
let branch = vec![1, 2, 5, 6];
|
let branch = vec![1, 2, 5, 6];
|
||||||
@@ -2844,8 +2879,14 @@ fn salted_event_hashes_survive_node_split() {
|
|||||||
.node(tc.arena.resolve(split_child).expect("live test node"))
|
.node(tc.arena.resolve(split_child).expect("live test node"))
|
||||||
.parent();
|
.parent();
|
||||||
let split_parent = tc.arena.node(split_parent_idx).id;
|
let split_parent = tc.arena.node(split_parent_idx).id;
|
||||||
assert_eq!(tc.salted_event_hashes[&split_parent], original_hashes[..1]);
|
assert_eq!(
|
||||||
assert_eq!(tc.salted_event_hashes[&split_child], original_hashes[1..]);
|
tc.namespaced_event_hashes[&split_parent],
|
||||||
|
original_hashes[..1]
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
tc.namespaced_event_hashes[&split_child],
|
||||||
|
original_hashes[1..]
|
||||||
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -2863,9 +2904,9 @@ fn salted_event_hash_walk_is_iterative_and_on_demand() {
|
|||||||
)
|
)
|
||||||
.unwrap();
|
.unwrap();
|
||||||
}
|
}
|
||||||
assert!(tc.salted_event_hashes.is_empty());
|
assert!(tc.namespaced_event_hashes.is_empty());
|
||||||
tc.ensure_salted_event_hashes_(parent);
|
tc.ensure_namespaced_event_hashes_(parent);
|
||||||
assert_eq!(tc.salted_event_hashes.len(), 1100);
|
assert_eq!(tc.namespaced_event_hashes.len(), 1100);
|
||||||
}
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
|
|||||||
@@ -542,9 +542,8 @@ pub struct UnifiedTreeCore<K: ChildKeyType> {
|
|||||||
pub(crate) enable_kv_cache_events: bool,
|
pub(crate) enable_kv_cache_events: bool,
|
||||||
/// Queued placement events, drained by take_events.
|
/// Queued placement events, drained by take_events.
|
||||||
pub(crate) kv_event_queue: Vec<KvCacheEvent<K::Atom>>,
|
pub(crate) kv_event_queue: Vec<KvCacheEvent<K::Atom>>,
|
||||||
/// Namespace-aware event hashes, populated only for salted nodes whose
|
/// Namespaced event hashes, seeded only by cache_salt; events omit extra_key.
|
||||||
/// placement events are requested. Storage hashes remain on the nodes.
|
pub(crate) namespaced_event_hashes: HashMap<NodeId, Vec<HashDigest>>,
|
||||||
pub(crate) salted_event_hashes: HashMap<NodeId, Vec<HashDigest>>,
|
|
||||||
/// Hit count at which a node earns a host write-through backup.
|
/// Hit count at which a node earns a host write-through backup.
|
||||||
pub(crate) write_through_threshold: i64,
|
pub(crate) write_through_threshold: i64,
|
||||||
|
|
||||||
@@ -723,7 +722,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
has_swa_host_pool: params.has_swa_host_pool,
|
has_swa_host_pool: params.has_swa_host_pool,
|
||||||
enable_kv_cache_events: params.enable_kv_cache_events,
|
enable_kv_cache_events: params.enable_kv_cache_events,
|
||||||
kv_event_queue: Vec::new(),
|
kv_event_queue: Vec::new(),
|
||||||
salted_event_hashes: HashMap::new(),
|
namespaced_event_hashes: HashMap::new(),
|
||||||
write_through_threshold: params.write_through_threshold,
|
write_through_threshold: params.write_through_threshold,
|
||||||
swa_uuid_counter: 1,
|
swa_uuid_counter: 1,
|
||||||
device: params.device,
|
device: params.device,
|
||||||
@@ -752,7 +751,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
self.write_back_coexist_reclaim_digest = 0;
|
self.write_back_coexist_reclaim_digest = 0;
|
||||||
self.lru_lists = Self::new_lru_lists();
|
self.lru_lists = Self::new_lru_lists();
|
||||||
self.full_evict_device_heap.clear();
|
self.full_evict_device_heap.clear();
|
||||||
self.salted_event_hashes.clear();
|
self.namespaced_event_hashes.clear();
|
||||||
self.ongoing_insert_walk_state = None;
|
self.ongoing_insert_walk_state = None;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1798,13 +1797,13 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
child.hash_value = child_hash;
|
child.hash_value = child_hash;
|
||||||
self.arena.node_mut(new_node_id).hash_value = new_node_hash;
|
self.arena.node_mut(new_node_id).hash_value = new_node_hash;
|
||||||
let child_handle = self.arena.node(child_id).id;
|
let child_handle = self.arena.node(child_id).id;
|
||||||
if let Some(mut parent_event_hashes) = self.salted_event_hashes.remove(&child_handle) {
|
if let Some(mut parent_event_hashes) = self.namespaced_event_hashes.remove(&child_handle) {
|
||||||
let child_event_hashes = parent_event_hashes.split_off(split_len / self.page_size);
|
let child_event_hashes = parent_event_hashes.split_off(split_len / self.page_size);
|
||||||
parent_event_hashes.shrink_to_fit();
|
parent_event_hashes.shrink_to_fit();
|
||||||
let new_node_handle = self.arena.node(new_node_id).id;
|
let new_node_handle = self.arena.node(new_node_id).id;
|
||||||
self.salted_event_hashes
|
self.namespaced_event_hashes
|
||||||
.insert(new_node_handle, parent_event_hashes);
|
.insert(new_node_handle, parent_event_hashes);
|
||||||
self.salted_event_hashes
|
self.namespaced_event_hashes
|
||||||
.insert(child_handle, child_event_hashes);
|
.insert(child_handle, child_event_hashes);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -2530,7 +2529,7 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
pub fn remove_leaf_from_parent_(&mut self, node_id: NodeIdx_) {
|
pub fn remove_leaf_from_parent_(&mut self, node_id: NodeIdx_) {
|
||||||
// Arena slots are reused, so discard tracking before freeing the node.
|
// Arena slots are reused, so discard tracking before freeing the node.
|
||||||
self.full_coexisting_host_nodes.discard(node_id);
|
self.full_coexisting_host_nodes.discard(node_id);
|
||||||
self.salted_event_hashes
|
self.namespaced_event_hashes
|
||||||
.remove(&self.arena.node(node_id).id);
|
.remove(&self.arena.node(node_id).id);
|
||||||
// The arena is the registry: freeing detaches by page key and recycles the slot.
|
// The arena is the registry: freeing detaches by page key and recycles the slot.
|
||||||
self.arena
|
self.arena
|
||||||
@@ -2767,18 +2766,13 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Fill the salted external event-hash chain through `node_id`.
|
/// Fill the event chain through `node_id`, using only cache_salt.
|
||||||
fn ensure_salted_event_hashes_(&mut self, node_id: NodeIdx_) {
|
fn ensure_namespaced_event_hashes_(&mut self, node_id: NodeIdx_) {
|
||||||
let cache_salt = self
|
let node = self.arena.node(node_id);
|
||||||
.arena
|
if self.namespaced_event_hashes.contains_key(&node.id) {
|
||||||
.node(node_id)
|
|
||||||
.namespace
|
|
||||||
.cache_salt_arc()
|
|
||||||
.expect("salted event hashing requires cache_salt");
|
|
||||||
let node_handle = self.arena.node(node_id).id;
|
|
||||||
if self.salted_event_hashes.contains_key(&node_handle) {
|
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
let namespace = node.namespace.clone();
|
||||||
|
|
||||||
let mut missing = Vec::new();
|
let mut missing = Vec::new();
|
||||||
let mut cursor = Some(node_id);
|
let mut cursor = Some(node_id);
|
||||||
@@ -2789,11 +2783,10 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
node.namespace.cache_salt(),
|
node.namespace, namespace,
|
||||||
Some(cache_salt.as_ref()),
|
"radix path contains mismatched cache namespaces"
|
||||||
"radix path contains mismatched cache_salt values"
|
|
||||||
);
|
);
|
||||||
if let Some(hashes) = self.salted_event_hashes.get(&node.id) {
|
if let Some(hashes) = self.namespaced_event_hashes.get(&node.id) {
|
||||||
prior = hashes.last().copied();
|
prior = hashes.last().copied();
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
@@ -2801,12 +2794,14 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
cursor = node.try_parent();
|
cursor = node.try_parent();
|
||||||
}
|
}
|
||||||
|
|
||||||
let mut prior = prior.unwrap_or_else(|| {
|
if prior.is_none() {
|
||||||
let mut hasher = Sha256::new();
|
prior = namespace.cache_salt().map(|cache_salt| {
|
||||||
hasher.update(b"sglang-cache-salt-v1\0");
|
let mut hasher = Sha256::new();
|
||||||
hasher.update(cache_salt.as_bytes());
|
hasher.update(b"sglang-cache-salt-v1\0");
|
||||||
hasher.finalize().into()
|
hasher.update(cache_salt.as_bytes());
|
||||||
});
|
hasher.finalize().into()
|
||||||
|
});
|
||||||
|
}
|
||||||
for id in missing.into_iter().rev() {
|
for id in missing.into_iter().rev() {
|
||||||
let (handle, hashes) = {
|
let (handle, hashes) = {
|
||||||
let node = self.arena.node(id);
|
let node = self.arena.node(id);
|
||||||
@@ -2814,15 +2809,15 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
node.id,
|
node.id,
|
||||||
crate::node::get_hash_digests::<K>(
|
crate::node::get_hash_digests::<K>(
|
||||||
node.key.as_ref(),
|
node.key.as_ref(),
|
||||||
Some(&prior),
|
prior.as_ref(),
|
||||||
self.page_size,
|
self.page_size,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
};
|
};
|
||||||
if let Some(last) = hashes.last() {
|
if let Some(last) = hashes.last() {
|
||||||
prior = *last;
|
prior = Some(*last);
|
||||||
}
|
}
|
||||||
self.salted_event_hashes.insert(handle, hashes);
|
self.namespaced_event_hashes.insert(handle, hashes);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -2836,15 +2831,16 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
self.arena.node_mut(node_id).hash_value = Some(hash_values);
|
self.arena.node_mut(node_id).hash_value = Some(hash_values);
|
||||||
}
|
}
|
||||||
let cache_salt = self.arena.node(node_id).namespace.cache_salt_arc();
|
let cache_salt = self.arena.node(node_id).namespace.cache_salt_arc();
|
||||||
if cache_salt.is_some() {
|
let namespaced = self.arena.node(node_id).namespace != KeyNamespace::default();
|
||||||
self.ensure_salted_event_hashes_(node_id);
|
if namespaced {
|
||||||
|
self.ensure_namespaced_event_hashes_(node_id);
|
||||||
}
|
}
|
||||||
let events = {
|
let events = {
|
||||||
let node = self.arena.node(node_id);
|
let node = self.arena.node(node_id);
|
||||||
let mut parent_block_hash = node.parent.and_then(|parent_id| {
|
let mut parent_block_hash = node.parent.and_then(|parent_id| {
|
||||||
let parent = self.arena.node(parent_id);
|
let parent = self.arena.node(parent_id);
|
||||||
if cache_salt.is_some() {
|
if namespaced {
|
||||||
self.salted_event_hashes
|
self.namespaced_event_hashes
|
||||||
.get(&parent.id)
|
.get(&parent.id)
|
||||||
.and_then(|hashes| hashes.last())
|
.and_then(|hashes| hashes.last())
|
||||||
.map(crate::node::hash_digest_to_int64)
|
.map(crate::node::hash_digest_to_int64)
|
||||||
@@ -2867,8 +2863,8 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
});
|
});
|
||||||
parent_block_hash = Some(block_hash);
|
parent_block_hash = Some(block_hash);
|
||||||
};
|
};
|
||||||
if cache_salt.is_some() {
|
if namespaced {
|
||||||
let hashes = &self.salted_event_hashes[&node.id];
|
let hashes = &self.namespaced_event_hashes[&node.id];
|
||||||
assert!(
|
assert!(
|
||||||
hashes.len() >= num_pages,
|
hashes.len() >= num_pages,
|
||||||
"store event: {} page hashes for {num_pages} pages",
|
"store event: {} page hashes for {num_pages} pages",
|
||||||
@@ -2904,14 +2900,14 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
let hash_values = self.arena.compute_node_hash_values(node_id, self.page_size);
|
let hash_values = self.arena.compute_node_hash_values(node_id, self.page_size);
|
||||||
self.arena.node_mut(node_id).hash_value = Some(hash_values);
|
self.arena.node_mut(node_id).hash_value = Some(hash_values);
|
||||||
}
|
}
|
||||||
let cache_salt = self.arena.node(node_id).namespace.cache_salt_arc();
|
let namespaced = self.arena.node(node_id).namespace != KeyNamespace::default();
|
||||||
if cache_salt.is_some() {
|
if namespaced {
|
||||||
self.ensure_salted_event_hashes_(node_id);
|
self.ensure_namespaced_event_hashes_(node_id);
|
||||||
}
|
}
|
||||||
let node = self.arena.node(node_id);
|
let node = self.arena.node(node_id);
|
||||||
let num_pages = node.key.atom_len().div_ceil(self.page_size);
|
let num_pages = node.key.atom_len().div_ceil(self.page_size);
|
||||||
let block_hashes: Vec<i64> = if cache_salt.is_some() {
|
let block_hashes: Vec<i64> = if namespaced {
|
||||||
self.salted_event_hashes[&node.id][..num_pages]
|
self.namespaced_event_hashes[&node.id][..num_pages]
|
||||||
.iter()
|
.iter()
|
||||||
.map(crate::node::hash_digest_to_int64)
|
.map(crate::node::hash_digest_to_int64)
|
||||||
.collect()
|
.collect()
|
||||||
@@ -3955,23 +3951,23 @@ impl<K: ChildKeyType> UnifiedTreeCore<K> {
|
|||||||
&orphans[..orphans.len().min(5)]
|
&orphans[..orphans.len().min(5)]
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
for (&node_handle, hashes) in &self.salted_event_hashes {
|
for (&node_handle, hashes) in &self.namespaced_event_hashes {
|
||||||
let Ok(node_id) = self.arena.resolve(node_handle) else {
|
let Ok(node_id) = self.arena.resolve(node_handle) else {
|
||||||
errors.push(format!(
|
errors.push(format!(
|
||||||
"[Events] salted hashes reference freed node {node_handle}"
|
"[Events] event hashes reference freed node {node_handle}"
|
||||||
));
|
));
|
||||||
continue;
|
continue;
|
||||||
};
|
};
|
||||||
let node = self.arena.node(node_id);
|
let node = self.arena.node(node_id);
|
||||||
if node.namespace.cache_salt().is_none() {
|
if node.namespace == KeyNamespace::default() {
|
||||||
errors.push(format!(
|
errors.push(format!(
|
||||||
"[Events] unsalted node {node_handle} carries salted hashes"
|
"[Events] default-namespace node {node_handle} carries event hashes"
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
let expected_pages = node.key.atom_len().div_ceil(self.page_size);
|
let expected_pages = node.key.atom_len().div_ceil(self.page_size);
|
||||||
if hashes.len() != expected_pages {
|
if hashes.len() != expected_pages {
|
||||||
errors.push(format!(
|
errors.push(format!(
|
||||||
"[Events] node {node_handle} has {} salted hashes for {expected_pages} pages",
|
"[Events] node {node_handle} has {} event hashes for {expected_pages} pages",
|
||||||
hashes.len()
|
hashes.len()
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,156 @@
|
|||||||
|
"""Storage round trips preserve LoRA and salt isolation."""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import random
|
||||||
|
import shutil
|
||||||
|
import tempfile
|
||||||
|
import unittest
|
||||||
|
from typing import Dict, Optional
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
from sglang.benchmark.utils import get_tokenizer
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_MODEL_NAME_FOR_TEST,
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
DEFAULT_URL_FOR_TEST,
|
||||||
|
CustomTestCase,
|
||||||
|
popen_launch_server,
|
||||||
|
terminate_and_kill_process_tree,
|
||||||
|
)
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=300, stage="base-b", runner_config="1-gpu-large")
|
||||||
|
|
||||||
|
LORA_NAME = "sql"
|
||||||
|
LORA_PATH = "philschmid/code-llama-3-1-8b-text-to-sql-lora"
|
||||||
|
PAGE_SIZE = 64
|
||||||
|
PROMPT_TOKENS = 768
|
||||||
|
|
||||||
|
|
||||||
|
class TestHiCacheStorageLoRAIsolation(CustomTestCase):
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.temp_dir = tempfile.mkdtemp()
|
||||||
|
cls.model = DEFAULT_MODEL_NAME_FOR_TEST
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
cls.tokenizer = get_tokenizer(cls.model)
|
||||||
|
|
||||||
|
extra_config = {"hicache_storage_pass_prefix_keys": True}
|
||||||
|
other_args = [
|
||||||
|
"--enable-hierarchical-cache",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.6",
|
||||||
|
"--hicache-ratio",
|
||||||
|
"1.2",
|
||||||
|
"--page-size",
|
||||||
|
str(PAGE_SIZE),
|
||||||
|
"--enable-cache-report",
|
||||||
|
"--hicache-storage-prefetch-policy",
|
||||||
|
"wait_complete",
|
||||||
|
"--hicache-storage-backend",
|
||||||
|
"file",
|
||||||
|
"--hicache-storage-backend-extra-config",
|
||||||
|
json.dumps(extra_config),
|
||||||
|
"--enable-lora",
|
||||||
|
"--lora-paths",
|
||||||
|
f"{LORA_NAME}={LORA_PATH}",
|
||||||
|
"--max-loras-per-batch",
|
||||||
|
"2",
|
||||||
|
# Triton keeps radix caching enabled under deterministic inference.
|
||||||
|
"--enable-deterministic-inference",
|
||||||
|
"--attention-backend",
|
||||||
|
"triton",
|
||||||
|
]
|
||||||
|
env = {**os.environ, "SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR": cls.temp_dir}
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=other_args,
|
||||||
|
env=env,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
if getattr(cls, "process", None):
|
||||||
|
terminate_and_kill_process_tree(cls.process)
|
||||||
|
shutil.rmtree(cls.temp_dir, ignore_errors=True)
|
||||||
|
|
||||||
|
def send_request(
|
||||||
|
self,
|
||||||
|
prompt: str,
|
||||||
|
lora_path: Optional[str],
|
||||||
|
max_tokens: int = 32,
|
||||||
|
cache_salt: Optional[str] = None,
|
||||||
|
) -> Dict:
|
||||||
|
payload = {
|
||||||
|
"text": prompt,
|
||||||
|
"sampling_params": {
|
||||||
|
"temperature": 0.0,
|
||||||
|
"max_new_tokens": max_tokens,
|
||||||
|
"ignore_eos": True,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if lora_path is not None:
|
||||||
|
payload["lora_path"] = lora_path
|
||||||
|
if cache_salt is not None:
|
||||||
|
payload["cache_salt"] = cache_salt
|
||||||
|
response = requests.post(f"{self.base_url}/generate", json=payload, timeout=120)
|
||||||
|
self.assertEqual(response.status_code, 200, response.text)
|
||||||
|
return response.json()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def cached_tokens(response_json: Dict) -> int:
|
||||||
|
return int(response_json.get("meta_info", {}).get("cached_tokens", 0))
|
||||||
|
|
||||||
|
def flush_device_cache(self):
|
||||||
|
# A short unrelated request first so the pages of interest get offloaded.
|
||||||
|
self.send_request(self.gen_prompt(1), lora_path=None, max_tokens=150)
|
||||||
|
res = requests.post(
|
||||||
|
f"{self.base_url}/flush_cache", params={"timeout": 30}, timeout=40
|
||||||
|
)
|
||||||
|
res.raise_for_status()
|
||||||
|
|
||||||
|
def gen_prompt(self, token_num: int) -> str:
|
||||||
|
vocab = list(self.tokenizer.get_vocab().values())
|
||||||
|
return self.tokenizer.decode(random.choices(vocab, k=token_num))
|
||||||
|
|
||||||
|
def test_adapter_pages_are_isolated_in_storage(self):
|
||||||
|
prompt = self.gen_prompt(PROMPT_TOKENS)
|
||||||
|
hit_floor = PROMPT_TOKENS - 2 * PAGE_SIZE
|
||||||
|
|
||||||
|
# Cold pass with the adapter populates host and storage.
|
||||||
|
lora_first = self.send_request(prompt, lora_path=LORA_NAME)
|
||||||
|
self.flush_device_cache()
|
||||||
|
|
||||||
|
# Read adapter pages before any base request stores the same prompt.
|
||||||
|
lora_again = self.send_request(prompt, lora_path=LORA_NAME)
|
||||||
|
self.assertGreater(
|
||||||
|
self.cached_tokens(lora_again),
|
||||||
|
hit_floor,
|
||||||
|
"the adapter's pages were not served from storage after the flush",
|
||||||
|
)
|
||||||
|
self.assertEqual(lora_first["text"], lora_again["text"])
|
||||||
|
self.flush_device_cache()
|
||||||
|
|
||||||
|
# Base and salted pages must miss existing namespaces, then round-trip.
|
||||||
|
for cache_salt in (None, "tenant-a"):
|
||||||
|
with self.subTest(cache_salt=cache_salt):
|
||||||
|
first = self.send_request(prompt, lora_path=None, cache_salt=cache_salt)
|
||||||
|
self.assertLess(self.cached_tokens(first), PAGE_SIZE)
|
||||||
|
self.flush_device_cache()
|
||||||
|
|
||||||
|
again = self.send_request(prompt, lora_path=None, cache_salt=cache_salt)
|
||||||
|
self.assertGreater(self.cached_tokens(again), hit_floor)
|
||||||
|
self.assertEqual(first["text"], again["text"])
|
||||||
|
self.flush_device_cache()
|
||||||
|
|
||||||
|
lora_third = self.send_request(prompt, lora_path=LORA_NAME)
|
||||||
|
self.assertGreater(self.cached_tokens(lora_third), hit_floor)
|
||||||
|
self.assertEqual(lora_first["text"], lora_third["text"])
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -41,7 +41,7 @@ class TestDecodeHiCacheTreeCore(CustomTestCase):
|
|||||||
rid="req-0",
|
rid="req-0",
|
||||||
origin_input_ids=[0, 1, 2, 3, 4, 5, 6, 7],
|
origin_input_ids=[0, 1, 2, 3, 4, 5, 6, 7],
|
||||||
extra_key="model",
|
extra_key="model",
|
||||||
cache_salt=None,
|
cache_salt="tenant-a",
|
||||||
)
|
)
|
||||||
result = SimpleNamespace(
|
result = SimpleNamespace(
|
||||||
device_indices=torch.tensor([10, 11]),
|
device_indices=torch.tensor([10, 11]),
|
||||||
@@ -56,7 +56,12 @@ class TestDecodeHiCacheTreeCore(CustomTestCase):
|
|||||||
|
|
||||||
self.assertEqual(prefix_match.l3_storage_hit_length, 2)
|
self.assertEqual(prefix_match.l3_storage_hit_length, 2)
|
||||||
tree_cache.query_storage_hit_length.assert_called_once_with(
|
tree_cache.query_storage_hit_length.assert_called_once_with(
|
||||||
22, [4, 5, 6, 7], "h2", ["h0", "h1"]
|
22,
|
||||||
|
[4, 5, 6, 7],
|
||||||
|
"h2",
|
||||||
|
["h0", "h1"],
|
||||||
|
extra_key="model",
|
||||||
|
cache_salt="tenant-a",
|
||||||
)
|
)
|
||||||
|
|
||||||
DecodeHiCachePreallocMixin._start_hicache_prefetch(harness, req, prefix_match)
|
DecodeHiCachePreallocMixin._start_hicache_prefetch(harness, req, prefix_match)
|
||||||
@@ -69,7 +74,7 @@ class TestDecodeHiCacheTreeCore(CustomTestCase):
|
|||||||
"h2",
|
"h2",
|
||||||
["h0", "h1"],
|
["h0", "h1"],
|
||||||
extra_key="model",
|
extra_key="model",
|
||||||
cache_salt=None,
|
cache_salt="tenant-a",
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_stale_prefetch_anchor_degrades_to_l2(self):
|
def test_stale_prefetch_anchor_degrades_to_l2(self):
|
||||||
|
|||||||
@@ -26,6 +26,8 @@ from sglang.srt.managers.scheduler_components.batch_result_processor import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||||
|
from sglang.srt.mem_cache.radix_cache import RadixKey
|
||||||
|
from sglang.srt.mem_cache.utils import get_hash_str, get_storage_hash_str
|
||||||
from sglang.srt.runtime_context import get_context
|
from sglang.srt.runtime_context import get_context
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
@@ -44,6 +46,8 @@ def _make_mock_req(
|
|||||||
"""Create a mock Req with the KV cache state needed for testing."""
|
"""Create a mock Req with the KV cache state needed for testing."""
|
||||||
req = MagicMock()
|
req = MagicMock()
|
||||||
req.rid = rid
|
req.rid = rid
|
||||||
|
req.extra_key = None # base traffic: storage hashes chain from tokens alone
|
||||||
|
req.cache_salt = None
|
||||||
req.origin_input_ids = list(range(origin_len))
|
req.origin_input_ids = list(range(origin_len))
|
||||||
req.kv = ReqKvInfo(
|
req.kv = ReqKvInfo(
|
||||||
req_pool_idx=req_pool_idx,
|
req_pool_idx=req_pool_idx,
|
||||||
@@ -121,6 +125,26 @@ class _FinishedEvent:
|
|||||||
class TestReleaseFinishedReq(unittest.TestCase):
|
class TestReleaseFinishedReq(unittest.TestCase):
|
||||||
"""Tests for _release_finished_req overallocation cleanup."""
|
"""Tests for _release_finished_req overallocation cleanup."""
|
||||||
|
|
||||||
|
def test_decode_offload_hash_chain_matches_prefill(self):
|
||||||
|
"""Decode pages must keep the prefill namespace across offload chunks."""
|
||||||
|
manager, _ = _make_manager(pool_size=8, page_size=2)
|
||||||
|
manager.cache_controller = MagicMock(get_hash_str=get_hash_str)
|
||||||
|
tokens = [1, 2, 3, 4, 5, 6]
|
||||||
|
for extra_key, cache_salt in [
|
||||||
|
(None, None),
|
||||||
|
("lora-a", None),
|
||||||
|
(None, "tenant-a"),
|
||||||
|
("lora-a", "tenant-a"),
|
||||||
|
]:
|
||||||
|
with self.subTest(extra_key=extra_key, cache_salt=cache_salt):
|
||||||
|
namespace = dict(extra_key=extra_key, cache_salt=cache_salt)
|
||||||
|
prefix = manager._compute_prefix_hash(tokens[:4], **namespace)
|
||||||
|
tail = manager._compute_prefix_hash(tokens[4:], prefix[-1], **namespace)
|
||||||
|
self.assertEqual(
|
||||||
|
prefix + tail,
|
||||||
|
get_storage_hash_str(RadixKey(tokens, **namespace), page_size=2),
|
||||||
|
)
|
||||||
|
|
||||||
def test_no_overallocation(self):
|
def test_no_overallocation(self):
|
||||||
"""Without spec v2, kv_committed == kv_allocated; no extra free."""
|
"""Without spec v2, kv_committed == kv_allocated; no extra free."""
|
||||||
manager, freed = _make_manager(pool_size=32)
|
manager, freed = _make_manager(pool_size=32)
|
||||||
|
|||||||
@@ -56,10 +56,11 @@ def _legacy_page_hashes(key, page_size, prior_hash=None):
|
|||||||
|
|
||||||
|
|
||||||
class _HashKey:
|
class _HashKey:
|
||||||
def __init__(self, token_ids, is_bigram=False, cache_salt=None):
|
def __init__(self, token_ids, is_bigram=False, cache_salt=None, extra_key=None):
|
||||||
self.token_ids = token_ids
|
self.token_ids = token_ids
|
||||||
self.is_bigram = is_bigram
|
self.is_bigram = is_bigram
|
||||||
self.cache_salt = cache_salt
|
self.cache_salt = cache_salt
|
||||||
|
self.extra_key = extra_key
|
||||||
|
|
||||||
def __len__(self):
|
def __len__(self):
|
||||||
if self.is_bigram:
|
if self.is_bigram:
|
||||||
@@ -75,8 +76,13 @@ class _HashKey:
|
|||||||
self.token_ids[start : stop + 1],
|
self.token_ids[start : stop + 1],
|
||||||
is_bigram=True,
|
is_bigram=True,
|
||||||
cache_salt=self.cache_salt,
|
cache_salt=self.cache_salt,
|
||||||
|
extra_key=self.extra_key,
|
||||||
)
|
)
|
||||||
return _HashKey(self.token_ids[start:stop], cache_salt=self.cache_salt)
|
return _HashKey(
|
||||||
|
self.token_ids[start:stop],
|
||||||
|
cache_salt=self.cache_salt,
|
||||||
|
extra_key=self.extra_key,
|
||||||
|
)
|
||||||
if self.is_bigram:
|
if self.is_bigram:
|
||||||
return (self.token_ids[index], self.token_ids[index + 1])
|
return (self.token_ids[index], self.token_ids[index + 1])
|
||||||
return self.token_ids[index]
|
return self.token_ids[index]
|
||||||
@@ -267,6 +273,52 @@ class TestGetHashStr(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestStorageHashNamespace(unittest.TestCase):
|
||||||
|
def test_node_hashes_isolate_namespaces_and_continue_the_chain(self):
|
||||||
|
root = SimpleNamespace(parent=None, key=_HashKey(array("q")), hash_value=None)
|
||||||
|
tokens = array("q", range(1, 129))
|
||||||
|
|
||||||
|
def child(extra_key=None, cache_salt=None):
|
||||||
|
return SimpleNamespace(
|
||||||
|
parent=root,
|
||||||
|
key=_HashKey(tokens, extra_key=extra_key, cache_salt=cache_salt),
|
||||||
|
hash_value=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
plain = compute_node_hash_values(child(), page_size=64)
|
||||||
|
self.assertEqual(plain, get_hash_str(tokens, None, page_size=64))
|
||||||
|
# Also guard ambiguous concatenations: ("a", "bc") vs ("ab", "c").
|
||||||
|
namespaced = [
|
||||||
|
compute_node_hash_values(child(*namespace), page_size=64)
|
||||||
|
for namespace in [
|
||||||
|
("lora-a", None),
|
||||||
|
("lora-b", None),
|
||||||
|
(None, "tenant-a"),
|
||||||
|
("lora-a", "tenant-a"),
|
||||||
|
("a", "bc"),
|
||||||
|
("ab", "c"),
|
||||||
|
]
|
||||||
|
]
|
||||||
|
for i in range(len(plain)):
|
||||||
|
page_hashes = {plain[i], *(hashes[i] for hashes in namespaced)}
|
||||||
|
self.assertEqual(len(page_hashes), 1 + len(namespaced))
|
||||||
|
|
||||||
|
# Continue the parent chain without re-seeding.
|
||||||
|
parent = child("lora-a", "tenant-a")
|
||||||
|
parent.hash_value = namespaced[3]
|
||||||
|
grand = SimpleNamespace(
|
||||||
|
parent=parent,
|
||||||
|
key=_HashKey(
|
||||||
|
array("q", range(200, 264)), extra_key="lora-a", cache_salt="tenant-a"
|
||||||
|
),
|
||||||
|
hash_value=None,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
compute_node_hash_values(grand, page_size=64),
|
||||||
|
get_hash_str(array("q", range(200, 264)), namespaced[3][-1], page_size=64),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestHashStrToInt64(unittest.TestCase):
|
class TestHashStrToInt64(unittest.TestCase):
|
||||||
def test_zero_hash(self):
|
def test_zero_hash(self):
|
||||||
result = hash_str_to_int64("0" * 64)
|
result = hash_str_to_int64("0" * 64)
|
||||||
@@ -333,10 +385,6 @@ class TestComputeNodeHashValues(unittest.TestCase):
|
|||||||
compute_node_event_hash_values(self._make_node(key), page_size=8),
|
compute_node_event_hash_values(self._make_node(key), page_size=8),
|
||||||
_legacy_page_hashes(key, page_size=8, prior_hash=seed),
|
_legacy_page_hashes(key, page_size=8, prior_hash=seed),
|
||||||
)
|
)
|
||||||
self.assertEqual(
|
|
||||||
compute_node_hash_values(self._make_node(key), page_size=8),
|
|
||||||
_legacy_page_hashes(key, page_size=8),
|
|
||||||
)
|
|
||||||
|
|
||||||
other = _HashKey(array("q", range(1, 17)), cache_salt="tenant-b")
|
other = _HashKey(array("q", range(1, 17)), cache_salt="tenant-b")
|
||||||
self.assertNotEqual(
|
self.assertNotEqual(
|
||||||
|
|||||||
@@ -796,6 +796,32 @@ class TestRadixCache(CustomTestCase):
|
|||||||
]
|
]
|
||||||
self.assertNotEqual(unsalted_hashes, stored[0].block_hashes)
|
self.assertNotEqual(unsalted_hashes, stored[0].block_hashes)
|
||||||
|
|
||||||
|
def test_extra_key_does_not_move_published_block_hashes(self):
|
||||||
|
"""Adding extra_key preserves event hashes and split-parent links."""
|
||||||
|
for cache_salt in (None, "tenant-a"):
|
||||||
|
published = []
|
||||||
|
for extra_key in (None, "lora-a"):
|
||||||
|
cache = RadixCache.create_simulated(
|
||||||
|
page_size=2, enable_kv_cache_events=True
|
||||||
|
)
|
||||||
|
namespace = dict(extra_key=extra_key, cache_salt=cache_salt)
|
||||||
|
for tokens in ([1, 2, 3, 4, 5, 6], [1, 2, 7, 8]):
|
||||||
|
cache.insert(
|
||||||
|
InsertParams(
|
||||||
|
key=RadixKey(array("q", tokens), **namespace),
|
||||||
|
value=torch.tensor(tokens, dtype=torch.int64),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
published.append(
|
||||||
|
[
|
||||||
|
(event.parent_block_hash, tuple(event.block_hashes))
|
||||||
|
for event in cache.take_events()
|
||||||
|
if isinstance(event, BlockStored)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
self.assertEqual(published[0], published[1])
|
||||||
|
self.assertIsNotNone(published[1][-1][0])
|
||||||
|
|
||||||
def test_cache_salt_event_hashes_are_preserved_across_node_split(self):
|
def test_cache_salt_event_hashes_are_preserved_across_node_split(self):
|
||||||
cache = RadixCache.create_simulated(page_size=2, enable_kv_cache_events=True)
|
cache = RadixCache.create_simulated(page_size=2, enable_kv_cache_events=True)
|
||||||
original = RadixKey(array("q", [1, 2, 3, 4]), cache_salt="tenant-a")
|
original = RadixKey(array("q", [1, 2, 3, 4]), cache_salt="tenant-a")
|
||||||
|
|||||||
@@ -52,7 +52,7 @@ from sglang.srt.mem_cache.unified_cache.cache_action import (
|
|||||||
SWARebuild,
|
SWARebuild,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.unified_cache.component_type import ComponentType
|
from sglang.srt.mem_cache.unified_cache.component_type import ComponentType
|
||||||
from sglang.srt.mem_cache.utils import hash_str_to_int64
|
from sglang.srt.mem_cache.utils import get_storage_hash_str, hash_str_to_int64
|
||||||
from sglang.srt.runtime_context import get_context
|
from sglang.srt.runtime_context import get_context
|
||||||
|
|
||||||
|
|
||||||
@@ -1058,21 +1058,25 @@ def test_storage_backup_spec_round_trips_the_backuped_node():
|
|||||||
core = _tree_core(page_size=2)
|
core = _tree_core(page_size=2)
|
||||||
core.set_hicache_enabled()
|
core.set_hicache_enabled()
|
||||||
core.enable_storage = True
|
core.enable_storage = True
|
||||||
_insert(core, [1, 2], [10, 11])
|
key = RadixKey(
|
||||||
_insert(core, [1, 2, 7, 8], [10, 11, 12, 13])
|
array("q", [1, 2, 7, 8]), extra_key="adapter-a", cache_salt="tenant-a"
|
||||||
parent = core.match_prefix(MatchPrefixParams(key=_key([1, 2]))).best_match_node
|
)
|
||||||
child = core.match_prefix(MatchPrefixParams(key=_key([1, 2, 7, 8]))).best_match_node
|
for length in (2, 4):
|
||||||
|
_pump_insert(
|
||||||
|
core,
|
||||||
|
InsertParams(key=key[:length], value=torch.arange(10, 10 + length)),
|
||||||
|
)
|
||||||
|
parent = core.match_prefix(MatchPrefixParams(key=key[:2])).best_match_node
|
||||||
|
child = core.match_prefix(MatchPrefixParams(key=key)).best_match_node
|
||||||
core.commit_backup(parent, torch.tensor([100, 101], dtype=torch.int64), {})
|
core.commit_backup(parent, torch.tensor([100, 101], dtype=torch.int64), {})
|
||||||
core.commit_backup(child, torch.tensor([102, 103], dtype=torch.int64), {})
|
core.commit_backup(child, torch.tensor([102, 103], dtype=torch.int64), {})
|
||||||
|
|
||||||
spec = core.build_storage_backup_spec(child, pass_prefix_keys=True)
|
spec = core.build_storage_backup_spec(child, pass_prefix_keys=True)
|
||||||
assert spec.host_value.tolist() == [102, 103]
|
assert spec.host_value.tolist() == [102, 103]
|
||||||
assert spec.token_ids == array("q", [7, 8])
|
assert spec.token_ids == array("q", [7, 8])
|
||||||
parent_hashes = mem_cache.get_hash_str(array("q", [1, 2]), None, 2)
|
hashes = get_storage_hash_str(key, page_size=2)
|
||||||
assert spec.prefix_keys == parent_hashes
|
assert spec.prefix_keys == hashes[:1]
|
||||||
assert spec.hash_value == mem_cache.get_hash_str(
|
assert spec.hash_value == hashes[1:]
|
||||||
array("q", [7, 8]), parent_hashes[-1], 2
|
|
||||||
)
|
|
||||||
assert spec.comp_xfers == {}
|
assert spec.comp_xfers == {}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -3912,10 +3912,18 @@ class UnifiedRadixCacheSuite:
|
|||||||
storage_dir, seq, extra_key=extra_key, cache_salt=cache_salt
|
storage_dir, seq, extra_key=extra_key, cache_salt=cache_salt
|
||||||
)
|
)
|
||||||
|
|
||||||
# A root anchor has no namespace of its own. The fetched span must use
|
# A root anchor has no namespace; probe and prefetch must use the request's.
|
||||||
# the request namespace supplied to prefetch_from_storage.
|
|
||||||
cons, _, _ = build_fixture(self.cfg)
|
cons, _, _ = build_fixture(self.cfg)
|
||||||
self._init_buffer_hicache(cons, storage_dir)
|
self._init_buffer_hicache(cons, storage_dir)
|
||||||
|
self.assertEqual(
|
||||||
|
cons.query_storage_hit_length(
|
||||||
|
cons.root_node_handle(),
|
||||||
|
array("q", seq),
|
||||||
|
extra_key=extra_key,
|
||||||
|
cache_salt=cache_salt,
|
||||||
|
),
|
||||||
|
len(seq),
|
||||||
|
)
|
||||||
root_req = "salted-root-prefetch"
|
root_req = "salted-root-prefetch"
|
||||||
cons.prefetch_from_storage(
|
cons.prefetch_from_storage(
|
||||||
root_req,
|
root_req,
|
||||||
|
|||||||
Reference in New Issue
Block a user