From 8922bb98e28d3c576108f8629fa4dcb74cf9aa9f Mon Sep 17 00:00:00 2001 From: cctry Date: Sun, 16 Aug 2026 00:33:34 -0700 Subject: [PATCH] refactor(hicache): flatten L2 transfer execution (#34793) GB300 test fails unrelated --- .../decode_kvcache_offload_manager.py | 34 +- .../sglang/srt/managers/cache_controller.py | 227 +++++------ .../sglang/srt/mem_cache/base_prefix_cache.py | 2 + .../hybrid_cache/hybrid_cache_controller.py | 361 +++++++----------- .../hybrid_cache/hybrid_pool_assembler.py | 198 ++++++---- python/sglang/srt/mem_cache/l2_transfer.py | 127 ++++++ .../sglang/srt/mem_cache/memory_pool_host.py | 114 ------ .../test_hicache_load_back_timing.py | 14 +- ...test_hicache_staged_write_back_dispatch.py | 337 +++++++++++----- ...test_supplied_instance_exposure_ratchet.py | 1 - 10 files changed, 767 insertions(+), 648 deletions(-) create mode 100644 python/sglang/srt/mem_cache/l2_transfer.py diff --git a/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py b/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py index 600877617..9787e0231 100644 --- a/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py +++ b/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py @@ -13,14 +13,14 @@ from sglang.srt.environ import envs from sglang.srt.managers.cache_controller import HiCacheController from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache +from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import ( + build_kv_host_pool, +) from sglang.srt.mem_cache.memory_pool import ( MHATokenToKVPool, MLATokenToKVPool, ReqToTokenPool, ) -from sglang.srt.mem_cache.pool_host.common import get_allocator_type -from sglang.srt.mem_cache.pool_host.mha import get_mha_host_pool_cls -from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost from sglang.srt.runtime_context import get_schedule from sglang.srt.server_args import ServerArgs from sglang.srt.utils.common import ceil_align @@ -55,28 +55,14 @@ class DecodeKVCacheOffloadManager: self.page_size, (env_stride // self.page_size) * self.page_size ) kv_cache = self.token_to_kv_pool_allocator.get_kvcache() - allocator_type = get_allocator_type(server_args) - - if isinstance(kv_cache, MHATokenToKVPool): - self.decode_host_mem_pool = get_mha_host_pool_cls(kv_cache)( - kv_cache, - server_args.hicache_ratio, - server_args.hicache_size, - self.page_size, - server_args.hicache_mem_layout, - allocator_type=allocator_type, - ) - elif isinstance(kv_cache, MLATokenToKVPool): - self.decode_host_mem_pool = MLATokenToKVPoolHost( - kv_cache, - server_args.hicache_ratio, - server_args.hicache_size, - self.page_size, - server_args.hicache_mem_layout, - allocator_type=allocator_type, - ) - else: + if not isinstance(kv_cache, (MHATokenToKVPool, MLATokenToKVPool)): raise ValueError("Unsupported KV cache type for decode offload") + self.decode_host_mem_pool = build_kv_host_pool( + kv_pool=kv_cache, + page_size=self.page_size, + server_args=server_args, + use_mla=isinstance(kv_cache, MLATokenToKVPool), + ) self.tp_group = tp_group self.tp_world_size = torch.distributed.get_world_size(group=self.tp_group) diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index 29ddc12b9..6a1089254 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -16,7 +16,6 @@ limitations under the License. import logging import threading import time -from functools import cache from queue import Empty, Queue from typing import TYPE_CHECKING, List, NamedTuple, Optional @@ -38,6 +37,7 @@ from sglang.srt.layers.dp_attention import ( get_attention_dp_rank, is_dp_attention_enabled, ) +from sglang.srt.mem_cache.l2_transfer import L2Transfer, L2TransferEngine from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import get_device_module @@ -47,26 +47,6 @@ logger = logging.getLogger(__name__) 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: def __init__(self, num_layers: int): self._num_layers = num_layers @@ -126,30 +106,66 @@ class CacheOperation: device_indices: torch.Tensor, node_id: int, priority: Optional[int] = None, + pool_transfers: Optional[List[PoolTransfer]] = None, ): self.host_indices = host_indices self.device_indices = device_indices self.node_ids = [node_id] self.data = None + self.pool_transfers = pool_transfers self.id = CacheOperation.counter CacheOperation.counter += 1 # default priority is the order of creation self.priority = priority if priority is not None else self.id + @staticmethod + def _merge_pool_transfers( + ops: List[CacheOperation], + ) -> Optional[List[PoolTransfer]]: + grouped: dict[tuple[PoolName, Optional[PoolName]], List[PoolTransfer]] = {} + for op in ops: + for transfer in op.pool_transfers or []: + grouped.setdefault( + (transfer.name, transfer.indices_from_pool), [] + ).append(transfer) + if not grouped: + return None + + def cat_or_none(tensors): + parts = [tensor for tensor in tensors if tensor is not None] + return torch.cat(parts) if parts else None + + return [ + PoolTransfer( + name=transfers[0].name, + host_indices=cat_or_none(t.host_indices for t in transfers), + device_indices=cat_or_none(t.device_indices for t in transfers), + keys=[key for t in transfers if t.keys for key in t.keys] or None, + hit_policy=transfers[0].hit_policy, + indices_from_pool=transfers[0].indices_from_pool, + ) + for transfers in grouped.values() + ] + @staticmethod def merge_ops(ops: List[CacheOperation]) -> CacheOperation: - assert len(ops) > 0 + assert ops if len(ops) == 1: return ops[0] - host_indices = torch.cat([op.host_indices for op in ops]) device_indices = torch.cat([op.device_indices for op in ops]) node_ids = [] priority = min(op.priority for op in ops) for op in ops: node_ids.extend(op.node_ids) - merged_op = CacheOperation(host_indices, device_indices, -1, priority) + merged_op = CacheOperation( + host_indices, + device_indices, + -1, + priority, + pool_transfers=CacheOperation._merge_pool_transfers(ops), + ) merged_op.node_ids = node_ids return merged_op @@ -302,8 +318,7 @@ class HiCacheController: self.ack_load_queue: List[HiCacheAck] = [] self.ack_write_queue: List[HiCacheAck] = [] - self.write_stream = device_module.Stream() - self.load_stream = device_module.Stream() + self.l2_transfer_engine = L2TransferEngine(io_backend) # If a storage backend is provided at startup, treat it as an implicit attach, # so init/runtime share the same lifecycle semantics and code paths. @@ -698,60 +713,21 @@ class HiCacheController: return op = CacheOperation.merge_ops(self.write_queue) - # Kernel write-back keeps host indices on CPU only for page_first AND only - # when the staged JIT write-back kernel is available (it stages through - # device memory and accepts CPU destination indices). Otherwise we fall back - # to the plain transfer kernel, whose CUDA/HIP implementation requires - # device-resident destination indices -- so the indices must be moved to the - # device first. Without the can_use_write_back_jit check this crashes on - # backends where the JIT kernel is unavailable, with - # "Destination indices must be a CUDA tensor". - if ( - self.io_backend == "kernel" - and self.mem_pool_host.layout == "page_first" - and getattr(self.mem_pool_host, "can_use_write_back_jit", False) - ): - host_indices, device_indices = op.host_indices, op.device_indices - else: - host_indices, device_indices = self.move_indices( - op.host_indices, op.device_indices - ) + host_indices, device_indices, pool_transfers = self._move_write_operation(op) self.write_queue.clear() - start_event = device_module.Event() - ack_start_event, ack_finish_event, timing_enabled = make_timing_event_pair() - - start_event.record() - with device_module.stream(self.write_stream): - start_event.wait(self.write_stream) - ack_start_event.record() - self.mem_pool_host.backup_from_device_all_layer( - self.mem_pool_device, host_indices, device_indices, self.io_backend - ) - if self.has_draft: - self.mem_pool_host_draft.backup_from_device_all_layer( - self.mem_pool_device_draft, - host_indices, - device_indices, - self.io_backend, - ) - ack_finish_event.record() - # NOTE: We must save the host indices and device indices here, - # this is because we need to guarantee that these tensors are - # still alive when the write stream is executing. - if host_indices.is_cuda: - host_indices.record_stream(self.write_stream) - if device_indices.is_cuda: - device_indices.record_stream(self.write_stream) + completion = self.l2_transfer_engine.submit_device_to_host( + self._l2_transfers(host_indices, device_indices, pool_transfers) + ) self.ack_write_queue.append( HiCacheAck( - start_event=ack_start_event, - finish_event=ack_finish_event, + start_event=completion.start_event, + finish_event=completion.finish_event, node_ids=op.node_ids, num_tokens=len(op.device_indices), - timing_enabled=timing_enabled, - num_tokens_by_pool={PoolName.KV.value: len(op.device_indices)}, + timing_enabled=completion.timing_enabled, + num_tokens_by_pool=self._num_tokens_by_pool(op), num_bytes=self._transfer_num_bytes(op), ) ) @@ -764,6 +740,9 @@ class HiCacheController: num_bytes += num_tokens * self.mem_pool_host_draft.size_per_token return num_bytes + def _num_tokens_by_pool(self, op: CacheOperation) -> dict[str, int]: + return {PoolName.KV.value: len(op.device_indices)} + def load( self, host_indices: torch.Tensor, @@ -781,7 +760,9 @@ class HiCacheController: ) return device_indices - def move_indices(self, host_indices: torch.Tensor, device_indices: torch.Tensor): + def move_indices( + self, host_indices: torch.Tensor, device_indices: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor]: # move indices to GPU if using kernels, to host if using direct indexing if self.io_backend == "kernel": if not host_indices.is_cuda: @@ -803,58 +784,82 @@ class HiCacheController: else: raise ValueError(f"Unsupported io backend") + def _move_write_operation( + self, op: CacheOperation + ) -> tuple[torch.Tensor, torch.Tensor, Optional[List[PoolTransfer]]]: + """Keep CPU host indices only for page-first staged write-back.""" + if ( + self.io_backend == "kernel" + and self.mem_pool_host.layout == "page_first" + and getattr(self.mem_pool_host, "can_use_write_back_jit", False) + ): + return op.host_indices, op.device_indices, op.pool_transfers + return self._move_op_indices(op) + + def _move_op_indices( + self, op: CacheOperation + ) -> tuple[torch.Tensor, torch.Tensor, Optional[List[PoolTransfer]]]: + return (*self.move_indices(op.host_indices, op.device_indices), None) + + def _l2_transfers( + self, + host_indices: torch.Tensor, + device_indices: torch.Tensor, + pool_transfers: Optional[List[PoolTransfer]] = None, + ) -> list[L2Transfer]: + transfers = [ + L2Transfer( + host_pool=self.mem_pool_host, + device_pool=self.mem_pool_device, + host_indices=host_indices, + device_indices=device_indices, + ) + ] + if self.has_draft and host_indices.numel() > 0: + transfers.append( + L2Transfer( + host_pool=self.mem_pool_host_draft, + device_pool=self.mem_pool_device_draft, + host_indices=host_indices, + device_indices=device_indices, + ) + ) + return transfers + + def _l2_load_transfers( + self, + host_indices: torch.Tensor, + device_indices: torch.Tensor, + pool_transfers: Optional[List[PoolTransfer]] = None, + ) -> list[L2Transfer]: + return self._l2_transfers(host_indices, device_indices, pool_transfers) + def start_loading(self) -> int: if len(self.load_queue) == 0: return -1 producer_id = self.layer_done_counter.update_producer() op = CacheOperation.merge_ops(self.load_queue) - host_indices, device_indices = self.move_indices( - op.host_indices, op.device_indices - ) + host_indices, device_indices, pool_transfers = self._move_op_indices(op) self.load_queue.clear() producer_event = self.layer_done_counter.events[producer_id] producer_event.start_event.record() - ack_start_event, ack_finish_event, timing_enabled = make_timing_event_pair() - - with device_module.stream(self.load_stream): - producer_event.start_event.wait(self.load_stream) - ack_start_event.record() - for i in range(self.layer_num): - self.mem_pool_host.load_to_device_per_layer( - self.mem_pool_device, - host_indices, - device_indices, - i, - self.io_backend, - ) - if self.has_draft and i < self.mem_pool_host_draft.layer_num: - self.mem_pool_host_draft.load_to_device_per_layer( - self.mem_pool_device_draft, - host_indices, - device_indices, - i, - self.io_backend, - ) - producer_event.complete(i) - ack_finish_event.record() - # NOTE: We must save the host indices and device indices here, - # this is because we need to guarantee that these tensors are - # still alive when the load stream is executing. - if host_indices.is_cuda: - host_indices.record_stream(self.load_stream) - if device_indices.is_cuda: - device_indices.record_stream(self.load_stream) + completion = self.l2_transfer_engine.submit_host_to_device( + self._l2_load_transfers(host_indices, device_indices, pool_transfers), + start_event=producer_event.start_event, + on_layer_done=producer_event.complete, + layer_num=self.layer_num, + ) self.ack_load_queue.append( HiCacheAck( - start_event=ack_start_event, - finish_event=ack_finish_event, + start_event=completion.start_event, + finish_event=completion.finish_event, node_ids=op.node_ids, num_tokens=len(op.device_indices), - timing_enabled=timing_enabled, - num_tokens_by_pool={PoolName.KV.value: len(op.device_indices)}, + timing_enabled=completion.timing_enabled, + num_tokens_by_pool=self._num_tokens_by_pool(op), num_bytes=self._transfer_num_bytes(op), ) ) diff --git a/python/sglang/srt/mem_cache/base_prefix_cache.py b/python/sglang/srt/mem_cache/base_prefix_cache.py index 0c24cb2f1..5ef8e5ddf 100644 --- a/python/sglang/srt/mem_cache/base_prefix_cache.py +++ b/python/sglang/srt/mem_cache/base_prefix_cache.py @@ -26,6 +26,7 @@ from sglang.srt.observability.metrics_collector import ( from sglang.srt.runtime_context import get_observability if TYPE_CHECKING: + from sglang.srt.managers.cache_controller import HiCacheController from sglang.srt.managers.schedule_batch import Req from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.mem_cache.unified_cache.cache_action import ( @@ -233,6 +234,7 @@ class BasePrefixCache(ABC, PrefixCacheTrait): metrics_collector: Optional[RadixCacheMetricsCollector] = ( None # metrics collector for the cache ) + cache_controller: Optional[HiCacheController] = None def init_metrics_collector(self): from sglang.srt.runtime_context import get_server_args diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py index bc19336d4..2c47e326d 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py @@ -5,14 +5,14 @@ import logging import os import threading import time +from dataclasses import replace from queue import Empty, Queue from typing import TYPE_CHECKING, Any, Callable, List, Optional import torch -from sglang.srt.managers.cache_controller import CacheOperation as BaseCacheOperation from sglang.srt.managers.cache_controller import ( - HiCacheAck, + CacheOperation, ) from sglang.srt.managers.cache_controller import ( HiCacheController as BaseHiCacheController, @@ -23,9 +23,6 @@ from sglang.srt.managers.cache_controller import ( from sglang.srt.managers.cache_controller import ( StorageOperation as BaseStorageOperation, ) -from sglang.srt.managers.cache_controller import ( - make_timing_event_pair, -) from sglang.srt.mem_cache.hicache_storage import ( HiCacheStorageExtraInfo, PoolHitPolicy, @@ -33,75 +30,14 @@ from sglang.srt.mem_cache.hicache_storage import ( PoolTransfer, PoolTransferResult, ) +from sglang.srt.mem_cache.l2_transfer import L2Transfer from sglang.srt.mem_cache.memory_pool_host import HostPoolGroup, PoolEntry from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost -from sglang.srt.utils import get_device_module if TYPE_CHECKING: from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator logger = logging.getLogger(__name__) -device_module = get_device_module() - - -class CacheOperation(BaseCacheOperation): - def __init__( - self, - host_indices: torch.Tensor, - device_indices: torch.Tensor, - node_id: int, - priority: Optional[int] = None, - pool_transfers: Optional[list[PoolTransfer]] = None, - ): - super().__init__(host_indices, device_indices, node_id, priority) - self.pool_transfers = pool_transfers - - @staticmethod - def merge_pool_transfers( - ops: List[CacheOperation], - ) -> Optional[list[PoolTransfer]]: - grouped: dict[tuple[PoolName, Optional[PoolName]], list[PoolTransfer]] = {} - for op in ops: - for t in op.pool_transfers or []: - grouped.setdefault((t.name, t.indices_from_pool), []).append(t) - if not grouped: - return None - - def cat_or_none(tensors): - parts = [x for x in tensors if x is not None] - return torch.cat(parts) if parts else None - - return [ - PoolTransfer( - name=ts[0].name, - host_indices=cat_or_none(t.host_indices for t in ts), - device_indices=cat_or_none(t.device_indices for t in ts), - keys=[k for t in ts if t.keys for k in t.keys] or None, - hit_policy=ts[0].hit_policy, - indices_from_pool=ts[0].indices_from_pool, - ) - for ts in grouped.values() - ] - - @staticmethod - def merge_ops(ops: List[CacheOperation]) -> CacheOperation: - if len(ops) == 1: - return ops[0] - host_indices = torch.cat([op.host_indices for op in ops]) - device_indices = torch.cat([op.device_indices for op in ops]) - node_ids = [] - priority = min(op.priority for op in ops) - for op in ops: - node_ids.extend(op.node_ids) - merged = CacheOperation( - host_indices, - device_indices, - -1, - priority, - pool_transfers=CacheOperation.merge_pool_transfers(ops), - ) - merged.node_ids = node_ids - return merged class StorageOperation(BaseStorageOperation): @@ -403,69 +339,134 @@ class HybridCacheController(BaseHiCacheController): self.start_writing() return host_indices - def start_writing(self) -> None: - if not self.write_queue: - return - op = CacheOperation.merge_ops(self.write_queue) - # Page-first staged write-back kernels need CPU destination host indices. - # A HostPoolGroup may mix staged and non-staged child pools, so let it - # normalize indices per child instead of moving the whole operation here. - if ( - self.io_backend == "kernel" - and self.mem_pool_host.layout == "page_first" - and ( - getattr(self.mem_pool_host, "can_use_write_back_jit", False) - or getattr( - self.mem_pool_host, "supports_per_pool_backup_indices", False - ) - ) - ): - host_indices = op.host_indices - device_indices = op.device_indices - resolved_pool_transfers = op.pool_transfers - else: - host_indices, device_indices, resolved_pool_transfers = ( - self.move_hybrid_indices(op) - ) - self.write_queue.clear() - start_event = device_module.Event() - ack_start_event, ack_finish_event, timing_enabled = make_timing_event_pair() - start_event.record() - with device_module.stream(self.write_stream): - start_event.wait(self.write_stream) - ack_start_event.record() - self.mem_pool_host.backup_from_device_all_layer( - self.mem_pool_device, - host_indices, - device_indices, - self.io_backend, - pool_transfers=resolved_pool_transfers, - ) - if self.has_draft and host_indices.numel() > 0: - self.mem_pool_host_draft.backup_from_device_all_layer( - self.mem_pool_device_draft, - host_indices, - device_indices, - self.io_backend, - ) - ack_finish_event.record() - self._record_transfer_indices_on_stream( - self.write_stream, - host_indices, - device_indices, - resolved_pool_transfers, - ) - self.ack_write_queue.append( - HiCacheAck( - start_event=ack_start_event, - finish_event=ack_finish_event, - node_ids=op.node_ids, - num_tokens=len(op.device_indices), - timing_enabled=timing_enabled, - num_tokens_by_pool=self._num_tokens_by_pool(op), - num_bytes=self._transfer_num_bytes(op), - ) + def _move_op_indices( + self, op: CacheOperation + ) -> tuple[torch.Tensor, torch.Tensor, Optional[list[PoolTransfer]]]: + return self.move_hybrid_indices(op) + + def _move_write_operation( + self, op: CacheOperation + ) -> tuple[torch.Tensor, torch.Tensor, Optional[list[PoolTransfer]]]: + host_group = self.mem_pool_host + if self.io_backend != "kernel" or host_group.layout != "page_first": + return self.move_hybrid_indices(op) + if not getattr(host_group, "supports_per_pool_backup_indices", False): + if not getattr(host_group, "can_use_write_back_jit", False): + return self.move_hybrid_indices(op) + return op.host_indices, op.device_indices, op.pool_transfers + + def move_for_pool(host_pool, host_indices, device_indices): + if getattr(host_pool, "can_use_write_back_jit", False): + if host_indices.is_cuda: + host_indices = host_indices.cpu() + return host_indices, device_indices + return self.move_indices(host_indices, device_indices) + + host_indices, device_indices = move_for_pool( + host_group.anchor_entry.host_pool, + op.host_indices, + op.device_indices, ) + pool_transfers = [] + for transfer in op.pool_transfers or []: + entry = host_group.entry_map[transfer.name] + transfer_host_indices, transfer_device_indices = move_for_pool( + entry.host_pool, + transfer.host_indices, + transfer.device_indices, + ) + pool_transfers.append( + replace( + transfer, + host_indices=transfer_host_indices, + device_indices=transfer_device_indices, + ) + ) + return host_indices, device_indices, pool_transfers or None + + def _l2_transfers( + self, + host_indices: torch.Tensor, + device_indices: torch.Tensor, + pool_transfers: Optional[list[PoolTransfer]] = None, + ) -> list[L2Transfer]: + anchor = self.mem_pool_host.anchor_entry + transfers = [] + if host_indices.numel() > 0: + transfers.append( + L2Transfer( + host_pool=anchor.host_pool, + device_pool=anchor.device_pool, + host_indices=host_indices, + device_indices=device_indices, + layer_mapper=anchor.layer_mapper, + ) + ) + for pool_transfer in pool_transfers or []: + if ( + pool_transfer.host_indices is None + or pool_transfer.device_indices is None + ): + raise ValueError(f"Unresolved L2 transfer for {pool_transfer.name}.") + entry = self.mem_pool_host.entry_map[pool_transfer.name] + transfers.append( + L2Transfer( + host_pool=entry.host_pool, + device_pool=entry.device_pool, + host_indices=pool_transfer.host_indices, + device_indices=pool_transfer.device_indices, + layer_mapper=entry.layer_mapper, + ) + ) + if self.has_draft and host_indices.numel() > 0: + transfers.append( + L2Transfer( + host_pool=self.mem_pool_host_draft, + device_pool=self.mem_pool_device_draft, + host_indices=host_indices, + device_indices=device_indices, + ) + ) + return transfers + + def _l2_load_transfers( + self, + host_indices: torch.Tensor, + device_indices: torch.Tensor, + pool_transfers: Optional[list[PoolTransfer]] = None, + ) -> list[L2Transfer]: + transfers = self._l2_transfers(host_indices, device_indices, pool_transfers) + if getattr(self, "has_mtp_draft", False): + target_transfers = list(transfers) + for depth, draft_device_pool in enumerate(self.mtp_draft_device_pools): + for transfer in target_transfers: + if transfer.layer_mapper is None: + continue + draft_host_layer = transfer.layer_mapper(self.layer_num + depth) + if draft_host_layer is None: + continue + + def draft_layer_mapper( + layer_id: int, + *, + expected_layer_id: int = depth, + host_layer_id: int = draft_host_layer, + ) -> Optional[int]: + if layer_id == expected_layer_id: + return host_layer_id + return None + + transfers.append( + L2Transfer( + host_pool=transfer.host_pool, + device_pool=draft_device_pool, + host_indices=transfer.host_indices, + device_indices=transfer.device_indices, + layer_mapper=draft_layer_mapper, + is_draft=True, + ) + ) + return transfers def _num_tokens_by_pool(self, op: CacheOperation) -> dict[str, int]: """Per-pool token counts for a merged transfer op (anchor + extra @@ -546,108 +547,6 @@ class HybridCacheController(BaseHiCacheController): ) return device_indices - def start_loading(self) -> int: - if not self.load_queue: - return -1 - producer_id = self.layer_done_counter.update_producer() - op = CacheOperation.merge_ops(self.load_queue) - host_indices, device_indices, resolved_pool_transfers = ( - self.move_hybrid_indices(op) - ) - self.load_queue.clear() - producer_event = self.layer_done_counter.events[producer_id] - producer_event.start_event.record() - - ack_start_event, ack_finish_event, timing_enabled = make_timing_event_pair() - - with device_module.stream(self.load_stream): - producer_event.start_event.wait(self.load_stream) - ack_start_event.record() - target_device_pool = self.mem_pool_host.anchor_entry.device_pool - for i in range(self.layer_num): - self.mem_pool_host.load_to_device_per_layer( - target_device_pool, - host_indices, - device_indices, - i, - self.io_backend, - pool_transfers=resolved_pool_transfers, - ) - if ( - self.has_draft - and host_indices.numel() > 0 - and i < self.mem_pool_host_draft.layer_num - ): - self.mem_pool_host_draft.load_to_device_per_layer( - self.mem_pool_device_draft, - host_indices, - device_indices, - i, - self.io_backend, - ) - - # HiCache now supports draft caches through two paths: - # - # - Packed: standard NextN/MTP models (DeepSeek-V3.2, GLM-5.x, - # DeepSeek-V4, MiMo-V2.5) and DeepSeek-V4 DSpark. Draft KV/indexer/SWA - # buffers are appended to the matching target host pools as tail layers - # and share their slot mappings. D2H/H2D therefore moves target and draft - # in the same cache operation; the branch below restores the tail layers. - # - # - Sidecar: standalone EAGLE/EAGLE3 (for example Llama-2/Llama-3.1), - # DFlash (for example Gemma-4), and non-DeepSeek-V4 DSpark. Draft - # KV/indexer/SWA gets a separate host-pool entry sized to its source target - # pool. Its PoolTransfer follows the target KV or SWA indices and is - # attached to the same cache operation. - - if self.has_mtp_draft and i < len(self.mtp_draft_device_pools): - self.mem_pool_host.load_to_device_per_layer( - self.mtp_draft_device_pools[i], - host_indices, - device_indices, - self.layer_num + i, - self.io_backend, - pool_transfers=resolved_pool_transfers, - is_draft=True, - ) - producer_event.complete(i) - ack_finish_event.record() - self._record_transfer_indices_on_stream( - self.load_stream, - host_indices, - device_indices, - resolved_pool_transfers, - ) - self.ack_load_queue.append( - HiCacheAck( - ack_start_event, - ack_finish_event, - op.node_ids, - num_tokens=len(op.device_indices), - timing_enabled=timing_enabled, - num_tokens_by_pool=self._num_tokens_by_pool(op), - num_bytes=self._transfer_num_bytes(op), - ) - ) - return producer_id - - def _record_transfer_indices_on_stream( - self, - stream: torch.Stream, - host_indices: torch.Tensor, - device_indices: torch.Tensor, - pool_transfers: Optional[list[PoolTransfer]] = None, - ) -> None: - if host_indices.is_cuda: - host_indices.record_stream(stream) - if device_indices.is_cuda: - device_indices.record_stream(stream) - for transfer in pool_transfers or []: - if transfer.host_indices is not None and transfer.host_indices.is_cuda: - transfer.host_indices.record_stream(stream) - if transfer.device_indices is not None and transfer.device_indices.is_cuda: - transfer.device_indices.record_stream(stream) - def prefetch( self, request_id: str, diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index 84b1e84fc..0f1e194e0 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -151,6 +151,120 @@ def build_pool_entry( ) +def build_kv_only_group( + *, + page_size: int, + server_args: ServerArgs, + kv_pool: Any, + full_layer_mapping: dict[int, int], + use_mla: bool, + override_kv_cache_dim: Optional[int] = None, + host_size: Optional[float] = None, + mtp_draft_device_pools: tuple[Any, ...] = (), +) -> HostPoolGroup: + """Anchor-only host pool group for a flat MHA/MLA device pool.""" + transfer_layer_num = len(full_layer_mapping) + kv_host_pool = build_kv_host_pool( + kv_pool=kv_pool, + page_size=page_size, + server_args=server_args, + use_mla=use_mla, + override_kv_cache_dim=override_kv_cache_dim, + host_size=host_size, + mtp_draft_device_pools=mtp_draft_device_pools, + ) + if mtp_draft_device_pools: + full_layer_mapping = _with_mtp_layer_mapping( + full_layer_mapping, + transfer_layer_start=transfer_layer_num, + target_device_layer_num=kv_pool.layer_num, + draft_layer_num=len(mtp_draft_device_pools), + ) + return HostPoolGroup( + [ + build_pool_entry( + name=PoolName.KV, + host_pool=kv_host_pool, + device_pool=kv_pool, + layer_mapping=full_layer_mapping, + transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools), + is_anchor=True, + ) + ] + ) + + +def build_hybrid_swa_group( + *, + page_size: int, + server_args: ServerArgs, + full_kv_pool: Any, + swa_kv_pool: Any, + full_layer_mapping: dict[int, int], + swa_layer_mapping: dict[int, int], + use_mla: bool, + kv_host_size: Optional[float] = None, + swa_host_size: Optional[float] = None, + host_swa_evict_fn: Optional[Callable[[int], Any]] = None, + device_swa_evict_fn: Optional[Callable[[int], Any]] = None, + swa_attn_allocator: Any = None, + mtp_swa_device_pools: tuple[Any, ...] = (), +) -> HostPoolGroup: + """Anchor (full) + SWA host pool group for a hybrid-SWA device pool.""" + transfer_layer_num = len(full_layer_mapping | swa_layer_mapping) + kv_host_pool = build_kv_host_pool( + kv_pool=full_kv_pool, + page_size=page_size, + server_args=server_args, + use_mla=use_mla, + host_size=kv_host_size, + pool_label="full", + ) + swa_host_pool = build_kv_host_pool( + kv_pool=swa_kv_pool, + page_size=page_size, + server_args=server_args, + use_mla=use_mla, + host_size=swa_host_size, + mtp_draft_device_pools=mtp_swa_device_pools, + pool_label="swa", + ) + if mtp_swa_device_pools: + swa_layer_mapping = _with_mtp_layer_mapping( + swa_layer_mapping, + transfer_layer_start=transfer_layer_num, + target_device_layer_num=swa_kv_pool.layer_num, + draft_layer_num=len(mtp_swa_device_pools), + ) + return HostPoolGroup( + [ + build_pool_entry( + name=PoolName.KV, + host_pool=kv_host_pool, + device_pool=full_kv_pool, + layer_mapping=full_layer_mapping, + transfer_layer_num=transfer_layer_num, + is_anchor=True, + ), + build_pool_entry( + name=PoolName.SWA, + host_pool=swa_host_pool, + device_pool=swa_kv_pool, + layer_mapping=swa_layer_mapping, + transfer_layer_num=transfer_layer_num + len(mtp_swa_device_pools), + host_evict_fn=host_swa_evict_fn, + device_evict_fn=device_swa_evict_fn, + device_alloc_fn=( + swa_attn_allocator.alloc if swa_attn_allocator is not None else None + ), + device_free_fn=( + swa_attn_allocator.free if swa_attn_allocator is not None else None + ), + ), + ] + ) + + def build_kv_only_stack( *, params: CacheInitParams, @@ -167,33 +281,15 @@ def build_kv_only_stack( enable_storage_metrics: bool = False, ) -> tuple[HostPoolGroup, HybridCacheController]: transfer_layer_num = len(full_layer_mapping) - kv_host_pool = build_kv_host_pool( - kv_pool=kv_pool, + host_pool_group = build_kv_only_group( page_size=params.page_size, server_args=server_args, + kv_pool=kv_pool, + full_layer_mapping=full_layer_mapping, use_mla=use_mla, override_kv_cache_dim=override_kv_cache_dim, mtp_draft_device_pools=params.mtp_draft_device_pools, ) - if params.mtp_draft_device_pools: - full_layer_mapping = _with_mtp_layer_mapping( - full_layer_mapping, - transfer_layer_start=transfer_layer_num, - target_device_layer_num=kv_pool.layer_num, - draft_layer_num=len(params.mtp_draft_device_pools), - ) - - entries = [ - build_pool_entry( - name=PoolName.KV, - host_pool=kv_host_pool, - device_pool=kv_pool, - layer_mapping=full_layer_mapping, - transfer_layer_num=transfer_layer_num + len(params.mtp_draft_device_pools), - is_anchor=True, - ) - ] - host_pool_group = HostPoolGroup(entries) cache_controller = HybridCacheController( params.token_to_kv_pool_allocator, host_pool_group, @@ -248,56 +344,22 @@ def build_hybrid_swa_stack( server_args.hicache_size, (full_kv_pool, swa_kv_pool) ) - kv_host_pool = build_kv_host_pool( - kv_pool=full_kv_pool, + host_pool_group = build_hybrid_swa_group( page_size=params.page_size, server_args=server_args, + full_kv_pool=full_kv_pool, + swa_kv_pool=swa_kv_pool, + full_layer_mapping=full_layer_mapping, + swa_layer_mapping=swa_layer_mapping, use_mla=use_mla, - host_size=kv_host_size, - pool_label="full", + kv_host_size=kv_host_size, + swa_host_size=swa_host_size, + host_swa_evict_fn=host_swa_evict_fn, + device_swa_evict_fn=device_swa_evict_fn, + # For SWA hybrid, device allocation goes through the inner allocator. + swa_attn_allocator=params.token_to_kv_pool_allocator.swa_attn_allocator, + mtp_swa_device_pools=mtp_swa_device_pools, ) - swa_host_pool = build_kv_host_pool( - kv_pool=swa_kv_pool, - page_size=params.page_size, - server_args=server_args, - use_mla=use_mla, - host_size=swa_host_size, - mtp_draft_device_pools=mtp_swa_device_pools, - pool_label="swa", - ) - - if mtp_swa_device_pools: - swa_layer_mapping = _with_mtp_layer_mapping( - swa_layer_mapping, - transfer_layer_start=transfer_layer_num, - target_device_layer_num=swa_kv_pool.layer_num, - draft_layer_num=len(mtp_swa_device_pools), - ) - - # For SWA hybrid, the device alloc/free goes through the inner swa_attn_allocator - swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator - entries = [ - build_pool_entry( - name=PoolName.KV, - host_pool=kv_host_pool, - device_pool=full_kv_pool, - layer_mapping=full_layer_mapping, - transfer_layer_num=transfer_layer_num, - is_anchor=True, - ), - build_pool_entry( - name=PoolName.SWA, - host_pool=swa_host_pool, - device_pool=swa_kv_pool, - layer_mapping=swa_layer_mapping, - transfer_layer_num=transfer_layer_num + len(mtp_swa_device_pools), - host_evict_fn=host_swa_evict_fn, - device_evict_fn=device_swa_evict_fn, - device_alloc_fn=swa_attn_allocator.alloc, - device_free_fn=swa_attn_allocator.free, - ), - ] - host_pool_group = HostPoolGroup(entries) cache_controller = HybridCacheController( params.token_to_kv_pool_allocator, host_pool_group, @@ -846,7 +908,7 @@ def build_anchor_sidecar_stack( mtp_draft_device_pools=mtp_draft_device_pools, ) sidecar_host_pool = sidecar_host_pool_factory(kv_host_pool) - # Let HostPoolGroup dispatch packed MTP tail layers through the normal path. + # Expose packed MTP tail layers to the controller's flat transfer builder. if mtp_draft_device_pools: full_layer_mapping = _with_mtp_layer_mapping( full_layer_mapping, diff --git a/python/sglang/srt/mem_cache/l2_transfer.py b/python/sglang/srt/mem_cache/l2_transfer.py new file mode 100644 index 000000000..f80afd9c7 --- /dev/null +++ b/python/sglang/srt/mem_cache/l2_transfer.py @@ -0,0 +1,127 @@ +from __future__ import annotations + +import logging +from functools import cache +from typing import Any, Callable, NamedTuple, Optional + +import torch + +from sglang.srt.utils import get_device_module + +logger = logging.getLogger(__name__) +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 timing; L2 transfer timing is disabled", + 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 L2Transfer(NamedTuple): + host_pool: Any + device_pool: Any + host_indices: torch.Tensor + device_indices: torch.Tensor + layer_mapper: Optional[Callable[[int], Optional[int]]] = None + is_draft: bool = False + + +class TransferCompletion(NamedTuple): + start_event: Any + finish_event: Any + timing_enabled: bool + + +class L2TransferEngine: + """Runs resolved device↔host transfers without owning cache state.""" + + def __init__(self, io_backend: str): + self.io_backend = io_backend + self.device_to_host_stream = device_module.Stream() + self.host_to_device_stream = device_module.Stream() + + def submit_device_to_host(self, transfers: list[L2Transfer]) -> TransferCompletion: + start_event = self._start_event(None) + ack_start, ack_finish, timing_enabled = make_timing_event_pair() + with device_module.stream(self.device_to_host_stream): + start_event.wait(self.device_to_host_stream) + ack_start.record() + for transfer in transfers: + transfer.host_pool.backup_from_device_all_layer( + transfer.device_pool, + transfer.host_indices, + transfer.device_indices, + self.io_backend, + ) + ack_finish.record() + self._record_stream(transfers, self.device_to_host_stream) + return TransferCompletion(ack_start, ack_finish, timing_enabled) + + def submit_host_to_device( + self, + transfers: list[L2Transfer], + *, + layer_num: int, + start_event=None, + on_layer_done=None, + ) -> TransferCompletion: + start_event = self._start_event(start_event) + ack_start, ack_finish, timing_enabled = make_timing_event_pair() + primary = transfers[0] if transfers else None + with device_module.stream(self.host_to_device_stream): + start_event.wait(self.host_to_device_stream) + ack_start.record() + for layer_id in range(layer_num): + for transfer in transfers: + local_layer_id = ( + transfer.layer_mapper(layer_id) + if transfer.layer_mapper is not None + else layer_id + ) + if local_layer_id is None or ( + transfer is not primary + and transfer.layer_mapper is None + and layer_id >= transfer.host_pool.layer_num + ): + continue + transfer.host_pool.load_to_device_per_layer( + transfer.device_pool, + transfer.host_indices, + transfer.device_indices, + local_layer_id, + self.io_backend, + is_draft=transfer.is_draft, + ) + if on_layer_done is not None: + on_layer_done(layer_id) + ack_finish.record() + self._record_stream(transfers, self.host_to_device_stream) + return TransferCompletion(ack_start, ack_finish, timing_enabled) + + @staticmethod + def _start_event(start_event): + if start_event is None: + start_event = device_module.Event() + start_event.record() + return start_event + + @staticmethod + def _record_stream(transfers: list[L2Transfer], stream) -> None: + for transfer in transfers: + for indices in (transfer.host_indices, transfer.device_indices): + if indices.is_cuda: + indices.record_stream(stream) diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index f9c88bd26..0024ebc6e 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -1662,120 +1662,6 @@ class HostPoolGroup: def set_from_flat_data_page(self, index: int, data_page) -> None: return self.anchor_entry.host_pool.set_from_flat_data_page(index, data_page) - def load_to_device_per_layer( - self, - device_pool, - host_indices, - device_indices, - layer_id, - io_backend, - pool_transfers: Optional[list] = None, - *, - is_draft: bool = False, - ) -> None: - # 1. Anchor (KV) transfer - anchor = self.anchor_entry - local_layer_id = anchor.layer_mapper(layer_id) - if local_layer_id is not None and host_indices.numel() > 0: - anchor.host_pool.load_to_device_per_layer( - device_pool if is_draft else anchor.device_pool, - host_indices, - device_indices, - local_layer_id, - io_backend, - is_draft=is_draft, - ) - - # 2. Extra pool transfers - for transfer in pool_transfers or []: - entry = self.entry_map.get(transfer.name) - if entry is None or transfer.host_indices is None: - continue - local_layer_id = entry.layer_mapper(layer_id) - if local_layer_id is None: - continue - entry.host_pool.load_to_device_per_layer( - device_pool if is_draft else entry.device_pool, - transfer.host_indices, - transfer.device_indices, - local_layer_id, - io_backend, - is_draft=is_draft, - ) - - def _backup_uses_cpu_host_indices(self, host_pool, io_backend) -> bool: - return ( - io_backend == "kernel" - and getattr(host_pool, "layout", None) == "page_first" - and getattr(host_pool, "can_use_write_back_jit", False) - ) - - def _kernel_index_device(self, entry, device_indices): - if device_indices is not None and device_indices.is_cuda: - return device_indices.device - return getattr(entry.device_pool, "device", None) - - def _normalize_backup_indices( - self, entry, host_indices, device_indices, io_backend - ): - if io_backend != "kernel": - return host_indices, device_indices - - if self._backup_uses_cpu_host_indices(entry.host_pool, io_backend): - if host_indices.is_cuda: - host_indices = host_indices.cpu() - return host_indices, device_indices - - if not host_indices.is_cuda: - target_device = self._kernel_index_device(entry, device_indices) - if target_device is not None: - host_indices = host_indices.to(target_device, non_blocking=True) - if host_indices.is_cuda: - host_indices.record_stream( - torch.cuda.current_stream(host_indices.device) - ) - return host_indices, device_indices - - def backup_from_device_all_layer( - self, - device_pool, - host_indices, - device_indices, - io_backend, - pool_transfers: Optional[list] = None, - ) -> None: - # 1. Anchor (KV) backup - # A zero-length anchor denotes a component-only backup. - if host_indices.numel() > 0: - anchor_host_indices, anchor_device_indices = self._normalize_backup_indices( - self.anchor_entry, host_indices, device_indices, io_backend - ) - self.anchor_entry.host_pool.backup_from_device_all_layer( - self.anchor_entry.device_pool, - anchor_host_indices, - anchor_device_indices, - io_backend, - ) - # 2. Extra pool backup - for transfer in pool_transfers or []: - entry = self.entry_map.get(transfer.name) - if entry is None or transfer.host_indices is None: - continue - transfer_host_indices, transfer_device_indices = ( - self._normalize_backup_indices( - entry, - transfer.host_indices, - transfer.device_indices, - io_backend, - ) - ) - entry.host_pool.backup_from_device_all_layer( - entry.device_pool, - transfer_host_indices, - transfer_device_indices, - io_backend, - ) - class DSAIndexerPoolHost(HostKVCache): """Host-side DSA index buffers only. Slot layout matches the anchor MLA host pool.""" diff --git a/test/registered/unit/mem_cache/test_hicache_load_back_timing.py b/test/registered/unit/mem_cache/test_hicache_load_back_timing.py index b92e63475..46c5db8ec 100644 --- a/test/registered/unit/mem_cache/test_hicache_load_back_timing.py +++ b/test/registered/unit/mem_cache/test_hicache_load_back_timing.py @@ -16,12 +16,14 @@ register_cuda_ci(est_time=5, stage="base-b", runner_config="1-gpu-small") class TestLoadBackDurationMetric(CustomTestCase): def setUp(self): from sglang.srt.managers import cache_controller as cc + from sglang.srt.mem_cache import l2_transfer as transfer - cc._timing_events_supported.cache_clear() + transfer._timing_events_supported.cache_clear() self.cc = cc + self.transfer = transfer def _completed_pair(self, payload_floats=1024 * 1024): - start, finish, timing_enabled = self.cc.make_timing_event_pair() + start, finish, timing_enabled = self.transfer.make_timing_event_pair() self.assertTrue(timing_enabled) stream = torch.cuda.Stream() start.record() @@ -46,9 +48,11 @@ class TestLoadBackDurationMetric(CustomTestCase): 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() + with patch.object( + self.transfer.device_module, "Event", side_effect=create_event + ): + self.transfer._timing_events_supported.cache_clear() + start, finish, timing_enabled = self.transfer.make_timing_event_pair() self.assertFalse(timing_enabled) self.assertIs(start, events[0]) diff --git a/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py b/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py index 0a846b2ec..fc07cfb81 100644 --- a/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py +++ b/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py @@ -7,17 +7,17 @@ from unittest import mock import torch -from sglang.srt.managers import cache_controller as manager_cache_controller -from sglang.srt.managers.cache_controller import CacheOperation as ManagerCacheOperation -from sglang.srt.managers.cache_controller import ( - HiCacheController, +from sglang.srt.managers.cache_controller import CacheOperation, HiCacheController +from sglang.srt.mem_cache import l2_transfer as transfer_module +from sglang.srt.mem_cache.hicache_storage import ( + PoolHitPolicy, + PoolName, + PoolTransfer, ) -from sglang.srt.mem_cache.hicache_storage import PoolName, PoolTransfer -from sglang.srt.mem_cache.hybrid_cache import hybrid_cache_controller from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( - CacheOperation, HybridCacheController, ) +from sglang.srt.mem_cache.l2_transfer import L2Transfer, L2TransferEngine from sglang.srt.mem_cache.memory_pool_host import ( DeepSeekV4PagedHostPool, DeepSeekV4StateHostPool, @@ -30,6 +30,7 @@ from sglang.srt.mem_cache.memory_pool_host import ( from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=3, suite="base-a-test-cpu") @@ -59,6 +60,33 @@ def _device_pool_stub(*, layer_num: int, **fields) -> SimpleNamespace: ) +def _host_group_stub(captured, *, can_use_write_back_jit: bool) -> SimpleNamespace: + class FakeHostPool: + size_per_token = 2 + + def backup_from_device_all_layer( + self, device_pool, host_indices, device_indices, io_backend + ): + captured.append(host_indices) + + entries = [ + PoolEntry( + name=name, + host_pool=FakeHostPool(), + device_pool=None, + layer_mapper=lambda layer_id: layer_id, + is_primary_index_anchor=name == PoolName.KV, + ) + for name in (PoolName.KV, PoolName.SWA, PoolName.DEEPSEEK_V4_C4) + ] + return SimpleNamespace( + layout="page_first", + can_use_write_back_jit=can_use_write_back_jit, + anchor_entry=entries[0], + entry_map={entry.name: entry for entry in entries}, + ) + + def _cpu_staged_lf_pf_copy( src_registry, *, @@ -152,21 +180,204 @@ class _FakeEvent: class _FakeDeviceModule: Event = _FakeEvent + @staticmethod + def Stream(): + return object() + @staticmethod @contextmanager def stream(stream): yield -class TestHiCacheStagedWriteBackDispatch(unittest.TestCase): +class TestHiCacheStagedWriteBackDispatch(CustomTestCase): def setUp(self): - # start_writing probes timing support via a module-cached check; - # clear it on both sides so results from (or against) the fake - # device module never leak across tests. - manager_cache_controller._timing_events_supported.cache_clear() + transfer_module._timing_events_supported.cache_clear() + self.addCleanup(transfer_module._timing_events_supported.cache_clear) - def tearDown(self): - manager_cache_controller._timing_events_supported.cache_clear() + @staticmethod + def _start_writing(controller): + with mock.patch.object(transfer_module, "device_module", _FakeDeviceModule): + controller.l2_transfer_engine = L2TransferEngine("kernel") + controller.start_writing() + + def test_hybrid_load_forwards_merged_pool_transfers(self): + transfer = PoolTransfer( + name=PoolName.SWA, + host_indices=_indices(0, 2), + device_indices=_indices(2, 4), + keys=["page-key"], + hit_policy=PoolHitPolicy.TRAILING_PAGES, + ) + op = CacheOperation(_indices(0, 4), _indices(4, 8), 7) + op.pool_transfers = [transfer] + controller = mock.Mock(spec=HybridCacheController) + controller.load_queue = [op, op] + controller.layer_done_counter = mock.MagicMock() + controller.layer_done_counter.update_producer.return_value = 0 + controller._move_op_indices.side_effect = lambda op: ( + op.host_indices, + op.device_indices, + op.pool_transfers, + ) + controller.mem_pool_host = _host_group_stub([], can_use_write_back_jit=False) + controller.has_draft = False + controller.has_mtp_draft = False + controller._l2_transfers.side_effect = lambda *args: ( + HybridCacheController._l2_transfers(controller, *args) + ) + controller._l2_load_transfers.side_effect = lambda *args: ( + HybridCacheController._l2_load_transfers(controller, *args) + ) + controller._num_tokens_by_pool.return_value = {} + controller._transfer_num_bytes.return_value = 0 + controller.l2_transfer_engine = mock.Mock() + completion = SimpleNamespace( + start_event=object(), finish_event=object(), timing_enabled=False + ) + controller.l2_transfer_engine.submit_host_to_device.return_value = completion + controller.layer_num = 2 + controller.ack_load_queue = [] + + self.assertEqual(HybridCacheController.start_loading(controller), 0) + + merged_op = controller._move_op_indices.call_args.args[0] + merged_transfer = merged_op.pool_transfers[0] + self.assertEqual(merged_transfer.host_indices.tolist(), [0, 1, 0, 1]) + self.assertEqual(merged_transfer.keys, ["page-key", "page-key"]) + self.assertEqual(merged_transfer.hit_policy, PoolHitPolicy.TRAILING_PAGES) + controller._l2_load_transfers.assert_called_once() + l2_transfers = ( + controller.l2_transfer_engine.submit_host_to_device.call_args.args[0] + ) + self.assertEqual(len(l2_transfers), 2) + self.assertEqual(l2_transfers[1].host_indices.tolist(), [0, 1, 0, 1]) + self.assertEqual( + len( + HybridCacheController._l2_transfers( + controller, _indices(0, 0), _indices(0, 0), [merged_transfer] + ) + ), + 1, + ) + controller._num_tokens_by_pool.assert_called_once_with(merged_op) + self.assertEqual(controller.ack_load_queue[0].node_ids, [7, 7]) + + def test_l2_transfer_maps_global_layers(self): + host_pool = mock.Mock() + transfer = L2Transfer( + host_pool=host_pool, + device_pool=mock.sentinel.device_pool, + host_indices=_indices(0, 2), + device_indices=_indices(2, 4), + layer_mapper={1: 0, 3: 1}.get, + ) + with mock.patch.object(transfer_module, "device_module", _FakeDeviceModule): + L2TransferEngine("kernel").submit_host_to_device([transfer], layer_num=4) + + self.assertEqual( + [ + call.args[3] + for call in host_pool.load_to_device_per_layer.call_args_list + ], + [0, 1], + ) + + def test_packed_draft_load_is_flattened_into_l2_transfers(self): + host_pool = mock.Mock() + controller = HybridCacheController.__new__(HybridCacheController) + controller.mem_pool_host = SimpleNamespace( + anchor_entry=PoolEntry( + name=PoolName.KV, + host_pool=host_pool, + device_pool=mock.sentinel.target_device_pool, + layer_mapper={0: 0, 1: 1, 2: 2}.get, + is_primary_index_anchor=True, + ), + entry_map={}, + ) + controller.layer_num = 2 + controller.has_mtp_draft = True + controller.mtp_draft_device_pools = (mock.sentinel.draft_device_pool,) + controller.has_draft = False + + self.assertEqual( + len(controller._l2_transfers(_indices(0, 2), _indices(2, 4))), 1 + ) + transfers = controller._l2_load_transfers(_indices(0, 2), _indices(2, 4)) + + self.assertEqual(len(transfers), 2) + self.assertFalse(transfers[0].is_draft) + self.assertTrue(transfers[1].is_draft) + with mock.patch.object(transfer_module, "device_module", _FakeDeviceModule): + L2TransferEngine("kernel").submit_host_to_device(transfers, layer_num=2) + self.assertEqual( + [ + call.args[3] + for call in host_pool.load_to_device_per_layer.call_args_list + ], + [0, 2, 1], + ) + self.assertIs( + host_pool.load_to_device_per_layer.call_args_list[1].args[0], + mock.sentinel.draft_device_pool, + ) + self.assertTrue( + host_pool.load_to_device_per_layer.call_args_list[1].kwargs["is_draft"] + ) + + def test_mixed_staged_write_resolves_indices_per_pool(self): + anchor_host_pool = SimpleNamespace(can_use_write_back_jit=True) + extra_host_pool = SimpleNamespace(can_use_write_back_jit=False) + anchor_entry = PoolEntry( + name=PoolName.KV, + host_pool=anchor_host_pool, + device_pool=None, + layer_mapper=lambda layer_id: layer_id, + is_primary_index_anchor=True, + ) + extra_entry = PoolEntry( + name=PoolName.SWA, + host_pool=extra_host_pool, + device_pool=None, + layer_mapper=lambda layer_id: layer_id, + ) + host_group = SimpleNamespace( + layout="page_first", + can_use_write_back_jit=False, + supports_per_pool_backup_indices=True, + anchor_entry=anchor_entry, + entry_map={PoolName.KV: anchor_entry, PoolName.SWA: extra_entry}, + ) + transfer = PoolTransfer( + name=PoolName.SWA, + host_indices=_indices(4, 6), + device_indices=_indices(6, 8), + ) + op = CacheOperation( + host_indices=_indices(0, 2), + device_indices=_indices(2, 4), + node_id=1, + pool_transfers=[transfer], + ) + controller = HybridCacheController.__new__(HybridCacheController) + controller.io_backend = "kernel" + controller.mem_pool_host = host_group + controller.move_indices = mock.Mock( + return_value=(mock.sentinel.host_indices, mock.sentinel.device_indices) + ) + + host_indices, device_indices, pool_transfers = controller._move_write_operation( + op + ) + + self.assertIs(host_indices, op.host_indices) + self.assertIs(device_indices, op.device_indices) + controller.move_indices.assert_called_once_with( + transfer.host_indices, transfer.device_indices + ) + self.assertIs(pool_transfers[0].host_indices, mock.sentinel.host_indices) + self.assertIs(pool_transfers[0].device_indices, mock.sentinel.device_indices) def _patched_transfers(self, src_registry=None, module=MEMORY_POOL_HOST_MODULE): staged_side_effect = None @@ -701,26 +912,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase): self.assertIsNone(group.destroy()) def test_write_back_jit_hybrid_write_keeps_extra_host_indices_on_cpu(self): - captured = {} - - class FakeHostGroup: - layout = "page_first" - can_use_write_back_jit = True - anchor_entry = SimpleNamespace( - name=PoolName.KV, host_pool=SimpleNamespace(size_per_token=2) - ) - entry_map = {} - - def backup_from_device_all_layer( - self, - device_pool, - host_indices, - device_indices, - io_backend, - pool_transfers=None, - ): - captured["host_indices"] = host_indices - captured["pool_transfers"] = pool_transfers + captured = [] controller = HybridCacheController.__new__(HybridCacheController) controller.write_queue = [ @@ -738,53 +930,25 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase): ) ] controller.io_backend = "kernel" - controller.mem_pool_host = FakeHostGroup() + controller.mem_pool_host = _host_group_stub( + captured, can_use_write_back_jit=True + ) controller.mem_pool_device = None controller.has_draft = False - controller.write_stream = object() controller.ack_write_queue = [] - controller._record_transfer_indices_on_stream = lambda *args: None controller.move_hybrid_indices = mock.Mock( side_effect=AssertionError( "write-back JIT kernel write should not move indices" ) ) - with ( - mock.patch.object( - hybrid_cache_controller, "device_module", _FakeDeviceModule - ), - mock.patch.object( - manager_cache_controller, "device_module", _FakeDeviceModule - ), - ): - controller.start_writing() + self._start_writing(controller) controller.move_hybrid_indices.assert_not_called() - self.assertEqual(captured["host_indices"].device.type, "cpu") - self.assertEqual(captured["pool_transfers"][0].host_indices.device.type, "cpu") + self.assertEqual([indices.device.type for indices in captured], ["cpu", "cpu"]) def test_hybrid_write_moves_indices_without_write_back_jit(self): - captured = {} - - class FakeHostGroup: - layout = "page_first" - can_use_write_back_jit = False - anchor_entry = SimpleNamespace( - name=PoolName.KV, host_pool=SimpleNamespace(size_per_token=2) - ) - entry_map = {} - - def backup_from_device_all_layer( - self, - device_pool, - host_indices, - device_indices, - io_backend, - pool_transfers=None, - ): - captured["host_indices"] = host_indices - captured["pool_transfers"] = pool_transfers + captured = [] op = CacheOperation( host_indices=_indices(0, 4), @@ -801,29 +965,20 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase): controller = HybridCacheController.__new__(HybridCacheController) controller.write_queue = [op] controller.io_backend = "kernel" - controller.mem_pool_host = FakeHostGroup() + controller.mem_pool_host = _host_group_stub( + captured, can_use_write_back_jit=False + ) controller.mem_pool_device = None controller.has_draft = False - controller.write_stream = object() controller.ack_write_queue = [] - controller._record_transfer_indices_on_stream = lambda *args: None controller.move_hybrid_indices = mock.Mock( return_value=(op.host_indices, op.device_indices, op.pool_transfers) ) - with ( - mock.patch.object( - hybrid_cache_controller, "device_module", _FakeDeviceModule - ), - mock.patch.object( - manager_cache_controller, "device_module", _FakeDeviceModule - ), - ): - controller.start_writing() + self._start_writing(controller) controller.move_hybrid_indices.assert_called_once() - self.assertEqual(captured["host_indices"].device.type, "cpu") - self.assertEqual(captured["pool_transfers"][0].host_indices.device.type, "cpu") + self.assertEqual([indices.device.type for indices in captured], ["cpu", "cpu"]) def test_write_back_jit_cache_controller_keeps_host_indices_on_cpu(self): captured = {} @@ -840,7 +995,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase): controller = HiCacheController.__new__(HiCacheController) controller.write_queue = [ - ManagerCacheOperation( + CacheOperation( host_indices=_indices(0, 4), device_indices=_indices(4, 8), node_id=1, @@ -850,7 +1005,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase): controller.mem_pool_host = FakeHostPool() controller.mem_pool_device = None controller.has_draft = False - controller.write_stream = object() + controller.device = "cuda" controller.ack_write_queue = [] controller.move_indices = mock.Mock( side_effect=AssertionError( @@ -858,10 +1013,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase): ) ) - with mock.patch.object( - manager_cache_controller, "device_module", _FakeDeviceModule - ): - controller.start_writing() + self._start_writing(controller) controller.move_indices.assert_not_called() self.assertEqual(captured["host_indices"].device.type, "cpu") @@ -879,7 +1031,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase): ): captured["host_indices"] = host_indices - op = ManagerCacheOperation( + op = CacheOperation( host_indices=_indices(0, 4), device_indices=_indices(4, 8), node_id=1, @@ -890,16 +1042,13 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase): controller.mem_pool_host = FakeHostPool() controller.mem_pool_device = None controller.has_draft = False - controller.write_stream = object() + controller.device = "cuda" controller.ack_write_queue = [] controller.move_indices = mock.Mock( return_value=(op.host_indices, op.device_indices) ) - with mock.patch.object( - manager_cache_controller, "device_module", _FakeDeviceModule - ): - controller.start_writing() + self._start_writing(controller) controller.move_indices.assert_called_once() self.assertEqual(captured["host_indices"].device.type, "cpu") diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index 3c720e403..f1998915d 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -148,7 +148,6 @@ _EXPOSED = { ("disaggregation/common/conn.py", "disaggregation_bootstrap_port"), ("disaggregation/common/conn.py", "pp_size"), ("disaggregation/decode_kvcache_offload_manager.py", "hicache_io_backend"), - ("disaggregation/decode_kvcache_offload_manager.py", "hicache_mem_layout"), ("disaggregation/decode_kvcache_offload_manager.py", "served_model_name"), ("disaggregation/encode_receiver.py", "disaggregation_ib_device"), ("disaggregation/encode_receiver.py", "encoder_transfer_backend"),