Remove revoke queue after hit-then-alloc refactoring (#33613)
This commit is contained in:
@@ -382,7 +382,6 @@ class HiCacheController:
|
||||
self.backup_queue = Queue()
|
||||
|
||||
self.prefetch_hit_queue: Queue[StorageOperation] = Queue()
|
||||
self.prefetch_revoke_queue: Queue[str] = Queue()
|
||||
self.ack_backup_queue: Queue[StorageOperation] = Queue()
|
||||
self.host_mem_release_queue: Queue[torch.Tensor] = Queue()
|
||||
|
||||
@@ -659,7 +658,6 @@ class HiCacheController:
|
||||
self.backup_thread.join()
|
||||
self.prefetch_queue.queue.clear()
|
||||
self.backup_queue.queue.clear()
|
||||
self.prefetch_revoke_queue.queue.clear()
|
||||
self.prefetch_hit_queue.queue.clear()
|
||||
self.ack_backup_queue.queue.clear()
|
||||
self.host_mem_release_queue.queue.clear()
|
||||
@@ -1103,19 +1101,13 @@ class HiCacheController:
|
||||
)
|
||||
storage_hit_count = storage_hit_count_tensor.item()
|
||||
|
||||
if storage_hit_count < self.prefetch_threshold:
|
||||
# not to prefetch if not enough benefits
|
||||
self.prefetch_revoke_queue.put(operation.request_id)
|
||||
logger.debug(
|
||||
f"Revoking prefetch for request {operation.request_id} due to insufficient hits ({storage_hit_count})."
|
||||
)
|
||||
else:
|
||||
# Record hit count, so the scheduler thread will know the exact memory to allocate
|
||||
operation.hash_value = hash_value[
|
||||
: (storage_hit_count // self.page_size)
|
||||
]
|
||||
operation.storage_hit_count = storage_hit_count
|
||||
self.prefetch_hit_queue.put(operation)
|
||||
# Record the TP-synced hit count; the scheduler thread decides
|
||||
# at drain time whether to revoke (below threshold) or allocate.
|
||||
operation.hash_value = hash_value[
|
||||
: (storage_hit_count // self.page_size)
|
||||
]
|
||||
operation.storage_hit_count = storage_hit_count
|
||||
self.prefetch_hit_queue.put(operation)
|
||||
|
||||
except Empty:
|
||||
continue
|
||||
|
||||
@@ -1474,7 +1474,6 @@ class HiMambaRadixCache(MambaRadixCache):
|
||||
|
||||
def _drain_storage_control_queues_local(self):
|
||||
self._drain_storage_control_queues_impl(
|
||||
n_revoke=None,
|
||||
n_storage_hit=0,
|
||||
n_backup=None,
|
||||
n_release=None,
|
||||
@@ -1495,7 +1494,6 @@ class HiMambaRadixCache(MambaRadixCache):
|
||||
|
||||
def _drain_storage_control_queues_impl(
|
||||
self,
|
||||
n_revoke: Optional[int],
|
||||
n_storage_hit: Optional[int],
|
||||
n_backup: Optional[int],
|
||||
n_release: Optional[int],
|
||||
@@ -1513,10 +1511,6 @@ class HiMambaRadixCache(MambaRadixCache):
|
||||
drained += 1
|
||||
yield item
|
||||
|
||||
def _drain_revoke():
|
||||
for req_id in _drain_queue(cc.prefetch_revoke_queue, n_revoke):
|
||||
self._revoke_pending_prefetch(req_id)
|
||||
|
||||
def _drain_and_alloc_storage_hit():
|
||||
# The L3 hit count is now known, so reserve exactly that much host
|
||||
# KV memory. NOTE: alloc/evict here is rank-local but deterministic
|
||||
@@ -1534,6 +1528,10 @@ class HiMambaRadixCache(MambaRadixCache):
|
||||
# request was aborted while the storage query was in flight
|
||||
self._revoke_pending_prefetch(req_id)
|
||||
continue
|
||||
if operation.storage_hit_count < self.prefetch_threshold:
|
||||
# not to prefetch if not enough benefits
|
||||
self._revoke_pending_prefetch(req_id)
|
||||
continue
|
||||
|
||||
alloc_len = operation.storage_hit_count
|
||||
host_indices = cc.mem_pool_host.alloc(alloc_len)
|
||||
@@ -1587,7 +1585,6 @@ class HiMambaRadixCache(MambaRadixCache):
|
||||
host_indices = torch.cat(host_indices_list, dim=0)
|
||||
cc.mem_pool_host.free(host_indices)
|
||||
|
||||
_drain_revoke()
|
||||
_drain_and_alloc_storage_hit()
|
||||
_drain_backup()
|
||||
_drain_release()
|
||||
@@ -1698,7 +1695,6 @@ class HiMambaRadixCache(MambaRadixCache):
|
||||
|
||||
qsizes = torch.tensor(
|
||||
[
|
||||
cc.prefetch_revoke_queue.qsize(),
|
||||
cc.prefetch_hit_queue.qsize(),
|
||||
cc.ack_backup_queue.qsize(),
|
||||
cc.host_mem_release_queue.qsize(),
|
||||
@@ -1710,9 +1706,8 @@ class HiMambaRadixCache(MambaRadixCache):
|
||||
qsizes, op=torch.distributed.ReduceOp.MIN, group=self.tp_group
|
||||
)
|
||||
|
||||
n_revoke, n_storage_hit, n_backup, n_release = map(int, qsizes.tolist())
|
||||
n_storage_hit, n_backup, n_release = map(int, qsizes.tolist())
|
||||
self._drain_storage_control_queues_impl(
|
||||
n_revoke=n_revoke,
|
||||
n_storage_hit=n_storage_hit,
|
||||
n_backup=n_backup,
|
||||
n_release=n_release,
|
||||
|
||||
@@ -579,7 +579,6 @@ class HiRadixCache(RadixCache):
|
||||
cleanup even if queue sizes temporarily differ across ranks.
|
||||
"""
|
||||
self._drain_storage_control_queues_impl(
|
||||
n_revoke=None,
|
||||
n_storage_hit=0,
|
||||
n_backup=None,
|
||||
n_release=None,
|
||||
@@ -588,7 +587,6 @@ class HiRadixCache(RadixCache):
|
||||
|
||||
def _drain_storage_control_queues_impl(
|
||||
self,
|
||||
n_revoke: Optional[int],
|
||||
n_storage_hit: Optional[int],
|
||||
n_backup: Optional[int],
|
||||
n_release: Optional[int],
|
||||
@@ -606,10 +604,6 @@ class HiRadixCache(RadixCache):
|
||||
drained += 1
|
||||
yield item
|
||||
|
||||
def _drain_revoke():
|
||||
for req_id in _drain_queue(cc.prefetch_revoke_queue, n_revoke):
|
||||
self._revoke_pending_prefetch(req_id)
|
||||
|
||||
def _drain_and_alloc_storage_hit():
|
||||
# The L3 hit count is now known, so reserve exactly that much host
|
||||
# memory (this is the whole point: no over-allocation up front).
|
||||
@@ -628,6 +622,13 @@ class HiRadixCache(RadixCache):
|
||||
# request was aborted while the storage query was in flight
|
||||
self._revoke_pending_prefetch(req_id)
|
||||
continue
|
||||
if operation.storage_hit_count < self.prefetch_threshold:
|
||||
# not to prefetch if not enough benefits
|
||||
self._revoke_pending_prefetch(req_id)
|
||||
logger.debug(
|
||||
f"Revoking prefetch for request {req_id} due to insufficient hits ({operation.storage_hit_count})."
|
||||
)
|
||||
continue
|
||||
|
||||
alloc_len = operation.storage_hit_count
|
||||
host_indices = cc.mem_pool_host.alloc(alloc_len)
|
||||
@@ -676,7 +677,6 @@ class HiRadixCache(RadixCache):
|
||||
host_indices = torch.cat(host_indices_list, dim=0)
|
||||
cc.mem_pool_host.free(host_indices)
|
||||
|
||||
_drain_revoke()
|
||||
_drain_and_alloc_storage_hit()
|
||||
_drain_backup()
|
||||
_drain_release()
|
||||
@@ -997,7 +997,6 @@ class HiRadixCache(RadixCache):
|
||||
cache_controller = self.cache_controller
|
||||
storage_queue_sizes = (
|
||||
(
|
||||
cache_controller.prefetch_revoke_queue.qsize(),
|
||||
cache_controller.prefetch_hit_queue.qsize(),
|
||||
cache_controller.ack_backup_queue.qsize(),
|
||||
cache_controller.host_mem_release_queue.qsize(),
|
||||
@@ -1531,9 +1530,8 @@ class HiRadixCache(RadixCache):
|
||||
self.loading_check(finish_count=load_finish_count)
|
||||
|
||||
if self.enable_storage and storage_queue_sizes:
|
||||
n_revoke, n_storage_hit, n_backup, n_release = storage_queue_sizes[:4]
|
||||
n_storage_hit, n_backup, n_release = storage_queue_sizes[:3]
|
||||
self._drain_storage_control_queues_impl(
|
||||
n_revoke=n_revoke,
|
||||
n_storage_hit=n_storage_hit,
|
||||
n_backup=n_backup,
|
||||
n_release=n_release,
|
||||
@@ -1553,7 +1551,6 @@ class HiRadixCache(RadixCache):
|
||||
|
||||
qsizes = torch.tensor(
|
||||
[
|
||||
cc.prefetch_revoke_queue.qsize(),
|
||||
cc.prefetch_hit_queue.qsize(),
|
||||
cc.ack_backup_queue.qsize(),
|
||||
cc.host_mem_release_queue.qsize(),
|
||||
@@ -1562,9 +1559,8 @@ class HiRadixCache(RadixCache):
|
||||
)
|
||||
self._all_reduce_attn_groups(qsizes, torch.distributed.ReduceOp.MIN)
|
||||
|
||||
n_revoke, n_storage_hit, n_backup, n_release = map(int, qsizes.tolist())
|
||||
n_storage_hit, n_backup, n_release = map(int, qsizes.tolist())
|
||||
self._drain_storage_control_queues_impl(
|
||||
n_revoke=n_revoke,
|
||||
n_storage_hit=n_storage_hit,
|
||||
n_backup=n_backup,
|
||||
n_release=n_release,
|
||||
|
||||
@@ -1569,7 +1569,6 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
|
||||
def _drain_storage_control_queues_impl(
|
||||
self,
|
||||
n_revoke: Optional[int],
|
||||
n_storage_hit: Optional[int],
|
||||
n_backup: Optional[int],
|
||||
n_release: Optional[int],
|
||||
@@ -1588,10 +1587,6 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
drained += 1
|
||||
yield item
|
||||
|
||||
def _drain_revoke():
|
||||
for req_id in _drain_queue(cc.prefetch_revoke_queue, n_revoke):
|
||||
self._revoke_pending_prefetch(req_id)
|
||||
|
||||
def _drain_and_alloc_storage_hit():
|
||||
for operation in _drain_queue(cc.prefetch_hit_queue, n_storage_hit):
|
||||
req_id = operation.request_id
|
||||
@@ -1603,6 +1598,10 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
# request was aborted while the storage query was in flight
|
||||
self._revoke_pending_prefetch(req_id)
|
||||
continue
|
||||
if operation.storage_hit_count < self.prefetch_threshold:
|
||||
# not to prefetch if not enough benefits
|
||||
self._revoke_pending_prefetch(req_id)
|
||||
continue
|
||||
|
||||
alloc_len = operation.storage_hit_count
|
||||
host_indices = cc.mem_pool_host.alloc(alloc_len)
|
||||
@@ -1678,7 +1677,6 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
drained[pool_name] = (len(host_indices_list), released_tokens)
|
||||
return drained
|
||||
|
||||
_drain_revoke()
|
||||
_drain_and_alloc_storage_hit()
|
||||
_drain_backup()
|
||||
_drain_release()
|
||||
@@ -1689,7 +1687,6 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
extra_release_queues = getattr(cc, "extra_host_mem_release_queues", {})
|
||||
extra_pool_names = list(extra_release_queues)
|
||||
local_qsize_list = [
|
||||
cc.prefetch_revoke_queue.qsize(),
|
||||
cc.prefetch_hit_queue.qsize(),
|
||||
cc.ack_backup_queue.qsize(),
|
||||
cc.host_mem_release_queue.qsize(),
|
||||
@@ -1704,13 +1701,12 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
)
|
||||
self._all_reduce_attn_groups(qsizes, torch.distributed.ReduceOp.MIN)
|
||||
qsize_list = list(map(int, qsizes.tolist()))
|
||||
n_revoke, n_storage_hit, n_backup, n_release = qsize_list[:4]
|
||||
n_storage_hit, n_backup, n_release = qsize_list[:3]
|
||||
extra_release_counts = {
|
||||
pool_name: count
|
||||
for pool_name, count in zip(extra_pool_names, qsize_list[4:])
|
||||
for pool_name, count in zip(extra_pool_names, qsize_list[3:])
|
||||
}
|
||||
self._drain_storage_control_queues_impl(
|
||||
n_revoke=n_revoke,
|
||||
n_storage_hit=n_storage_hit,
|
||||
n_backup=n_backup,
|
||||
n_release=n_release,
|
||||
@@ -1834,7 +1830,6 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
)
|
||||
storage_queue_sizes = (
|
||||
(
|
||||
cc.prefetch_revoke_queue.qsize(),
|
||||
cc.prefetch_hit_queue.qsize(),
|
||||
cc.ack_backup_queue.qsize(),
|
||||
cc.host_mem_release_queue.qsize(),
|
||||
@@ -2025,16 +2020,15 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
self.loading_check(finish_count=load_finish_count)
|
||||
|
||||
if self.enable_storage and storage_queue_sizes:
|
||||
n_revoke, n_storage_hit, n_backup, n_release = storage_queue_sizes[:4]
|
||||
n_storage_hit, n_backup, n_release = storage_queue_sizes[:3]
|
||||
extra_release_counts = {
|
||||
pool_name: count
|
||||
for pool_name, count in zip(
|
||||
extra_pool_names,
|
||||
storage_queue_sizes[4:],
|
||||
storage_queue_sizes[3:],
|
||||
)
|
||||
}
|
||||
self._drain_storage_control_queues_impl(
|
||||
n_revoke=n_revoke,
|
||||
n_storage_hit=n_storage_hit,
|
||||
n_backup=n_backup,
|
||||
n_release=n_release,
|
||||
|
||||
Reference in New Issue
Block a user