[HiCache][LoRA] Simplify decode offload hash inputs (#39162)

Co-authored-by: Shuwen Wang <47200617+alphabetc1@users.noreply.github.com>
Co-authored-by: alphabetc1 <2508695655@qq.com>
This commit is contained in:
Yanbin Jiang
2026-09-14 11:23:31 +08:00
committed by GitHub
co-authored by Shuwen Wang alphabetc1
parent f539c1fc65
commit 4c71d14fba
2 changed files with 7 additions and 15 deletions
@@ -139,9 +139,7 @@ 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, 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
@@ -274,12 +272,7 @@ 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( page_hashes = self._compute_prefix_hash(req, incremental_tokens, prior_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,
@@ -288,12 +281,10 @@ 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( def _compute_prefix_hash(self, req: Req, tokens, prior_hash=""):
self, tokens, prior_hash="", extra_key=None, cache_salt=None
):
"""Match prefill storage hashes.""" """Match prefill storage hashes."""
page_hashes = [] page_hashes = []
last_hash = prior_hash or storage_namespace_seed(extra_key, cache_salt) last_hash = prior_hash or storage_namespace_seed(req.extra_key, req.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)
@@ -138,8 +138,9 @@ class TestReleaseFinishedReq(unittest.TestCase):
]: ]:
with self.subTest(extra_key=extra_key, cache_salt=cache_salt): with self.subTest(extra_key=extra_key, cache_salt=cache_salt):
namespace = dict(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) req = SimpleNamespace(**namespace)
tail = manager._compute_prefix_hash(tokens[4:], prefix[-1], **namespace) prefix = manager._compute_prefix_hash(req, tokens[:4])
tail = manager._compute_prefix_hash(req, tokens[4:], prefix[-1])
self.assertEqual( self.assertEqual(
prefix + tail, prefix + tail,
get_storage_hash_str(RadixKey(tokens, **namespace), page_size=2), get_storage_hash_str(RadixKey(tokens, **namespace), page_size=2),