[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:
co-authored by
Shuwen Wang
alphabetc1
parent
f539c1fc65
commit
4c71d14fba
@@ -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),
|
||||||
|
|||||||
Reference in New Issue
Block a user