[Perf] Walk the radix tree by offset instead of re-slicing token storage (ported from #36507) (#37324)
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
57c26a84e0
commit
cf3173aeb9
@@ -180,10 +180,18 @@ class RadixKey:
|
||||
|
||||
def match(self, other: RadixKey, page_size: int = 1) -> int:
|
||||
"""Logical-unit prefix length shared with ``other``. Result is rounded down to ``page_size``."""
|
||||
return self.match_at(other, offset=0, page_size=page_size)
|
||||
|
||||
def match_at(self, other: RadixKey, offset: int, page_size: int = 1) -> int:
|
||||
"""Match without slicing while preserving bigram boundaries and limit semantics."""
|
||||
self._check_compatible(other)
|
||||
if self.is_bigram != other.is_bigram:
|
||||
raise ValueError("RadixKey operations require matching bigram modes")
|
||||
if offset < 0 or offset > len(other):
|
||||
raise IndexError(f"RadixKey offset out of range: {offset}")
|
||||
t0, t1 = self.token_ids, other.token_ids
|
||||
assert type(t0) is type(t1), (type(t0), type(t1))
|
||||
n = min(len(t0), len(t1))
|
||||
n = min(self._raw_len(), other._raw_len() - offset)
|
||||
|
||||
# Exponential search for the first diverging token: gallop in doubling
|
||||
# windows (one C-level slice compare each), then binary-search the window
|
||||
@@ -193,10 +201,10 @@ class RadixKey:
|
||||
step = 1
|
||||
while lo < n:
|
||||
hi = lo + step if lo + step < n else n
|
||||
if t0[lo:hi] != t1[lo:hi]:
|
||||
if t0[lo:hi] != t1[offset + lo : offset + hi]:
|
||||
while hi - lo > 1:
|
||||
mid = (lo + hi) // 2
|
||||
if t0[lo:mid] == t1[lo:mid]:
|
||||
if t0[lo:mid] == t1[offset + lo : offset + mid]:
|
||||
lo = mid
|
||||
else:
|
||||
hi = mid
|
||||
@@ -206,24 +214,37 @@ class RadixKey:
|
||||
step *= 2
|
||||
|
||||
if self.is_bigram:
|
||||
matched = max(0, min(matched_tokens - 1, len(self), len(other)))
|
||||
matched = max(0, min(matched_tokens - 1, len(self), len(other) - offset))
|
||||
return (matched // page_size) * page_size if page_size > 1 else matched
|
||||
|
||||
matched_tokens = min(matched_tokens, len(self), len(other))
|
||||
matched_tokens = min(matched_tokens, len(self), len(other) - offset)
|
||||
if page_size == 1:
|
||||
return matched_tokens
|
||||
return (matched_tokens // page_size) * page_size
|
||||
|
||||
def child_key(self, page_size: int = 1):
|
||||
"""Hashable dict-key for the first ``page_size`` logical units, namespaced by ``extra_key``."""
|
||||
return self.child_key_at(offset=0, page_size=page_size)
|
||||
|
||||
def child_key_at(self, offset: int, page_size: int = 1):
|
||||
"""Hashable child key at ``offset`` without slicing token storage."""
|
||||
if offset < 0 or offset + page_size > len(self):
|
||||
raise IndexError(
|
||||
f"RadixKey child range out of bounds: offset={offset}, "
|
||||
f"page_size={page_size}, len={len(self)}"
|
||||
)
|
||||
t = self.token_ids
|
||||
if self.is_bigram:
|
||||
if page_size == 1:
|
||||
plain = (t[0], t[1])
|
||||
plain = (t[offset], t[offset + 1])
|
||||
else:
|
||||
plain = tuple((t[j], t[j + 1]) for j in range(page_size))
|
||||
plain = tuple(
|
||||
(t[j], t[j + 1]) for j in range(offset, offset + page_size)
|
||||
)
|
||||
else:
|
||||
plain = t[0] if page_size == 1 else tuple(t[:page_size])
|
||||
plain = (
|
||||
t[offset] if page_size == 1 else tuple(t[offset : offset + page_size])
|
||||
)
|
||||
if self.cache_salt is not None:
|
||||
return ((self.extra_key, self.cache_salt), plain)
|
||||
return plain if self.extra_key is None else (self.extra_key, plain)
|
||||
|
||||
@@ -737,7 +737,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
||||
# nodes can also match, so we separately track the best device-resident
|
||||
# match for scheduler prefix indices and locking.
|
||||
node = self.root_node
|
||||
child_key = key.child_key(self.page_size)
|
||||
key_offset = 0
|
||||
child_key = key.child_key_at(key_offset, self.page_size)
|
||||
value: list[torch.Tensor] = []
|
||||
best_match_node = node
|
||||
best_match_device_node = node
|
||||
@@ -778,14 +779,14 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
||||
best_match_device_value_len = len(value)
|
||||
best_match_device_node = node
|
||||
|
||||
while len(key) > 0 and child_key in node.children:
|
||||
while key_offset < len(key) and child_key in node.children:
|
||||
child = node.children[child_key]
|
||||
|
||||
# HiCache: dead node (evicted + not backuped) — stop traversal
|
||||
if child.evicted and not child.backuped:
|
||||
break
|
||||
|
||||
prefix_len = child.key.match(key, page_size=self.page_size)
|
||||
prefix_len = child.key.match_at(key, key_offset, page_size=self.page_size)
|
||||
full_kv_hit_length += prefix_len
|
||||
if prefix_len < len(child.key):
|
||||
node, action = self._split_node(child.key, child, prefix_len)
|
||||
@@ -798,9 +799,9 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
|
||||
value.append(child.component_data[BASE_COMPONENT_TYPE].value)
|
||||
node = child
|
||||
_update_best_if_valid(node)
|
||||
key = key[prefix_len:]
|
||||
if len(key):
|
||||
child_key = key.child_key(self.page_size)
|
||||
key_offset += prefix_len
|
||||
if key_offset < len(key):
|
||||
child_key = key.child_key_at(key_offset, self.page_size)
|
||||
|
||||
return (
|
||||
value,
|
||||
|
||||
Reference in New Issue
Block a user