[Bugfix][HiCache] measure load-back duration with CUDA events (#26411)
Co-authored-by: Claude Opus 4.7 <noreply@anthropic.com> Co-authored-by: vuuihc <vuuihc@users.noreply.github.com> Co-authored-by: Zhiqiang Xie <xiezhq@stanford.edu>
This commit is contained in:
co-authored by
Claude Opus 4.7
vuuihc
Zhiqiang Xie
parent
f8eac995aa
commit
095a817612
@@ -217,9 +217,9 @@ class DecodeKVCacheOffloadManager:
|
|||||||
def _check_offload_progress(self, finish_count):
|
def _check_offload_progress(self, finish_count):
|
||||||
"""Check the progress of offload from device to host."""
|
"""Check the progress of offload from device to host."""
|
||||||
while finish_count > 0:
|
while finish_count > 0:
|
||||||
_, finish_event, ack_list = self.cache_controller.ack_write_queue.pop(0)
|
ack = self.cache_controller.ack_write_queue.pop(0)
|
||||||
finish_event.synchronize()
|
ack.finish_event.synchronize()
|
||||||
for ack_id in ack_list:
|
for ack_id in ack.node_ids:
|
||||||
(
|
(
|
||||||
req,
|
req,
|
||||||
host_indices,
|
host_indices,
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ limitations under the License.
|
|||||||
import logging
|
import logging
|
||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
|
from functools import cache
|
||||||
from queue import Empty, Queue
|
from queue import Empty, Queue
|
||||||
from typing import TYPE_CHECKING, List, NamedTuple, Optional
|
from typing import TYPE_CHECKING, List, NamedTuple, Optional
|
||||||
|
|
||||||
@@ -46,6 +47,26 @@ logger = logging.getLogger(__name__)
|
|||||||
device_module = get_device_module()
|
device_module = get_device_module()
|
||||||
|
|
||||||
|
|
||||||
|
@cache
|
||||||
|
def _timing_events_supported() -> bool:
|
||||||
|
try:
|
||||||
|
device_module.Event(enable_timing=True)
|
||||||
|
return True
|
||||||
|
except (TypeError, NotImplementedError):
|
||||||
|
logger.warning(
|
||||||
|
"%s.Event does not support enable_timing=True; load-back "
|
||||||
|
"duration metric will be skipped on this backend.",
|
||||||
|
device_module.__name__,
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def make_timing_event_pair():
|
||||||
|
timing_enabled = _timing_events_supported()
|
||||||
|
kwargs = {"enable_timing": True} if timing_enabled else {}
|
||||||
|
return device_module.Event(**kwargs), device_module.Event(**kwargs), timing_enabled
|
||||||
|
|
||||||
|
|
||||||
class LayerLoadingEvent:
|
class LayerLoadingEvent:
|
||||||
def __init__(self, num_layers: int):
|
def __init__(self, num_layers: int):
|
||||||
self._num_layers = num_layers
|
self._num_layers = num_layers
|
||||||
@@ -140,6 +161,8 @@ class HiCacheAck(NamedTuple):
|
|||||||
start_event: device_module.Event
|
start_event: device_module.Event
|
||||||
finish_event: device_module.Event
|
finish_event: device_module.Event
|
||||||
node_ids: List[int]
|
node_ids: List[int]
|
||||||
|
num_tokens: int = 0
|
||||||
|
timing_enabled: bool = False
|
||||||
|
|
||||||
|
|
||||||
class StorageOperation:
|
class StorageOperation:
|
||||||
@@ -767,8 +790,11 @@ class HiCacheController:
|
|||||||
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()
|
||||||
|
|
||||||
|
ack_start_event, ack_finish_event, timing_enabled = make_timing_event_pair()
|
||||||
|
|
||||||
with device_module.stream(self.load_stream):
|
with device_module.stream(self.load_stream):
|
||||||
producer_event.start_event.wait(self.load_stream)
|
producer_event.start_event.wait(self.load_stream)
|
||||||
|
ack_start_event.record()
|
||||||
for i in range(self.layer_num):
|
for i in range(self.layer_num):
|
||||||
self.mem_pool_host.load_to_device_per_layer(
|
self.mem_pool_host.load_to_device_per_layer(
|
||||||
self.mem_pool_device,
|
self.mem_pool_device,
|
||||||
@@ -786,6 +812,7 @@ class HiCacheController:
|
|||||||
self.io_backend,
|
self.io_backend,
|
||||||
)
|
)
|
||||||
producer_event.complete(i)
|
producer_event.complete(i)
|
||||||
|
ack_finish_event.record()
|
||||||
# NOTE: We must save the host indices and device indices here,
|
# NOTE: We must save the host indices and device indices here,
|
||||||
# this is because we need to guarantee that these tensors are
|
# this is because we need to guarantee that these tensors are
|
||||||
# still alive when the load stream is executing.
|
# still alive when the load stream is executing.
|
||||||
@@ -796,9 +823,11 @@ class HiCacheController:
|
|||||||
|
|
||||||
self.ack_load_queue.append(
|
self.ack_load_queue.append(
|
||||||
HiCacheAck(
|
HiCacheAck(
|
||||||
start_event=producer_event.start_event,
|
start_event=ack_start_event,
|
||||||
finish_event=producer_event.finish_event,
|
finish_event=ack_finish_event,
|
||||||
node_ids=op.node_ids,
|
node_ids=op.node_ids,
|
||||||
|
num_tokens=len(op.device_indices),
|
||||||
|
timing_enabled=timing_enabled,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
return producer_id
|
return producer_id
|
||||||
|
|||||||
@@ -385,9 +385,9 @@ class HiMambaRadixCache(MambaRadixCache):
|
|||||||
if write_back:
|
if write_back:
|
||||||
# blocking till all write back complete
|
# blocking till all write back complete
|
||||||
while len(self.ongoing_write_through) > 0:
|
while len(self.ongoing_write_through) > 0:
|
||||||
for _, finish_event, ack_list in self.cache_controller.ack_write_queue:
|
for ack in self.cache_controller.ack_write_queue:
|
||||||
finish_event.synchronize()
|
ack.finish_event.synchronize()
|
||||||
for ack_id in ack_list:
|
for ack_id in ack.node_ids:
|
||||||
backuped_node = self.ongoing_write_through.pop(ack_id)
|
backuped_node = self.ongoing_write_through.pop(ack_id)
|
||||||
self._record_store_event(
|
self._record_store_event(
|
||||||
backuped_node, medium=StorageMedium.CPU
|
backuped_node, medium=StorageMedium.CPU
|
||||||
@@ -403,8 +403,8 @@ class HiMambaRadixCache(MambaRadixCache):
|
|||||||
# independently (no cross-rank sync).
|
# independently (no cross-rank sync).
|
||||||
finish_count = 0
|
finish_count = 0
|
||||||
if len(self.ongoing_write_through) > 0:
|
if len(self.ongoing_write_through) > 0:
|
||||||
for _, finish_event, ack_list in self.cache_controller.ack_write_queue:
|
for ack in self.cache_controller.ack_write_queue:
|
||||||
if not finish_event.query():
|
if not ack.finish_event.query():
|
||||||
break
|
break
|
||||||
finish_count += 1
|
finish_count += 1
|
||||||
|
|
||||||
@@ -418,9 +418,9 @@ class HiMambaRadixCache(MambaRadixCache):
|
|||||||
finish_count = int(queue_size.item())
|
finish_count = int(queue_size.item())
|
||||||
|
|
||||||
while finish_count > 0:
|
while finish_count > 0:
|
||||||
_, finish_event, ack_list = self.cache_controller.ack_write_queue.pop(0)
|
ack = self.cache_controller.ack_write_queue.pop(0)
|
||||||
finish_event.synchronize()
|
ack.finish_event.synchronize()
|
||||||
for ack_id in ack_list:
|
for ack_id in ack.node_ids:
|
||||||
backuped_node = self.ongoing_write_through.pop(ack_id)
|
backuped_node = self.ongoing_write_through.pop(ack_id)
|
||||||
self._record_store_event(backuped_node, medium=StorageMedium.CPU)
|
self._record_store_event(backuped_node, medium=StorageMedium.CPU)
|
||||||
self.dec_lock_ref(backuped_node)
|
self.dec_lock_ref(backuped_node)
|
||||||
@@ -432,8 +432,8 @@ class HiMambaRadixCache(MambaRadixCache):
|
|||||||
# Every rank must enter the all_reduce below; ongoing_load_back can
|
# Every rank must enter the all_reduce below; ongoing_load_back can
|
||||||
# diverge across ranks.
|
# diverge across ranks.
|
||||||
finish_count = 0
|
finish_count = 0
|
||||||
for _, finish_event, ack_list in self.cache_controller.ack_load_queue:
|
for ack in self.cache_controller.ack_load_queue:
|
||||||
if not finish_event.query():
|
if not ack.finish_event.query():
|
||||||
break
|
break
|
||||||
finish_count += 1
|
finish_count += 1
|
||||||
|
|
||||||
@@ -447,11 +447,19 @@ class HiMambaRadixCache(MambaRadixCache):
|
|||||||
finish_count = int(queue_size.item())
|
finish_count = int(queue_size.item())
|
||||||
|
|
||||||
while finish_count > 0:
|
while finish_count > 0:
|
||||||
_, finish_event, ack_list = self.cache_controller.ack_load_queue.pop(0)
|
ack = self.cache_controller.ack_load_queue.pop(0)
|
||||||
finish_event.synchronize()
|
ack.finish_event.synchronize()
|
||||||
for ack_id in ack_list:
|
for ack_id in ack.node_ids:
|
||||||
end_node = self.ongoing_load_back.pop(ack_id)
|
end_node = self.ongoing_load_back.pop(ack_id)
|
||||||
self.dec_lock_ref(end_node)
|
self.dec_lock_ref(end_node)
|
||||||
|
|
||||||
|
if self.metrics_collector is not None:
|
||||||
|
self.metrics_collector.increment_load_back_num_tokens(ack.num_tokens)
|
||||||
|
if ack.timing_enabled:
|
||||||
|
duration_ms = ack.start_event.elapsed_time(ack.finish_event)
|
||||||
|
self.metrics_collector.observe_load_back_duration(
|
||||||
|
duration_ms / 1000.0
|
||||||
|
)
|
||||||
finish_count -= 1
|
finish_count -= 1
|
||||||
|
|
||||||
def ready_to_load_host_cache(self) -> int:
|
def ready_to_load_host_cache(self) -> int:
|
||||||
|
|||||||
@@ -938,9 +938,9 @@ class HiRadixCache(RadixCache):
|
|||||||
if write_back:
|
if write_back:
|
||||||
# blocking till all write back complete
|
# blocking till all write back complete
|
||||||
while len(self.ongoing_write_through) > 0:
|
while len(self.ongoing_write_through) > 0:
|
||||||
for _, finish_event, ack_list in self.cache_controller.ack_write_queue:
|
for ack in self.cache_controller.ack_write_queue:
|
||||||
finish_event.synchronize()
|
ack.finish_event.synchronize()
|
||||||
for ack_id in ack_list:
|
for ack_id in ack.node_ids:
|
||||||
self._finish_write_through_ack(ack_id, release_lock=False)
|
self._finish_write_through_ack(ack_id, release_lock=False)
|
||||||
self.cache_controller.ack_write_queue.clear()
|
self.cache_controller.ack_write_queue.clear()
|
||||||
assert len(self.ongoing_write_through) == 0
|
assert len(self.ongoing_write_through) == 0
|
||||||
@@ -952,8 +952,8 @@ class HiRadixCache(RadixCache):
|
|||||||
# sequence and deadlocks under TP > 1. (Matches UnifiedRadixCache.)
|
# sequence and deadlocks under TP > 1. (Matches UnifiedRadixCache.)
|
||||||
finish_count = 0
|
finish_count = 0
|
||||||
if self.pp_rank == 0:
|
if self.pp_rank == 0:
|
||||||
for _, finish_event, ack_list in self.cache_controller.ack_write_queue:
|
for ack in self.cache_controller.ack_write_queue:
|
||||||
if not finish_event.query():
|
if not ack.finish_event.query():
|
||||||
break
|
break
|
||||||
finish_count += 1
|
finish_count += 1
|
||||||
finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu")
|
finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu")
|
||||||
@@ -963,17 +963,17 @@ class HiRadixCache(RadixCache):
|
|||||||
if finish_count > 0:
|
if finish_count > 0:
|
||||||
logger.debug(f"Process {finish_count} write back operations")
|
logger.debug(f"Process {finish_count} write back operations")
|
||||||
while finish_count > 0:
|
while finish_count > 0:
|
||||||
_, finish_event, ack_list = self.cache_controller.ack_write_queue.pop(0)
|
ack = self.cache_controller.ack_write_queue.pop(0)
|
||||||
finish_event.synchronize()
|
ack.finish_event.synchronize()
|
||||||
for ack_id in ack_list:
|
for ack_id in ack.node_ids:
|
||||||
self._finish_write_through_ack(ack_id, release_lock=True)
|
self._finish_write_through_ack(ack_id, release_lock=True)
|
||||||
finish_count -= 1
|
finish_count -= 1
|
||||||
|
|
||||||
def loading_check(self):
|
def loading_check(self):
|
||||||
finish_count = 0
|
finish_count = 0
|
||||||
if self.pp_rank == 0:
|
if self.pp_rank == 0:
|
||||||
for _, finish_event, ack_list in self.cache_controller.ack_load_queue:
|
for ack in self.cache_controller.ack_load_queue:
|
||||||
if not finish_event.query():
|
if not ack.finish_event.query():
|
||||||
break
|
break
|
||||||
finish_count += 1
|
finish_count += 1
|
||||||
finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu")
|
finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu")
|
||||||
@@ -983,11 +983,19 @@ class HiRadixCache(RadixCache):
|
|||||||
if finish_count > 0:
|
if finish_count > 0:
|
||||||
logger.debug(f"Process {finish_count} load operations")
|
logger.debug(f"Process {finish_count} load operations")
|
||||||
while finish_count > 0:
|
while finish_count > 0:
|
||||||
_, finish_event, ack_list = self.cache_controller.ack_load_queue.pop(0)
|
ack = self.cache_controller.ack_load_queue.pop(0)
|
||||||
finish_event.synchronize()
|
ack.finish_event.synchronize()
|
||||||
for ack_id in ack_list:
|
for ack_id in ack.node_ids:
|
||||||
end_node = self.ongoing_load_back.pop(ack_id)
|
end_node = self.ongoing_load_back.pop(ack_id)
|
||||||
self.dec_lock_ref(end_node)
|
self.dec_lock_ref(end_node)
|
||||||
|
|
||||||
|
if self.metrics_collector is not None:
|
||||||
|
self.metrics_collector.increment_load_back_num_tokens(ack.num_tokens)
|
||||||
|
if ack.timing_enabled:
|
||||||
|
duration_ms = ack.start_event.elapsed_time(ack.finish_event)
|
||||||
|
self.metrics_collector.observe_load_back_duration(
|
||||||
|
duration_ms / 1000.0
|
||||||
|
)
|
||||||
finish_count -= 1
|
finish_count -= 1
|
||||||
|
|
||||||
def is_load_back_event_done(self, consumer_index: int) -> bool:
|
def is_load_back_event_done(self, consumer_index: int) -> bool:
|
||||||
@@ -1243,7 +1251,6 @@ class HiRadixCache(RadixCache):
|
|||||||
self, node: TreeNode, mem_quota: Optional[int] = None
|
self, node: TreeNode, mem_quota: Optional[int] = None
|
||||||
) -> Optional[torch.Tensor]:
|
) -> Optional[torch.Tensor]:
|
||||||
|
|
||||||
start_time = time.perf_counter()
|
|
||||||
last_hit_node = node
|
last_hit_node = node
|
||||||
nodes_to_load = []
|
nodes_to_load = []
|
||||||
while node.evicted:
|
while node.evicted:
|
||||||
@@ -1311,12 +1318,6 @@ class HiRadixCache(RadixCache):
|
|||||||
self.evictable_size_ += len(device_indices)
|
self.evictable_size_ += len(device_indices)
|
||||||
self.inc_lock_ref(last_hit_node)
|
self.inc_lock_ref(last_hit_node)
|
||||||
|
|
||||||
if self.metrics_collector is not None:
|
|
||||||
self.metrics_collector.observe_load_back_duration(
|
|
||||||
time.perf_counter() - start_time
|
|
||||||
)
|
|
||||||
self.metrics_collector.increment_load_back_num_tokens(len(device_indices))
|
|
||||||
|
|
||||||
return device_indices
|
return device_indices
|
||||||
|
|
||||||
def init_load_back(
|
def init_load_back(
|
||||||
|
|||||||
@@ -23,6 +23,9 @@ from sglang.srt.managers.cache_controller import (
|
|||||||
from sglang.srt.managers.cache_controller import (
|
from sglang.srt.managers.cache_controller import (
|
||||||
StorageOperation as BaseStorageOperation,
|
StorageOperation as BaseStorageOperation,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.managers.cache_controller import (
|
||||||
|
make_timing_event_pair,
|
||||||
|
)
|
||||||
from sglang.srt.mem_cache.hicache_storage import (
|
from sglang.srt.mem_cache.hicache_storage import (
|
||||||
HiCacheStorageExtraInfo,
|
HiCacheStorageExtraInfo,
|
||||||
PoolHitPolicy,
|
PoolHitPolicy,
|
||||||
@@ -490,8 +493,12 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
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()
|
||||||
|
|
||||||
|
ack_start_event, ack_finish_event, timing_enabled = make_timing_event_pair()
|
||||||
|
|
||||||
with device_module.stream(self.load_stream):
|
with device_module.stream(self.load_stream):
|
||||||
producer_event.start_event.wait(self.load_stream)
|
producer_event.start_event.wait(self.load_stream)
|
||||||
|
ack_start_event.record()
|
||||||
for i in range(self.layer_num):
|
for i in range(self.layer_num):
|
||||||
self.mem_pool_host.load_to_device_per_layer(
|
self.mem_pool_host.load_to_device_per_layer(
|
||||||
self.mem_pool_device,
|
self.mem_pool_device,
|
||||||
@@ -514,6 +521,7 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
self.io_backend,
|
self.io_backend,
|
||||||
)
|
)
|
||||||
producer_event.complete(i)
|
producer_event.complete(i)
|
||||||
|
ack_finish_event.record()
|
||||||
self._record_transfer_indices_on_stream(
|
self._record_transfer_indices_on_stream(
|
||||||
self.load_stream,
|
self.load_stream,
|
||||||
host_indices,
|
host_indices,
|
||||||
@@ -522,9 +530,11 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
)
|
)
|
||||||
self.ack_load_queue.append(
|
self.ack_load_queue.append(
|
||||||
HiCacheAck(
|
HiCacheAck(
|
||||||
producer_event.start_event,
|
ack_start_event,
|
||||||
producer_event.finish_event,
|
ack_finish_event,
|
||||||
op.node_ids,
|
op.node_ids,
|
||||||
|
num_tokens=len(op.device_indices),
|
||||||
|
timing_enabled=timing_enabled,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
return producer_id
|
return producer_id
|
||||||
|
|||||||
@@ -1668,7 +1668,6 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
|
|||||||
if self.cache_controller is None:
|
if self.cache_controller is None:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
start_time = time.perf_counter()
|
|
||||||
host_anchor_params = self.inc_host_lock_ref(best_match_node).to_dec_params()
|
host_anchor_params = self.inc_host_lock_ref(best_match_node).to_dec_params()
|
||||||
# Build KV transfer
|
# Build KV transfer
|
||||||
kv_xfer = self.components[BASE_COMPONENT_TYPE].build_hicache_transfers(
|
kv_xfer = self.components[BASE_COMPONENT_TYPE].build_hicache_transfers(
|
||||||
@@ -1753,12 +1752,6 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
|
|||||||
host_anchor_params,
|
host_anchor_params,
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.metrics_collector is not None:
|
|
||||||
self.metrics_collector.observe_load_back_duration(
|
|
||||||
time.perf_counter() - start_time
|
|
||||||
)
|
|
||||||
self.metrics_collector.increment_load_back_num_tokens(len(device_indices))
|
|
||||||
|
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def _build_sidecar_transfers(
|
def _build_sidecar_transfers(
|
||||||
@@ -2350,9 +2343,9 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
|
|||||||
if write_back:
|
if write_back:
|
||||||
# Blocking: wait for all pending write-backs
|
# Blocking: wait for all pending write-backs
|
||||||
while self.ongoing_write_through:
|
while self.ongoing_write_through:
|
||||||
for _, finish_event, ack_list in cc.ack_write_queue:
|
for ack in cc.ack_write_queue:
|
||||||
finish_event.synchronize()
|
ack.finish_event.synchronize()
|
||||||
for ack_id in ack_list:
|
for ack_id in ack.node_ids:
|
||||||
if ack_id in self.ongoing_write_through:
|
if ack_id in self.ongoing_write_through:
|
||||||
self._finish_write_through_ack(ack_id)
|
self._finish_write_through_ack(ack_id)
|
||||||
cc.ack_write_queue.clear()
|
cc.ack_write_queue.clear()
|
||||||
@@ -2363,8 +2356,8 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
|
|||||||
# diverge across ranks (e.g. write_backup returning 0 on a subset).
|
# diverge across ranks (e.g. write_backup returning 0 on a subset).
|
||||||
finish_count = 0
|
finish_count = 0
|
||||||
if self.pp_rank == 0:
|
if self.pp_rank == 0:
|
||||||
for _, finish_event, ack_list in cc.ack_write_queue:
|
for ack in cc.ack_write_queue:
|
||||||
if not finish_event.query():
|
if not ack.finish_event.query():
|
||||||
break
|
break
|
||||||
finish_count += 1
|
finish_count += 1
|
||||||
|
|
||||||
@@ -2374,9 +2367,9 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
|
|||||||
|
|
||||||
# Process completed acks
|
# Process completed acks
|
||||||
while finish_count > 0:
|
while finish_count > 0:
|
||||||
_, finish_event, ack_list = cc.ack_write_queue.pop(0)
|
ack = cc.ack_write_queue.pop(0)
|
||||||
finish_event.synchronize()
|
ack.finish_event.synchronize()
|
||||||
for ack_id in ack_list:
|
for ack_id in ack.node_ids:
|
||||||
self._finish_write_through_ack(ack_id)
|
self._finish_write_through_ack(ack_id)
|
||||||
finish_count -= 1
|
finish_count -= 1
|
||||||
|
|
||||||
@@ -2389,8 +2382,8 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
|
|||||||
# diverge across ranks.
|
# diverge across ranks.
|
||||||
finish_count = 0
|
finish_count = 0
|
||||||
if self.pp_rank == 0:
|
if self.pp_rank == 0:
|
||||||
for _, finish_event, ack_list in cc.ack_load_queue:
|
for ack in cc.ack_load_queue:
|
||||||
if not finish_event.query():
|
if not ack.finish_event.query():
|
||||||
break
|
break
|
||||||
finish_count += 1
|
finish_count += 1
|
||||||
finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu")
|
finish_count_tensor = torch.tensor(finish_count, dtype=torch.int, device="cpu")
|
||||||
@@ -2398,12 +2391,20 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
|
|||||||
finish_count = finish_count_tensor.item()
|
finish_count = finish_count_tensor.item()
|
||||||
|
|
||||||
while finish_count > 0:
|
while finish_count > 0:
|
||||||
_, finish_event, ack_list = cc.ack_load_queue.pop(0)
|
ack = cc.ack_load_queue.pop(0)
|
||||||
finish_event.synchronize()
|
ack.finish_event.synchronize()
|
||||||
for ack_id in ack_list:
|
for ack_id in ack.node_ids:
|
||||||
node, lock_params, host_lock_params = self.ongoing_load_back.pop(ack_id)
|
node, lock_params, host_lock_params = self.ongoing_load_back.pop(ack_id)
|
||||||
self.dec_lock_ref(node, lock_params)
|
self.dec_lock_ref(node, lock_params)
|
||||||
self.dec_host_lock_ref(node, host_lock_params)
|
self.dec_host_lock_ref(node, host_lock_params)
|
||||||
|
|
||||||
|
if self.metrics_collector is not None:
|
||||||
|
self.metrics_collector.increment_load_back_num_tokens(ack.num_tokens)
|
||||||
|
if ack.timing_enabled:
|
||||||
|
duration_ms = ack.start_event.elapsed_time(ack.finish_event)
|
||||||
|
self.metrics_collector.observe_load_back_duration(
|
||||||
|
duration_ms / 1000.0
|
||||||
|
)
|
||||||
finish_count -= 1
|
finish_count -= 1
|
||||||
|
|
||||||
# ---- HiCache: Scheduler Entry Points ----
|
# ---- HiCache: Scheduler Entry Points ----
|
||||||
|
|||||||
@@ -1934,7 +1934,7 @@ class RadixCacheMetricsCollector(_StatLoggerDIMixin):
|
|||||||
|
|
||||||
self.load_back_duration_seconds = Histogram(
|
self.load_back_duration_seconds = Histogram(
|
||||||
name="sglang:load_back_duration_seconds",
|
name="sglang:load_back_duration_seconds",
|
||||||
documentation="Time taken to load memory from CPU to GPU in seconds.",
|
documentation="Time taken to load KV cache from CPU back to GPU in seconds.",
|
||||||
labelnames=labels.keys(),
|
labelnames=labels.keys(),
|
||||||
buckets=bucket_load_back_duration,
|
buckets=bucket_load_back_duration,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ from sglang.srt.disaggregation.decode_kvcache_offload_manager import (
|
|||||||
DecodeKVCacheOffloadManager,
|
DecodeKVCacheOffloadManager,
|
||||||
)
|
)
|
||||||
from sglang.srt.disaggregation.kv_events import OffloadedState
|
from sglang.srt.disaggregation.kv_events import OffloadedState
|
||||||
|
from sglang.srt.managers.cache_controller import HiCacheAck
|
||||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
|
|
||||||
register_cuda_ci(est_time=8, stage="base-b", runner_config="1-gpu-small")
|
register_cuda_ci(est_time=8, stage="base-b", runner_config="1-gpu-small")
|
||||||
@@ -280,7 +281,9 @@ class TestReleaseFinishedReq(unittest.TestCase):
|
|||||||
8,
|
8,
|
||||||
)
|
)
|
||||||
manager.cache_controller = MagicMock()
|
manager.cache_controller = MagicMock()
|
||||||
manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [7])]
|
manager.cache_controller.ack_write_queue = [
|
||||||
|
HiCacheAck(None, _FinishedEvent(), [7])
|
||||||
|
]
|
||||||
manager._trigger_backup = MagicMock(return_value="last_hash")
|
manager._trigger_backup = MagicMock(return_value="last_hash")
|
||||||
|
|
||||||
manager._check_offload_progress(1)
|
manager._check_offload_progress(1)
|
||||||
@@ -314,7 +317,9 @@ class TestReleaseFinishedReq(unittest.TestCase):
|
|||||||
self.assertEqual(manager.offloaded_state[req.rid].inc_len, 4)
|
self.assertEqual(manager.offloaded_state[req.rid].inc_len, 4)
|
||||||
manager.cache_controller.write.assert_called_once()
|
manager.cache_controller.write.assert_called_once()
|
||||||
|
|
||||||
manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [1])]
|
manager.cache_controller.ack_write_queue = [
|
||||||
|
HiCacheAck(None, _FinishedEvent(), [1])
|
||||||
|
]
|
||||||
manager._trigger_backup = MagicMock(return_value="last_hash")
|
manager._trigger_backup = MagicMock(return_value="last_hash")
|
||||||
|
|
||||||
manager._check_offload_progress(1)
|
manager._check_offload_progress(1)
|
||||||
@@ -357,7 +362,9 @@ class TestReleaseFinishedReq(unittest.TestCase):
|
|||||||
8,
|
8,
|
||||||
)
|
)
|
||||||
manager.cache_controller = MagicMock()
|
manager.cache_controller = MagicMock()
|
||||||
manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [8])]
|
manager.cache_controller.ack_write_queue = [
|
||||||
|
HiCacheAck(None, _FinishedEvent(), [8])
|
||||||
|
]
|
||||||
manager._trigger_backup = MagicMock(return_value="last_hash")
|
manager._trigger_backup = MagicMock(return_value="last_hash")
|
||||||
|
|
||||||
manager._check_offload_progress(1)
|
manager._check_offload_progress(1)
|
||||||
@@ -387,7 +394,9 @@ class TestReleaseFinishedReq(unittest.TestCase):
|
|||||||
12,
|
12,
|
||||||
)
|
)
|
||||||
manager.cache_controller = MagicMock()
|
manager.cache_controller = MagicMock()
|
||||||
manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [9])]
|
manager.cache_controller.ack_write_queue = [
|
||||||
|
HiCacheAck(None, _FinishedEvent(), [9])
|
||||||
|
]
|
||||||
manager._trigger_backup = MagicMock(return_value="last_hash")
|
manager._trigger_backup = MagicMock(return_value="last_hash")
|
||||||
|
|
||||||
manager._check_offload_progress(1)
|
manager._check_offload_progress(1)
|
||||||
|
|||||||
@@ -0,0 +1,124 @@
|
|||||||
|
"""Unit tests for the HiCache load-back duration metric."""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=5, stage="base-b", runner_config="1-gpu-small")
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(torch.cuda.is_available(), "CUDA required")
|
||||||
|
class TestLoadBackDurationMetric(CustomTestCase):
|
||||||
|
def setUp(self):
|
||||||
|
from sglang.srt.managers import cache_controller as cc
|
||||||
|
|
||||||
|
cc._timing_events_supported.cache_clear()
|
||||||
|
self.cc = cc
|
||||||
|
|
||||||
|
def _completed_pair(self, payload_floats=1024 * 1024):
|
||||||
|
start, finish, timing_enabled = self.cc.make_timing_event_pair()
|
||||||
|
self.assertTrue(timing_enabled)
|
||||||
|
stream = torch.cuda.Stream()
|
||||||
|
start.record()
|
||||||
|
with torch.cuda.stream(stream):
|
||||||
|
start.wait(stream)
|
||||||
|
torch.empty(payload_floats, device="cuda").fill_(0)
|
||||||
|
finish.record()
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
return start, finish
|
||||||
|
|
||||||
|
def test_elapsed_time_works(self):
|
||||||
|
start, finish = self._completed_pair()
|
||||||
|
self.assertGreater(start.elapsed_time(finish), 0.0)
|
||||||
|
|
||||||
|
def test_timing_fallback_uses_dedicated_events(self):
|
||||||
|
events = []
|
||||||
|
|
||||||
|
def create_event(*, enable_timing=False):
|
||||||
|
if enable_timing:
|
||||||
|
raise TypeError
|
||||||
|
event = MagicMock()
|
||||||
|
events.append(event)
|
||||||
|
return event
|
||||||
|
|
||||||
|
with patch.object(self.cc.device_module, "Event", side_effect=create_event):
|
||||||
|
self.cc._timing_events_supported.cache_clear()
|
||||||
|
start, finish, timing_enabled = self.cc.make_timing_event_pair()
|
||||||
|
|
||||||
|
self.assertFalse(timing_enabled)
|
||||||
|
self.assertIs(start, events[0])
|
||||||
|
self.assertIs(finish, events[1])
|
||||||
|
self.assertIsNot(start, finish)
|
||||||
|
|
||||||
|
def test_loading_check_observes_duration_and_tokens(self):
|
||||||
|
from sglang.srt.mem_cache.hiradix_cache import HiRadixCache
|
||||||
|
|
||||||
|
start, finish = self._completed_pair()
|
||||||
|
ack = self.cc.HiCacheAck(
|
||||||
|
start,
|
||||||
|
finish,
|
||||||
|
node_ids=[1, 2],
|
||||||
|
num_tokens=1024,
|
||||||
|
timing_enabled=True,
|
||||||
|
)
|
||||||
|
stub = SimpleNamespace(
|
||||||
|
cache_controller=SimpleNamespace(ack_load_queue=[ack]),
|
||||||
|
ongoing_load_back={1: object(), 2: object()},
|
||||||
|
dec_lock_ref=MagicMock(),
|
||||||
|
metrics_collector=MagicMock(),
|
||||||
|
pp_rank=0,
|
||||||
|
_all_reduce=MagicMock(),
|
||||||
|
)
|
||||||
|
|
||||||
|
HiRadixCache.loading_check(stub)
|
||||||
|
|
||||||
|
stub.metrics_collector.increment_load_back_num_tokens.assert_called_once_with(
|
||||||
|
1024
|
||||||
|
)
|
||||||
|
stub.metrics_collector.observe_load_back_duration.assert_called_once()
|
||||||
|
(observed,), _ = stub.metrics_collector.observe_load_back_duration.call_args
|
||||||
|
self.assertGreater(observed, 0.0)
|
||||||
|
self.assertEqual(stub.cache_controller.ack_load_queue, [])
|
||||||
|
|
||||||
|
def test_loading_check_fallback_when_timing_unsupported(self):
|
||||||
|
"""On backends without enable_timing, count tokens but skip duration."""
|
||||||
|
from sglang.srt.mem_cache.hiradix_cache import HiRadixCache
|
||||||
|
|
||||||
|
start = torch.cuda.Event()
|
||||||
|
finish = torch.cuda.Event()
|
||||||
|
start.record()
|
||||||
|
finish.record()
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
|
ack = self.cc.HiCacheAck(
|
||||||
|
start_event=start,
|
||||||
|
finish_event=finish,
|
||||||
|
node_ids=[7],
|
||||||
|
num_tokens=512,
|
||||||
|
timing_enabled=False,
|
||||||
|
)
|
||||||
|
stub = SimpleNamespace(
|
||||||
|
cache_controller=SimpleNamespace(ack_load_queue=[ack]),
|
||||||
|
ongoing_load_back={7: object()},
|
||||||
|
dec_lock_ref=MagicMock(),
|
||||||
|
metrics_collector=MagicMock(),
|
||||||
|
pp_rank=0,
|
||||||
|
_all_reduce=MagicMock(),
|
||||||
|
)
|
||||||
|
|
||||||
|
HiRadixCache.loading_check(stub)
|
||||||
|
|
||||||
|
stub.metrics_collector.increment_load_back_num_tokens.assert_called_once_with(
|
||||||
|
512
|
||||||
|
)
|
||||||
|
stub.metrics_collector.observe_load_back_duration.assert_not_called()
|
||||||
|
self.assertEqual(stub.cache_controller.ack_load_queue, [])
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -499,8 +499,8 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase):
|
|||||||
self.assertTrue(loaded)
|
self.assertTrue(loaded)
|
||||||
producer_id = cache.ready_to_load_host_cache()
|
producer_id = cache.ready_to_load_host_cache()
|
||||||
self.assertNotEqual(producer_id, -1)
|
self.assertNotEqual(producer_id, -1)
|
||||||
for _, finish_event, _ in list(cache.cache_controller.ack_load_queue):
|
for ack in list(cache.cache_controller.ack_load_queue):
|
||||||
finish_event.synchronize()
|
ack.finish_event.synchronize()
|
||||||
cache.loading_check()
|
cache.loading_check()
|
||||||
|
|
||||||
def test_kv_events_store_and_remove_full_blocks(self):
|
def test_kv_events_store_and_remove_full_blocks(self):
|
||||||
@@ -2673,8 +2673,8 @@ class UnifiedRadixCacheSuite:
|
|||||||
self.assertTrue(loaded)
|
self.assertTrue(loaded)
|
||||||
producer_id = cache.ready_to_load_host_cache()
|
producer_id = cache.ready_to_load_host_cache()
|
||||||
self.assertNotEqual(producer_id, -1)
|
self.assertNotEqual(producer_id, -1)
|
||||||
for _, finish_event, _ in list(cache.cache_controller.ack_load_queue):
|
for ack in list(cache.cache_controller.ack_load_queue):
|
||||||
finish_event.synchronize()
|
ack.finish_event.synchronize()
|
||||||
cache.loading_check()
|
cache.loading_check()
|
||||||
return node.component_data[ComponentType.FULL].value
|
return node.component_data[ComponentType.FULL].value
|
||||||
|
|
||||||
@@ -3127,8 +3127,8 @@ class UnifiedRadixCacheSuite:
|
|||||||
def _finish_pending_loads(self, cache):
|
def _finish_pending_loads(self, cache):
|
||||||
producer_id = cache.ready_to_load_host_cache()
|
producer_id = cache.ready_to_load_host_cache()
|
||||||
self.assertNotEqual(producer_id, -1)
|
self.assertNotEqual(producer_id, -1)
|
||||||
for _, finish_event, _ in list(cache.cache_controller.ack_load_queue):
|
for ack in list(cache.cache_controller.ack_load_queue):
|
||||||
finish_event.synchronize()
|
ack.finish_event.synchronize()
|
||||||
cache.loading_check()
|
cache.loading_check()
|
||||||
|
|
||||||
def _match_tokens_for_chain(self, chain):
|
def _match_tokens_for_chain(self, chain):
|
||||||
|
|||||||
Reference in New Issue
Block a user