[BugFix] fix(hicache): fix two slot-reuse races in DecodeKVCacheOffloadManager (#24226)

This commit is contained in:
Kevin Flansburg
2026-05-20 18:20:19 -07:00
committed by GitHub
parent c4a7d12092
commit 643d44d699
3 changed files with 291 additions and 21 deletions
@@ -104,8 +104,22 @@ class DecodeKVCacheOffloadManager:
self.ongoing_offload = {}
self.ongoing_backup = {}
self.offloaded_state = {}
self.offload_inflight = {}
logger.info("Enable offload kv cache for decode side")
def _mark_offload_started(self, rid):
self.offload_inflight[rid] = self.offload_inflight.get(rid, 0) + 1
def _mark_offload_finished(self, rid):
count = self.offload_inflight.get(rid, 0)
if count <= 1:
self.offload_inflight.pop(rid, None)
else:
self.offload_inflight[rid] = count - 1
def _has_inflight_offload(self, rid):
return self.offload_inflight.get(rid, 0) > 0
def offload_kv_cache(self, req) -> bool:
"""Offload incremental KV cache for decode side."""
@@ -153,9 +167,11 @@ class DecodeKVCacheOffloadManager:
incremental_tokens = all_tokens[start:end]
incremental_indices = token_indices[start:end]
# Early free prefill-offloaded GPU memory
if state.prefill_len > 0 and state.inc_len == 0:
self.token_to_kv_pool_allocator.free(token_indices[: state.prefill_len])
# Prefill-aligned GPU slots are freed at request finish in
# _release_finished_req, NOT here. The decoding request
# continues to attend to those slots via req_to_token; freeing
# them mid-decode races with concurrent admission, which can
# reuse the slots and produce cross-pollinated KV reads.
# Asynchronously offload incremental KV cache from device to host
self.request_counter += 1
@@ -168,6 +184,7 @@ class DecodeKVCacheOffloadManager:
logger.error(f"Not enough host memory for request {req.rid}")
return False
self._mark_offload_started(req.rid)
self.ongoing_offload[ack_id] = (
req,
host_indices,
@@ -214,14 +231,7 @@ class DecodeKVCacheOffloadManager:
end,
) = self.ongoing_offload.pop(ack_id)
if req.finished():
self._release_finished_req(req, start)
else:
kv_indices = self.req_to_token_pool.req_to_token[
req.req_pool_idx, start:end
]
self.token_to_kv_pool_allocator.free(kv_indices)
self._mark_offload_finished(req.rid)
prior_hash = (
self.offloaded_state[req.rid].last_hash
if req.rid in self.offloaded_state
@@ -232,10 +242,34 @@ class DecodeKVCacheOffloadManager:
)
if req.rid in self.offloaded_state:
self.offloaded_state[req.rid].last_hash = last_hash
if req.finished() and not self._has_inflight_offload(req.rid):
state = self.offloaded_state.get(req.rid)
start_offset = state.prefill_len if state is not None else start
self._release_finished_req(req, start_offset)
finish_count -= 1
def _release_finished_req(self, req: Req, start_offset: int):
# Defensive guard: ReqToTokenPool.free sets req_pool_idx to None,
# so a previously-released request must be skipped here to avoid
# non-idempotent side effects (e.g. tree_cache.protected_size_
# double-decrement, host pool double-free).
if req.req_pool_idx is None or req.req_pool_idx == -1:
return
kv_committed_len = req.pop_committed_kv_cache()
# Free the prefill-aligned slots. Previously this was done
# eagerly in offload_kv_cache (mid-decode), which raced with
# concurrent admission. Now consolidated here at request
# finish, where the request is guaranteed to no longer attend
# to those slots.
state = self.offloaded_state.get(req.rid)
if state is not None and state.prefill_len > 0:
prefill_indices = self.req_to_token_pool.req_to_token[
req.req_pool_idx, : state.prefill_len
]
self.token_to_kv_pool_allocator.free(prefill_indices)
start = start_offset
end = kv_committed_len
# Free the incremental part of the request (DSA-aware)
@@ -296,7 +330,9 @@ class DecodeKVCacheOffloadManager:
def finalize_release_on_finish(self, req: Req):
"""Free any remaining tail KV that was not offloaded due to non-aligned length."""
if req.req_pool_idx == -1:
# ReqToTokenPool.free sets req_pool_idx to None on release, so
# guard against both sentinels here.
if req.req_pool_idx is None or req.req_pool_idx == -1:
return
state = self.offloaded_state.get(req.rid)
if state is None:
@@ -305,13 +341,13 @@ class DecodeKVCacheOffloadManager:
else:
prefill_len = state.prefill_len
inc_len = state.inc_len
# If no incremental offload ever happened, the prefill-aligned part was never freed.
# Free the prefill portion on request finish to avoid leaks.
if prefill_len > 0 and inc_len == 0:
token_indices = self.req_to_token_pool.req_to_token[req.req_pool_idx]
self.token_to_kv_pool_allocator.free(token_indices[:prefill_len])
logger.info(
f"Finalize release: freed prefill-aligned KV for req {req.rid}, len:{prefill_len}"
# Prefill-aligned slots are freed by _release_finished_req. Make
# sure state exists so it can find prefill_len.
if state is None:
self.offloaded_state[req.rid] = OffloadedState(
prefill_len=prefill_len, inc_len=0, last_hash=None
)
start_offset = prefill_len + inc_len
if self._has_inflight_offload(req.rid):
return
start_offset = prefill_len
self._release_finished_req(req, start_offset)