[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:
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user