Preserve CP HiCache valid tails while padding physical pages
CP HiCache now keeps radix and scheduler-visible lengths as valid tokens while host/device transfers reserve and replay the padded physical page span. Exact valid-tail write, insertion, and match paths no longer fall back to page-flooring; the physical owner-lane contract still uses padded page metadata. Constraint: Scheduler prefix indices must never include padded tail locs. Constraint: Host/device transfer and owner-lane admission remain page-based. Rejected: Pad to cp_size or 2*cp_size pages | wastes KV and recreates short-tail fallback behavior. Rejected: Expose padded locs through load_cp return | would leak fake tokens into req.prefix_indices. Confidence: medium Scope-risk: moderate Directive: Do not implement split-inside-tail by duplicating page_owners without a page-sharing/refcount design. Tested: local py_compile for touched CP HiCache/radix/controller files and tests. Tested: remote g0034 CP HiCache impacted suites: 143 passed, 5 warnings. Tested: remote g0034 CP shared KV C1-C5 suite: 122 passed, 5 warnings. Not-tested: full local pytest, blocked by missing runtime dependencies such as orjson/starlette. Not-tested: CUDA E2E runtime for this commit. Co-authored-by: OmX <omx@oh-my-codex.dev>
This commit is contained in:
@@ -38,7 +38,10 @@ from sglang.srt.layers.dp_attention import (
|
||||
get_attention_tp_size,
|
||||
is_dp_attention_enabled,
|
||||
)
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import (
|
||||
CpSharedKVLayout,
|
||||
pad_token_locs_to_page_boundary,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool, NSATokenToKVPool
|
||||
from sglang.srt.mem_cache.page_index_utils import validate_page_aligned_token_indices
|
||||
from sglang.srt.utils import get_device_module
|
||||
@@ -931,21 +934,22 @@ class HiCacheController:
|
||||
) -> HiCacheWriteReservation | HiCacheWriteFailure:
|
||||
from sglang.srt.mem_cache.hiradix_cache import CpHiCacheNodeMetadata
|
||||
|
||||
page_size = self.page_size
|
||||
layout = self.cp_shared_kv_layout
|
||||
owned_mask = layout.owned_by_this_rank(device_indices)
|
||||
owned_positions = owned_mask.nonzero(as_tuple=True)[0].cpu()
|
||||
logical_len = len(device_indices)
|
||||
padded_device_indices = pad_token_locs_to_page_boundary(
|
||||
device_indices,
|
||||
page_size,
|
||||
name="CP HiCache write device_indices",
|
||||
)
|
||||
padded_len = int(padded_device_indices.numel())
|
||||
owned_mask = layout.owned_by_this_rank(padded_device_indices)
|
||||
owned_positions = owned_mask.nonzero(as_tuple=True)[0].cpu()
|
||||
# Capture the global owner pattern (one int8 per logical PAGE; identical
|
||||
# on all CP ranks since it's a pure function of the logical page ids).
|
||||
# load_cp will replay this pattern via alloc_pages_with_owners so the
|
||||
# saved owned_positions correctly index the new allocation.
|
||||
page_size = self.page_size
|
||||
if logical_len % page_size != 0:
|
||||
raise ValueError(
|
||||
f"_write_cp expects page-aligned device_indices, got "
|
||||
f"logical_len={logical_len} page_size={page_size}"
|
||||
)
|
||||
page_first_locs = device_indices[::page_size]
|
||||
page_first_locs = padded_device_indices[::page_size]
|
||||
logical_pages = torch.div(page_first_locs, page_size, rounding_mode="floor")
|
||||
page_owners = layout.owner_for_logical_pages(logical_pages).to(
|
||||
dtype=torch.int8, device="cpu"
|
||||
@@ -960,6 +964,7 @@ class HiCacheController:
|
||||
return HiCacheWriteReservation(
|
||||
metadata=CpHiCacheNodeMetadata(
|
||||
logical_len=logical_len,
|
||||
padded_len=padded_len,
|
||||
owned_positions=owned_positions,
|
||||
host_indices=empty,
|
||||
page_owners=page_owners,
|
||||
@@ -973,7 +978,7 @@ class HiCacheController:
|
||||
draft_host_indices=(empty.clone() if self.has_draft_hicache else None),
|
||||
)
|
||||
|
||||
owned_logical_indices = device_indices[owned_mask]
|
||||
owned_logical_indices = padded_device_indices[owned_mask]
|
||||
physical_device_indices = layout.logical_locs_to_physical(owned_logical_indices)
|
||||
host_indices = self.mem_pool_host.alloc(len(physical_device_indices))
|
||||
if host_indices is None:
|
||||
@@ -1019,6 +1024,7 @@ class HiCacheController:
|
||||
return HiCacheWriteReservation(
|
||||
metadata=CpHiCacheNodeMetadata(
|
||||
logical_len=logical_len,
|
||||
padded_len=padded_len,
|
||||
owned_positions=owned_positions,
|
||||
host_indices=host_indices.cpu(),
|
||||
page_owners=page_owners,
|
||||
@@ -1350,24 +1356,47 @@ class HiCacheController:
|
||||
if device_indices is None:
|
||||
return None
|
||||
|
||||
logical_len_expected = sum(node.host_len for node in nodes_to_load)
|
||||
if device_indices.numel() != logical_len_expected:
|
||||
padded_len_expected = sum(
|
||||
int(getattr(node.cp_hicache, "padded_len", node.host_len))
|
||||
for node in nodes_to_load
|
||||
)
|
||||
if device_indices.numel() != padded_len_expected:
|
||||
self.mem_pool_device_allocator.free(device_indices)
|
||||
raise RuntimeError(
|
||||
"alloc_pages_with_owners returned unexpected length: "
|
||||
f"got {device_indices.numel()}, expected {logical_len_expected}"
|
||||
f"got {device_indices.numel()}, expected {padded_len_expected}"
|
||||
)
|
||||
|
||||
host_chunks = []
|
||||
draft_host_chunks = []
|
||||
physical_chunks = []
|
||||
visible_chunks = []
|
||||
offset = 0
|
||||
for node in nodes_to_load:
|
||||
node_device_indices = device_indices[offset : offset + node.host_len]
|
||||
offset += node.host_len
|
||||
meta = node.cp_hicache
|
||||
valid_len = int(getattr(meta, "valid_len", node.host_len))
|
||||
padded_len = int(getattr(meta, "padded_len", node.host_len))
|
||||
if valid_len != int(node.host_len):
|
||||
self.mem_pool_device_allocator.free(device_indices)
|
||||
raise RuntimeError(
|
||||
"CP HiCache load metadata valid length mismatch: "
|
||||
f"node_id={getattr(node, 'id', '?')} "
|
||||
f"host_len={node.host_len} valid_len={valid_len}"
|
||||
)
|
||||
if padded_len < valid_len:
|
||||
self.mem_pool_device_allocator.free(device_indices)
|
||||
raise RuntimeError(
|
||||
"CP HiCache load metadata padded length is shorter than "
|
||||
"valid length: "
|
||||
f"node_id={getattr(node, 'id', '?')} "
|
||||
f"valid_len={valid_len} padded_len={padded_len}"
|
||||
)
|
||||
node_device_indices = device_indices[offset : offset + padded_len]
|
||||
offset += padded_len
|
||||
visible_chunks.append(node_device_indices[:valid_len])
|
||||
if self.has_draft_hicache:
|
||||
draft_host_indices = getattr(
|
||||
node.cp_hicache, "draft_host_indices", None
|
||||
meta, "draft_host_indices", None
|
||||
)
|
||||
if draft_host_indices is None:
|
||||
self.mem_pool_device_allocator.free(device_indices)
|
||||
@@ -1378,7 +1407,7 @@ class HiCacheController:
|
||||
else:
|
||||
draft_host_indices = None
|
||||
|
||||
owned_positions = node.cp_hicache.owned_positions.to(device_indices.device)
|
||||
owned_positions = meta.owned_positions.to(device_indices.device)
|
||||
if owned_positions.numel() == 0:
|
||||
continue
|
||||
selected_logical_locs = node_device_indices[owned_positions]
|
||||
@@ -1397,10 +1426,18 @@ class HiCacheController:
|
||||
physical_chunks.append(
|
||||
self.cp_shared_kv_layout.logical_locs_to_physical(selected_logical_locs)
|
||||
)
|
||||
host_chunks.append(node.cp_hicache.host_indices)
|
||||
host_chunks.append(meta.host_indices)
|
||||
if draft_host_indices is not None:
|
||||
draft_host_chunks.append(draft_host_indices)
|
||||
|
||||
visible_device_indices = (
|
||||
torch.cat(visible_chunks)
|
||||
if visible_chunks
|
||||
else torch.empty(
|
||||
(0,), dtype=device_indices.dtype, device=device_indices.device
|
||||
)
|
||||
)
|
||||
|
||||
if not host_chunks:
|
||||
# Keep CP load ACK rows identical across ranks. A zero-owned rank
|
||||
# still queues a zero-length op with the logical node id so a later
|
||||
@@ -1414,7 +1451,7 @@ class HiCacheController:
|
||||
node_id,
|
||||
)
|
||||
)
|
||||
return device_indices
|
||||
return visible_device_indices
|
||||
|
||||
host_indices = torch.cat(host_chunks)
|
||||
physical_device_indices = torch.cat(physical_chunks)
|
||||
@@ -1450,7 +1487,7 @@ class HiCacheController:
|
||||
)
|
||||
)
|
||||
|
||||
return device_indices
|
||||
return visible_device_indices
|
||||
|
||||
def move_indices(self, op: CacheOperation, mem_pool_host=None):
|
||||
mem_pool_host = mem_pool_host or self.mem_pool_host
|
||||
|
||||
@@ -6,6 +6,53 @@ import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
def pad_token_locs_to_page_boundary(
|
||||
token_locs: torch.Tensor,
|
||||
page_size: int,
|
||||
*,
|
||||
name: str = "token_locs",
|
||||
) -> torch.Tensor:
|
||||
"""Pad a valid token-loc span to the end of its physical tail page.
|
||||
|
||||
The returned tensor keeps the original valid locs first and appends only
|
||||
tail-page slack locs. It does not pad to CP size and it does not create
|
||||
fake request tokens.
|
||||
"""
|
||||
|
||||
if page_size <= 0:
|
||||
raise ValueError(f"page_size must be positive, got {page_size}")
|
||||
valid_len = int(token_locs.numel())
|
||||
if valid_len == 0:
|
||||
return token_locs
|
||||
|
||||
padding_len = (-valid_len) % page_size
|
||||
if padding_len == 0:
|
||||
return token_locs
|
||||
|
||||
if not token_locs.is_cuda:
|
||||
first_offset = int(torch.remainder(token_locs[0], page_size).item())
|
||||
if first_offset != 0:
|
||||
raise ValueError(
|
||||
f"{name} must start at a page boundary before padding, got "
|
||||
f"first_loc={int(token_locs[0].item())} page_size={page_size}"
|
||||
)
|
||||
last_offset = int(torch.remainder(token_locs[-1], page_size).item())
|
||||
if last_offset + padding_len != page_size - 1:
|
||||
raise ValueError(
|
||||
f"{name} does not end inside the expected tail page: "
|
||||
f"last_loc={int(token_locs[-1].item())} valid_len={valid_len} "
|
||||
f"padding_len={padding_len} page_size={page_size}"
|
||||
)
|
||||
|
||||
pad_offsets = torch.arange(
|
||||
1,
|
||||
padding_len + 1,
|
||||
dtype=token_locs.dtype,
|
||||
device=token_locs.device,
|
||||
)
|
||||
return torch.cat([token_locs, token_locs[-1:] + pad_offsets])
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CpSharedKVLayout:
|
||||
"""Map CP-group logical KV locations to per-rank physical KV locations.
|
||||
|
||||
@@ -33,7 +33,10 @@ from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
MatchPrefixParams,
|
||||
MatchResult,
|
||||
)
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import (
|
||||
CpSharedKVLayout,
|
||||
pad_token_locs_to_page_boundary,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool import (
|
||||
MHATokenToKVPool,
|
||||
MLATokenToKVPool,
|
||||
@@ -44,7 +47,6 @@ from sglang.srt.mem_cache.radix_cache import (
|
||||
RadixKey,
|
||||
TreeNode,
|
||||
compute_node_hash_values,
|
||||
page_align_keys,
|
||||
split_node_hash_value,
|
||||
)
|
||||
from sglang.srt.mem_cache.utils import convert_to_bigram_key
|
||||
@@ -124,6 +126,8 @@ def _compute_shared_hicache_token_capacities(
|
||||
|
||||
@dataclass
|
||||
class CpHiCacheNodeMetadata:
|
||||
# Legacy name kept for existing call sites. Under the page-aligned cache
|
||||
# contract this is the radix/scheduler-visible valid token length.
|
||||
logical_len: int
|
||||
owned_positions: torch.Tensor
|
||||
host_indices: torch.Tensor
|
||||
@@ -138,6 +142,9 @@ class CpHiCacheNodeMetadata:
|
||||
# NEW: needed at split() time to convert split_len into a page index
|
||||
# without plumbing it from the caller.
|
||||
page_size: int
|
||||
# Physical host/device span. When omitted, legacy callers keep the old
|
||||
# page-aligned invariant by using logical_len as the physical length.
|
||||
padded_len: Optional[int] = None
|
||||
draft_host_indices: Optional[torch.Tensor] = None
|
||||
|
||||
def __post_init__(self):
|
||||
@@ -145,9 +152,19 @@ class CpHiCacheNodeMetadata:
|
||||
raise ValueError(f"logical_len must be non-negative, got {self.logical_len}")
|
||||
if self.page_size <= 0:
|
||||
raise ValueError(f"page_size must be positive, got {self.page_size}")
|
||||
if self.logical_len % self.page_size != 0:
|
||||
if self.padded_len is None:
|
||||
self.padded_len = self.logical_len
|
||||
self.padded_len = int(self.padded_len)
|
||||
if self.padded_len < 0:
|
||||
raise ValueError(f"padded_len must be non-negative, got {self.padded_len}")
|
||||
if self.padded_len < self.logical_len:
|
||||
raise ValueError(
|
||||
f"logical_len ({self.logical_len}) must be a multiple of "
|
||||
f"padded_len ({self.padded_len}) must be >= logical_len "
|
||||
f"({self.logical_len})"
|
||||
)
|
||||
if self.padded_len % self.page_size != 0:
|
||||
raise ValueError(
|
||||
f"padded_len ({self.padded_len}) must be a multiple of "
|
||||
f"page_size ({self.page_size})"
|
||||
)
|
||||
self.owned_positions = self.owned_positions.to(
|
||||
@@ -163,11 +180,11 @@ class CpHiCacheNodeMetadata:
|
||||
self.draft_host_indices = self.draft_host_indices.to(
|
||||
device="cpu", dtype=torch.int64
|
||||
).detach().clone()
|
||||
expected_num_pages = self.logical_len // self.page_size
|
||||
expected_num_pages = self.padded_len // self.page_size
|
||||
if self.page_owners.numel() != expected_num_pages:
|
||||
raise ValueError(
|
||||
f"page_owners length ({self.page_owners.numel()}) must equal "
|
||||
f"logical_len/page_size ({expected_num_pages})"
|
||||
f"padded_len/page_size ({expected_num_pages})"
|
||||
)
|
||||
if expected_num_pages > 0 and bool((self.page_owners < 0).any()):
|
||||
raise ValueError("page_owners entries must be non-negative")
|
||||
@@ -186,14 +203,21 @@ class CpHiCacheNodeMetadata:
|
||||
)
|
||||
if self.owned_positions.numel() > 0:
|
||||
if torch.any(self.owned_positions < 0) or torch.any(
|
||||
self.owned_positions >= self.logical_len
|
||||
self.owned_positions >= self.padded_len
|
||||
):
|
||||
raise ValueError("owned_positions must be in [0, logical_len)")
|
||||
raise ValueError(
|
||||
"owned_positions must be in [0, padded_len) "
|
||||
"(legacy [0, logical_len) for unpadded metadata)"
|
||||
)
|
||||
if torch.any(self.owned_positions[1:] <= self.owned_positions[:-1]):
|
||||
raise ValueError(
|
||||
"owned_positions must be sorted and strictly increasing"
|
||||
)
|
||||
|
||||
@property
|
||||
def valid_len(self) -> int:
|
||||
return self.logical_len
|
||||
|
||||
def split(
|
||||
self, split_len: int
|
||||
) -> tuple["CpHiCacheNodeMetadata", "CpHiCacheNodeMetadata"]:
|
||||
@@ -216,6 +240,7 @@ class CpHiCacheNodeMetadata:
|
||||
host_indices=self.host_indices[parent_mask],
|
||||
page_owners=self.page_owners[:split_pages],
|
||||
page_size=self.page_size,
|
||||
padded_len=split_len,
|
||||
draft_host_indices=(
|
||||
self.draft_host_indices[parent_mask]
|
||||
if self.draft_host_indices is not None
|
||||
@@ -228,6 +253,7 @@ class CpHiCacheNodeMetadata:
|
||||
host_indices=self.host_indices[child_mask],
|
||||
page_owners=self.page_owners[split_pages:],
|
||||
page_size=self.page_size,
|
||||
padded_len=self.padded_len - split_len,
|
||||
draft_host_indices=(
|
||||
self.draft_host_indices[child_mask]
|
||||
if self.draft_host_indices is not None
|
||||
@@ -780,13 +806,13 @@ class HiRadixCache(RadixCache):
|
||||
logical_len = int(device_indices.numel())
|
||||
if logical_len == 0:
|
||||
return self._cp_zero_counts(layout.cp_size)
|
||||
if logical_len % self.page_size != 0:
|
||||
raise RuntimeError(
|
||||
"CP HiCache capacity planning requires page-aligned device indices: "
|
||||
f"logical_len={logical_len} page_size={self.page_size}"
|
||||
)
|
||||
padded_device_indices = pad_token_locs_to_page_boundary(
|
||||
device_indices,
|
||||
self.page_size,
|
||||
name="CP HiCache capacity device_indices",
|
||||
)
|
||||
logical_pages = torch.div(
|
||||
device_indices[:: self.page_size],
|
||||
padded_device_indices[:: self.page_size],
|
||||
self.page_size,
|
||||
rounding_mode="floor",
|
||||
)
|
||||
@@ -904,6 +930,7 @@ class HiRadixCache(RadixCache):
|
||||
|
||||
page_owners: List[int] = []
|
||||
host_hit_len = 0
|
||||
physical_hit_len = 0
|
||||
for node in nodes_to_load:
|
||||
metadata = getattr(node, "cp_hicache", None)
|
||||
if metadata is None:
|
||||
@@ -912,26 +939,35 @@ class HiRadixCache(RadixCache):
|
||||
f"load_node_id={getattr(node, 'id', '?')} root_node_id={node_id}"
|
||||
)
|
||||
node_host_len = int(self._node_host_len(node))
|
||||
if node_host_len % self.page_size != 0:
|
||||
node_padded_len = int(getattr(metadata, "padded_len", node_host_len))
|
||||
if node_padded_len % self.page_size != 0:
|
||||
raise RuntimeError(
|
||||
"CP HiCache load-back requires page-aligned host lengths: "
|
||||
"CP HiCache load-back requires page-aligned physical lengths: "
|
||||
f"load_node_id={getattr(node, 'id', '?')} "
|
||||
f"host_len={node_host_len} page_size={self.page_size}"
|
||||
f"padded_len={node_padded_len} page_size={self.page_size}"
|
||||
)
|
||||
if int(getattr(metadata, "logical_len", node_host_len)) != node_host_len:
|
||||
if int(getattr(metadata, "valid_len", node_host_len)) != node_host_len:
|
||||
raise RuntimeError(
|
||||
"CP HiCache load-back metadata length mismatch: "
|
||||
f"load_node_id={getattr(node, 'id', '?')} "
|
||||
f"host_len={node_host_len} metadata_len={metadata.logical_len}"
|
||||
f"host_len={node_host_len} metadata_len={metadata.valid_len}"
|
||||
)
|
||||
if node_padded_len < node_host_len:
|
||||
raise RuntimeError(
|
||||
"CP HiCache load-back physical length is shorter than "
|
||||
"valid host length: "
|
||||
f"load_node_id={getattr(node, 'id', '?')} "
|
||||
f"host_len={node_host_len} padded_len={node_padded_len}"
|
||||
)
|
||||
page_owners.extend(int(owner) for owner in metadata.page_owners.tolist())
|
||||
host_hit_len += node_host_len
|
||||
physical_hit_len += node_padded_len
|
||||
|
||||
expected_pages = host_hit_len // self.page_size
|
||||
expected_pages = physical_hit_len // self.page_size
|
||||
if len(page_owners) != expected_pages:
|
||||
raise RuntimeError(
|
||||
"CP HiCache load-back page owner count mismatch: "
|
||||
f"node_id={node_id} host_hit_len={host_hit_len} "
|
||||
f"node_id={node_id} physical_hit_len={physical_hit_len} "
|
||||
f"page_size={self.page_size} page_owners={len(page_owners)}"
|
||||
)
|
||||
|
||||
@@ -1956,9 +1992,6 @@ class HiRadixCache(RadixCache):
|
||||
return 0
|
||||
|
||||
key, _ = self.maybe_bigram_convert(key)
|
||||
if self.page_size != 1:
|
||||
page_aligned_len = len(key) // self.page_size * self.page_size
|
||||
key = key[:page_aligned_len]
|
||||
|
||||
if len(key) == 0:
|
||||
return 0
|
||||
@@ -2016,7 +2049,6 @@ class HiRadixCache(RadixCache):
|
||||
|
||||
token_ids = req.fill_ids
|
||||
keys = convert_to_bigram_key(token_ids) if self.is_eagle else token_ids
|
||||
keys = page_align_keys(keys, self.page_size)
|
||||
existing_prefix_len = self._probe_existing_radix_prefix_len_no_split(
|
||||
RadixKey(
|
||||
keys,
|
||||
@@ -2820,11 +2852,16 @@ class HiRadixCache(RadixCache):
|
||||
|
||||
self.ongoing_load_back[last_hit_node.id] = last_hit_node
|
||||
offset = 0
|
||||
physical_loaded_len = 0
|
||||
for loaded_node in nodes_to_load:
|
||||
host_len = self._node_host_len(loaded_node)
|
||||
loaded_node.value = device_indices[offset : offset + host_len].clone()
|
||||
offset += host_len
|
||||
self.evictable_size_ += len(device_indices)
|
||||
metadata = getattr(loaded_node, "cp_hicache", None)
|
||||
physical_loaded_len += int(
|
||||
getattr(metadata, "padded_len", host_len)
|
||||
)
|
||||
self.evictable_size_ += physical_loaded_len
|
||||
self.inc_lock_ref(last_hit_node)
|
||||
|
||||
if self.metrics_collector is not None:
|
||||
@@ -2836,9 +2873,10 @@ class HiRadixCache(RadixCache):
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"[HiCache-load] load_back CP SUCCESS: node_id=%d loaded_tokens=%d",
|
||||
"[HiCache-load] load_back CP SUCCESS: node_id=%d loaded_tokens=%d physical_tokens=%d",
|
||||
last_hit_node.id,
|
||||
len(device_indices),
|
||||
physical_loaded_len,
|
||||
)
|
||||
return device_indices
|
||||
|
||||
@@ -3150,7 +3188,7 @@ class HiRadixCache(RadixCache):
|
||||
)
|
||||
|
||||
page_aligned_len = len(key)
|
||||
if self.page_size != 1:
|
||||
if self.page_size != 1 and not self._uses_cp_hicache:
|
||||
page_aligned_len = len(key) // self.page_size * self.page_size
|
||||
key = key[:page_aligned_len]
|
||||
|
||||
@@ -3177,7 +3215,7 @@ class HiRadixCache(RadixCache):
|
||||
host_hit_length = 0
|
||||
last_host_node = last_node
|
||||
if self._uses_cp_hicache:
|
||||
while last_node.evicted:
|
||||
while last_node != self.root_node and last_node.evicted:
|
||||
host_hit_length += self._node_host_len(last_node)
|
||||
last_node = last_node.parent
|
||||
while (
|
||||
|
||||
@@ -209,9 +209,10 @@ def _key_match_paged(key0: RadixKey, key1: RadixKey, page_size: int):
|
||||
|
||||
i = 0
|
||||
while i < min_len:
|
||||
if key0.token_ids[i : i + page_size] != key1.token_ids[i : i + page_size]:
|
||||
step = min(page_size, min_len - i)
|
||||
if key0.token_ids[i : i + step] != key1.token_ids[i : i + step]:
|
||||
break
|
||||
i += page_size
|
||||
i += step
|
||||
|
||||
return i
|
||||
|
||||
@@ -486,7 +487,8 @@ class RadixCache(BasePrefixCache):
|
||||
|
||||
# Maybe convert to bigram keys for EAGLE
|
||||
keys = convert_to_bigram_key(token_ids) if self.is_eagle else token_ids
|
||||
keys = page_align_keys(keys, self.page_size)
|
||||
if not getattr(self, "_uses_cp_hicache", False):
|
||||
keys = page_align_keys(keys, self.page_size)
|
||||
values = kv_indices[: len(keys)].to(dtype=torch.int64, copy=True)
|
||||
radix_key = RadixKey(keys, req.extra_key, is_bigram=self.is_eagle)
|
||||
prepared_cp_backup = getattr(req, "cp_hicache_prepared_backup", None)
|
||||
@@ -542,7 +544,8 @@ class RadixCache(BasePrefixCache):
|
||||
|
||||
# Maybe convert to bigram keys for EAGLE
|
||||
keys = convert_to_bigram_key(token_ids) if self.is_eagle else token_ids
|
||||
keys = page_align_keys(keys, self.page_size)
|
||||
if not getattr(self, "_uses_cp_hicache", False):
|
||||
keys = page_align_keys(keys, self.page_size)
|
||||
values = kv_indices[: len(keys)].to(dtype=torch.int64, copy=True)
|
||||
radix_key = RadixKey(keys, req.extra_key, is_bigram=self.is_eagle)
|
||||
prepared_cp_backup = getattr(req, "cp_hicache_prepared_backup", None)
|
||||
|
||||
Reference in New Issue
Block a user