[HiCache] Prevent move_hybrid_indices from polluting radix-tree node host state (#23427)

Co-authored-by: hzh0425 <hzh0425@apache.org>
This commit is contained in:
Yujun Dong
2026-04-25 14:27:42 +08:00
committed by GitHub
co-authored by hzh0425
parent 82254bd9c5
commit 21835fb0af
@@ -261,7 +261,9 @@ class HybridCacheController(BaseHiCacheController):
if not self.write_queue: if not self.write_queue:
return return
op = CacheOperation.merge_ops(self.write_queue) op = CacheOperation.merge_ops(self.write_queue)
host_indices, device_indices = self.move_hybrid_indices(op) host_indices, device_indices, resolved_pool_transfers = (
self.move_hybrid_indices(op)
)
self.write_queue.clear() self.write_queue.clear()
start_event = device_module.Event() start_event = device_module.Event()
finish_event = device_module.Event() finish_event = device_module.Event()
@@ -273,14 +275,14 @@ class HybridCacheController(BaseHiCacheController):
host_indices, host_indices,
device_indices, device_indices,
self.io_backend, self.io_backend,
pool_transfers=op.pool_transfers, pool_transfers=resolved_pool_transfers,
) )
finish_event.record() finish_event.record()
self._record_transfer_indices_on_stream( self._record_transfer_indices_on_stream(
self.write_stream, self.write_stream,
host_indices, host_indices,
device_indices, device_indices,
op.pool_transfers, resolved_pool_transfers,
) )
self.ack_write_queue.append(HiCacheAck(start_event, finish_event, op.node_ids)) self.ack_write_queue.append(HiCacheAck(start_event, finish_event, op.node_ids))
@@ -326,7 +328,9 @@ class HybridCacheController(BaseHiCacheController):
return -1 return -1
producer_id = self.layer_done_counter.update_producer() producer_id = self.layer_done_counter.update_producer()
op = CacheOperation.merge_ops(self.load_queue) op = CacheOperation.merge_ops(self.load_queue)
host_indices, device_indices = self.move_hybrid_indices(op) host_indices, device_indices, resolved_pool_transfers = (
self.move_hybrid_indices(op)
)
self.load_queue.clear() self.load_queue.clear()
producer_event = self.layer_done_counter.events[producer_id] producer_event = self.layer_done_counter.events[producer_id]
producer_event.start_event.record() producer_event.start_event.record()
@@ -339,14 +343,14 @@ class HybridCacheController(BaseHiCacheController):
device_indices, device_indices,
i, i,
self.io_backend, self.io_backend,
pool_transfers=op.pool_transfers, pool_transfers=resolved_pool_transfers,
) )
producer_event.complete(i) producer_event.complete(i)
self._record_transfer_indices_on_stream( self._record_transfer_indices_on_stream(
self.load_stream, self.load_stream,
host_indices, host_indices,
device_indices, device_indices,
op.pool_transfers, resolved_pool_transfers,
) )
self.ack_load_queue.append( self.ack_load_queue.append(
HiCacheAck( HiCacheAck(
@@ -445,16 +449,32 @@ class HybridCacheController(BaseHiCacheController):
kv_hit_pages * self.page_size, kv_hit_pages * self.page_size,
) )
def move_hybrid_indices(self, operation): def move_hybrid_indices(
self, operation: CacheOperation
) -> tuple[torch.Tensor, torch.Tensor, Optional[list[PoolTransfer]]]:
host_indices, device_indices = self.move_indices( host_indices, device_indices = self.move_indices(
operation.host_indices, operation.device_indices operation.host_indices, operation.device_indices
) )
resolved_pool_transfers = None
if operation.pool_transfers: if operation.pool_transfers:
resolved_pool_transfers = []
for transfer in operation.pool_transfers: for transfer in operation.pool_transfers:
transfer.host_indices, transfer.device_indices = self.move_indices( transfer_host_indices, transfer_device_indices = self.move_indices(
transfer.host_indices, transfer.device_indices transfer.host_indices, transfer.device_indices
) )
return host_indices, device_indices # Keep the original PoolTransfer unchanged because tree-owned
# transfers may still reference radix-tree host state. The
# controller only needs a normalized execution-time copy.
resolved_pool_transfers.append(
PoolTransfer(
name=transfer.name,
host_indices=transfer_host_indices,
device_indices=transfer_device_indices,
keys=transfer.keys,
hit_policy=transfer.hit_policy,
)
)
return host_indices, device_indices, resolved_pool_transfers
def _page_transfer(self, operation): def _page_transfer(self, operation):
# Transfer extra pools # Transfer extra pools