refactor(hicache): flatten L2 transfer execution (#34793)

GB300 test fails unrelated
This commit is contained in:
cctry
2026-08-16 00:33:34 -07:00
committed by GitHub
parent 56a759cffc
commit 8922bb98e2
10 changed files with 767 additions and 648 deletions
@@ -13,14 +13,14 @@ from sglang.srt.environ import envs
from sglang.srt.managers.cache_controller import HiCacheController from sglang.srt.managers.cache_controller import HiCacheController
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache 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 ( from sglang.srt.mem_cache.memory_pool import (
MHATokenToKVPool, MHATokenToKVPool,
MLATokenToKVPool, MLATokenToKVPool,
ReqToTokenPool, 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.runtime_context import get_schedule
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils.common import ceil_align 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 self.page_size, (env_stride // self.page_size) * self.page_size
) )
kv_cache = self.token_to_kv_pool_allocator.get_kvcache() kv_cache = self.token_to_kv_pool_allocator.get_kvcache()
allocator_type = get_allocator_type(server_args) if not isinstance(kv_cache, (MHATokenToKVPool, MLATokenToKVPool)):
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:
raise ValueError("Unsupported KV cache type for decode offload") 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_group = tp_group
self.tp_world_size = torch.distributed.get_world_size(group=self.tp_group) self.tp_world_size = torch.distributed.get_world_size(group=self.tp_group)
+116 -111
View File
@@ -16,7 +16,6 @@ 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
@@ -38,6 +37,7 @@ from sglang.srt.layers.dp_attention import (
get_attention_dp_rank, get_attention_dp_rank,
is_dp_attention_enabled, 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.mem_cache.memory_pool import MLATokenToKVPool
from sglang.srt.runtime_context import get_parallel from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import get_device_module from sglang.srt.utils import get_device_module
@@ -47,26 +47,6 @@ 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
@@ -126,30 +106,66 @@ class CacheOperation:
device_indices: torch.Tensor, device_indices: torch.Tensor,
node_id: int, node_id: int,
priority: Optional[int] = None, priority: Optional[int] = None,
pool_transfers: Optional[List[PoolTransfer]] = None,
): ):
self.host_indices = host_indices self.host_indices = host_indices
self.device_indices = device_indices self.device_indices = device_indices
self.node_ids = [node_id] self.node_ids = [node_id]
self.data = None self.data = None
self.pool_transfers = pool_transfers
self.id = CacheOperation.counter self.id = CacheOperation.counter
CacheOperation.counter += 1 CacheOperation.counter += 1
# default priority is the order of creation # default priority is the order of creation
self.priority = priority if priority is not None else self.id 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 @staticmethod
def merge_ops(ops: List[CacheOperation]) -> CacheOperation: def merge_ops(ops: List[CacheOperation]) -> CacheOperation:
assert len(ops) > 0 assert ops
if len(ops) == 1: if len(ops) == 1:
return ops[0] return ops[0]
host_indices = torch.cat([op.host_indices for op in ops]) host_indices = torch.cat([op.host_indices for op in ops])
device_indices = torch.cat([op.device_indices for op in ops]) device_indices = torch.cat([op.device_indices for op in ops])
node_ids = [] node_ids = []
priority = min(op.priority for op in ops) priority = min(op.priority for op in ops)
for op in ops: for op in ops:
node_ids.extend(op.node_ids) 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 merged_op.node_ids = node_ids
return merged_op return merged_op
@@ -302,8 +318,7 @@ class HiCacheController:
self.ack_load_queue: List[HiCacheAck] = [] self.ack_load_queue: List[HiCacheAck] = []
self.ack_write_queue: List[HiCacheAck] = [] self.ack_write_queue: List[HiCacheAck] = []
self.write_stream = device_module.Stream() self.l2_transfer_engine = L2TransferEngine(io_backend)
self.load_stream = device_module.Stream()
# If a storage backend is provided at startup, treat it as an implicit attach, # 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. # so init/runtime share the same lifecycle semantics and code paths.
@@ -698,60 +713,21 @@ class HiCacheController:
return return
op = CacheOperation.merge_ops(self.write_queue) op = CacheOperation.merge_ops(self.write_queue)
# Kernel write-back keeps host indices on CPU only for page_first AND only host_indices, device_indices, pool_transfers = self._move_write_operation(op)
# 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
)
self.write_queue.clear() self.write_queue.clear()
start_event = device_module.Event() completion = self.l2_transfer_engine.submit_device_to_host(
ack_start_event, ack_finish_event, timing_enabled = make_timing_event_pair() self._l2_transfers(host_indices, device_indices, pool_transfers)
)
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)
self.ack_write_queue.append( self.ack_write_queue.append(
HiCacheAck( HiCacheAck(
start_event=ack_start_event, start_event=completion.start_event,
finish_event=ack_finish_event, finish_event=completion.finish_event,
node_ids=op.node_ids, node_ids=op.node_ids,
num_tokens=len(op.device_indices), num_tokens=len(op.device_indices),
timing_enabled=timing_enabled, timing_enabled=completion.timing_enabled,
num_tokens_by_pool={PoolName.KV.value: len(op.device_indices)}, num_tokens_by_pool=self._num_tokens_by_pool(op),
num_bytes=self._transfer_num_bytes(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 num_bytes += num_tokens * self.mem_pool_host_draft.size_per_token
return num_bytes return num_bytes
def _num_tokens_by_pool(self, op: CacheOperation) -> dict[str, int]:
return {PoolName.KV.value: len(op.device_indices)}
def load( def load(
self, self,
host_indices: torch.Tensor, host_indices: torch.Tensor,
@@ -781,7 +760,9 @@ class HiCacheController:
) )
return device_indices 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 # move indices to GPU if using kernels, to host if using direct indexing
if self.io_backend == "kernel": if self.io_backend == "kernel":
if not host_indices.is_cuda: if not host_indices.is_cuda:
@@ -803,58 +784,82 @@ class HiCacheController:
else: else:
raise ValueError(f"Unsupported io backend") 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: def start_loading(self) -> int:
if len(self.load_queue) == 0: if len(self.load_queue) == 0:
return -1 return -1
producer_id = self.layer_done_counter.update_producer() producer_id = self.layer_done_counter.update_producer()
op = CacheOperation.merge_ops(self.load_queue) op = CacheOperation.merge_ops(self.load_queue)
host_indices, device_indices = self.move_indices( host_indices, device_indices, pool_transfers = self._move_op_indices(op)
op.host_indices, op.device_indices
)
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() completion = self.l2_transfer_engine.submit_host_to_device(
self._l2_load_transfers(host_indices, device_indices, pool_transfers),
with device_module.stream(self.load_stream): start_event=producer_event.start_event,
producer_event.start_event.wait(self.load_stream) on_layer_done=producer_event.complete,
ack_start_event.record() layer_num=self.layer_num,
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)
self.ack_load_queue.append( self.ack_load_queue.append(
HiCacheAck( HiCacheAck(
start_event=ack_start_event, start_event=completion.start_event,
finish_event=ack_finish_event, finish_event=completion.finish_event,
node_ids=op.node_ids, node_ids=op.node_ids,
num_tokens=len(op.device_indices), num_tokens=len(op.device_indices),
timing_enabled=timing_enabled, timing_enabled=completion.timing_enabled,
num_tokens_by_pool={PoolName.KV.value: len(op.device_indices)}, num_tokens_by_pool=self._num_tokens_by_pool(op),
num_bytes=self._transfer_num_bytes(op), num_bytes=self._transfer_num_bytes(op),
) )
) )
@@ -26,6 +26,7 @@ from sglang.srt.observability.metrics_collector import (
from sglang.srt.runtime_context import get_observability from sglang.srt.runtime_context import get_observability
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.cache_controller import HiCacheController
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.mem_cache.unified_cache.cache_action import ( from sglang.srt.mem_cache.unified_cache.cache_action import (
@@ -233,6 +234,7 @@ class BasePrefixCache(ABC, PrefixCacheTrait):
metrics_collector: Optional[RadixCacheMetricsCollector] = ( metrics_collector: Optional[RadixCacheMetricsCollector] = (
None # metrics collector for the cache None # metrics collector for the cache
) )
cache_controller: Optional[HiCacheController] = None
def init_metrics_collector(self): def init_metrics_collector(self):
from sglang.srt.runtime_context import get_server_args from sglang.srt.runtime_context import get_server_args
@@ -5,14 +5,14 @@ import logging
import os import os
import threading import threading
import time import time
from dataclasses import replace
from queue import Empty, Queue from queue import Empty, Queue
from typing import TYPE_CHECKING, Any, Callable, List, Optional from typing import TYPE_CHECKING, Any, Callable, List, Optional
import torch import torch
from sglang.srt.managers.cache_controller import CacheOperation as BaseCacheOperation
from sglang.srt.managers.cache_controller import ( from sglang.srt.managers.cache_controller import (
HiCacheAck, CacheOperation,
) )
from sglang.srt.managers.cache_controller import ( from sglang.srt.managers.cache_controller import (
HiCacheController as BaseHiCacheController, HiCacheController as BaseHiCacheController,
@@ -23,9 +23,6 @@ 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,
@@ -33,75 +30,14 @@ from sglang.srt.mem_cache.hicache_storage import (
PoolTransfer, PoolTransfer,
PoolTransferResult, 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.memory_pool_host import HostPoolGroup, PoolEntry
from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost
from sglang.srt.utils import get_device_module
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
logger = logging.getLogger(__name__) 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): class StorageOperation(BaseStorageOperation):
@@ -403,69 +339,134 @@ class HybridCacheController(BaseHiCacheController):
self.start_writing() self.start_writing()
return host_indices return host_indices
def start_writing(self) -> None: def _move_op_indices(
if not self.write_queue: self, op: CacheOperation
return ) -> tuple[torch.Tensor, torch.Tensor, Optional[list[PoolTransfer]]]:
op = CacheOperation.merge_ops(self.write_queue) return self.move_hybrid_indices(op)
# Page-first staged write-back kernels need CPU destination host indices.
# A HostPoolGroup may mix staged and non-staged child pools, so let it def _move_write_operation(
# normalize indices per child instead of moving the whole operation here. self, op: CacheOperation
if ( ) -> tuple[torch.Tensor, torch.Tensor, Optional[list[PoolTransfer]]]:
self.io_backend == "kernel" host_group = self.mem_pool_host
and self.mem_pool_host.layout == "page_first" if self.io_backend != "kernel" or host_group.layout != "page_first":
and ( return self.move_hybrid_indices(op)
getattr(self.mem_pool_host, "can_use_write_back_jit", False) if not getattr(host_group, "supports_per_pool_backup_indices", False):
or getattr( if not getattr(host_group, "can_use_write_back_jit", False):
self.mem_pool_host, "supports_per_pool_backup_indices", 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):
host_indices = op.host_indices if getattr(host_pool, "can_use_write_back_jit", False):
device_indices = op.device_indices if host_indices.is_cuda:
resolved_pool_transfers = op.pool_transfers host_indices = host_indices.cpu()
else: return host_indices, device_indices
host_indices, device_indices, resolved_pool_transfers = ( return self.move_indices(host_indices, device_indices)
self.move_hybrid_indices(op)
) host_indices, device_indices = move_for_pool(
self.write_queue.clear() host_group.anchor_entry.host_pool,
start_event = device_module.Event() op.host_indices,
ack_start_event, ack_finish_event, timing_enabled = make_timing_event_pair() op.device_indices,
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),
)
) )
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]: def _num_tokens_by_pool(self, op: CacheOperation) -> dict[str, int]:
"""Per-pool token counts for a merged transfer op (anchor + extra """Per-pool token counts for a merged transfer op (anchor + extra
@@ -546,108 +547,6 @@ class HybridCacheController(BaseHiCacheController):
) )
return device_indices 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( def prefetch(
self, self,
request_id: str, request_id: str,
@@ -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( def build_kv_only_stack(
*, *,
params: CacheInitParams, params: CacheInitParams,
@@ -167,33 +281,15 @@ def build_kv_only_stack(
enable_storage_metrics: bool = False, enable_storage_metrics: bool = False,
) -> tuple[HostPoolGroup, HybridCacheController]: ) -> tuple[HostPoolGroup, HybridCacheController]:
transfer_layer_num = len(full_layer_mapping) transfer_layer_num = len(full_layer_mapping)
kv_host_pool = build_kv_host_pool( host_pool_group = build_kv_only_group(
kv_pool=kv_pool,
page_size=params.page_size, page_size=params.page_size,
server_args=server_args, server_args=server_args,
kv_pool=kv_pool,
full_layer_mapping=full_layer_mapping,
use_mla=use_mla, use_mla=use_mla,
override_kv_cache_dim=override_kv_cache_dim, override_kv_cache_dim=override_kv_cache_dim,
mtp_draft_device_pools=params.mtp_draft_device_pools, 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( cache_controller = HybridCacheController(
params.token_to_kv_pool_allocator, params.token_to_kv_pool_allocator,
host_pool_group, host_pool_group,
@@ -248,56 +344,22 @@ def build_hybrid_swa_stack(
server_args.hicache_size, (full_kv_pool, swa_kv_pool) server_args.hicache_size, (full_kv_pool, swa_kv_pool)
) )
kv_host_pool = build_kv_host_pool( host_pool_group = build_hybrid_swa_group(
kv_pool=full_kv_pool,
page_size=params.page_size, page_size=params.page_size,
server_args=server_args, 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, use_mla=use_mla,
host_size=kv_host_size, kv_host_size=kv_host_size,
pool_label="full", 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( cache_controller = HybridCacheController(
params.token_to_kv_pool_allocator, params.token_to_kv_pool_allocator,
host_pool_group, host_pool_group,
@@ -846,7 +908,7 @@ def build_anchor_sidecar_stack(
mtp_draft_device_pools=mtp_draft_device_pools, mtp_draft_device_pools=mtp_draft_device_pools,
) )
sidecar_host_pool = sidecar_host_pool_factory(kv_host_pool) 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: if mtp_draft_device_pools:
full_layer_mapping = _with_mtp_layer_mapping( full_layer_mapping = _with_mtp_layer_mapping(
full_layer_mapping, full_layer_mapping,
+127
View File
@@ -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)
@@ -1662,120 +1662,6 @@ class HostPoolGroup:
def set_from_flat_data_page(self, index: int, data_page) -> None: 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) 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): class DSAIndexerPoolHost(HostKVCache):
"""Host-side DSA index buffers only. Slot layout matches the anchor MLA host pool.""" """Host-side DSA index buffers only. Slot layout matches the anchor MLA host pool."""
@@ -16,12 +16,14 @@ register_cuda_ci(est_time=5, stage="base-b", runner_config="1-gpu-small")
class TestLoadBackDurationMetric(CustomTestCase): class TestLoadBackDurationMetric(CustomTestCase):
def setUp(self): def setUp(self):
from sglang.srt.managers import cache_controller as cc 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.cc = cc
self.transfer = transfer
def _completed_pair(self, payload_floats=1024 * 1024): 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) self.assertTrue(timing_enabled)
stream = torch.cuda.Stream() stream = torch.cuda.Stream()
start.record() start.record()
@@ -46,9 +48,11 @@ class TestLoadBackDurationMetric(CustomTestCase):
events.append(event) events.append(event)
return event return event
with patch.object(self.cc.device_module, "Event", side_effect=create_event): with patch.object(
self.cc._timing_events_supported.cache_clear() self.transfer.device_module, "Event", side_effect=create_event
start, finish, timing_enabled = self.cc.make_timing_event_pair() ):
self.transfer._timing_events_supported.cache_clear()
start, finish, timing_enabled = self.transfer.make_timing_event_pair()
self.assertFalse(timing_enabled) self.assertFalse(timing_enabled)
self.assertIs(start, events[0]) self.assertIs(start, events[0])
@@ -7,17 +7,17 @@ from unittest import mock
import torch import torch
from sglang.srt.managers import cache_controller as manager_cache_controller from sglang.srt.managers.cache_controller import CacheOperation, HiCacheController
from sglang.srt.managers.cache_controller import CacheOperation as ManagerCacheOperation from sglang.srt.mem_cache import l2_transfer as transfer_module
from sglang.srt.managers.cache_controller import ( from sglang.srt.mem_cache.hicache_storage import (
HiCacheController, 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 ( from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import (
CacheOperation,
HybridCacheController, HybridCacheController,
) )
from sglang.srt.mem_cache.l2_transfer import L2Transfer, L2TransferEngine
from sglang.srt.mem_cache.memory_pool_host import ( from sglang.srt.mem_cache.memory_pool_host import (
DeepSeekV4PagedHostPool, DeepSeekV4PagedHostPool,
DeepSeekV4StateHostPool, 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.mha import MHATokenToKVPoolHost
from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost
from sglang.test.ci.ci_register import register_cpu_ci 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") 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( def _cpu_staged_lf_pf_copy(
src_registry, src_registry,
*, *,
@@ -152,21 +180,204 @@ class _FakeEvent:
class _FakeDeviceModule: class _FakeDeviceModule:
Event = _FakeEvent Event = _FakeEvent
@staticmethod
def Stream():
return object()
@staticmethod @staticmethod
@contextmanager @contextmanager
def stream(stream): def stream(stream):
yield yield
class TestHiCacheStagedWriteBackDispatch(unittest.TestCase): class TestHiCacheStagedWriteBackDispatch(CustomTestCase):
def setUp(self): def setUp(self):
# start_writing probes timing support via a module-cached check; transfer_module._timing_events_supported.cache_clear()
# clear it on both sides so results from (or against) the fake self.addCleanup(transfer_module._timing_events_supported.cache_clear)
# device module never leak across tests.
manager_cache_controller._timing_events_supported.cache_clear()
def tearDown(self): @staticmethod
manager_cache_controller._timing_events_supported.cache_clear() 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): def _patched_transfers(self, src_registry=None, module=MEMORY_POOL_HOST_MODULE):
staged_side_effect = None staged_side_effect = None
@@ -701,26 +912,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
self.assertIsNone(group.destroy()) self.assertIsNone(group.destroy())
def test_write_back_jit_hybrid_write_keeps_extra_host_indices_on_cpu(self): def test_write_back_jit_hybrid_write_keeps_extra_host_indices_on_cpu(self):
captured = {} 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
controller = HybridCacheController.__new__(HybridCacheController) controller = HybridCacheController.__new__(HybridCacheController)
controller.write_queue = [ controller.write_queue = [
@@ -738,53 +930,25 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
) )
] ]
controller.io_backend = "kernel" 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.mem_pool_device = None
controller.has_draft = False controller.has_draft = False
controller.write_stream = object()
controller.ack_write_queue = [] controller.ack_write_queue = []
controller._record_transfer_indices_on_stream = lambda *args: None
controller.move_hybrid_indices = mock.Mock( controller.move_hybrid_indices = mock.Mock(
side_effect=AssertionError( side_effect=AssertionError(
"write-back JIT kernel write should not move indices" "write-back JIT kernel write should not move indices"
) )
) )
with ( self._start_writing(controller)
mock.patch.object(
hybrid_cache_controller, "device_module", _FakeDeviceModule
),
mock.patch.object(
manager_cache_controller, "device_module", _FakeDeviceModule
),
):
controller.start_writing()
controller.move_hybrid_indices.assert_not_called() controller.move_hybrid_indices.assert_not_called()
self.assertEqual(captured["host_indices"].device.type, "cpu") self.assertEqual([indices.device.type for indices in captured], ["cpu", "cpu"])
self.assertEqual(captured["pool_transfers"][0].host_indices.device.type, "cpu")
def test_hybrid_write_moves_indices_without_write_back_jit(self): def test_hybrid_write_moves_indices_without_write_back_jit(self):
captured = {} 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
op = CacheOperation( op = CacheOperation(
host_indices=_indices(0, 4), host_indices=_indices(0, 4),
@@ -801,29 +965,20 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
controller = HybridCacheController.__new__(HybridCacheController) controller = HybridCacheController.__new__(HybridCacheController)
controller.write_queue = [op] controller.write_queue = [op]
controller.io_backend = "kernel" 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.mem_pool_device = None
controller.has_draft = False controller.has_draft = False
controller.write_stream = object()
controller.ack_write_queue = [] controller.ack_write_queue = []
controller._record_transfer_indices_on_stream = lambda *args: None
controller.move_hybrid_indices = mock.Mock( controller.move_hybrid_indices = mock.Mock(
return_value=(op.host_indices, op.device_indices, op.pool_transfers) return_value=(op.host_indices, op.device_indices, op.pool_transfers)
) )
with ( self._start_writing(controller)
mock.patch.object(
hybrid_cache_controller, "device_module", _FakeDeviceModule
),
mock.patch.object(
manager_cache_controller, "device_module", _FakeDeviceModule
),
):
controller.start_writing()
controller.move_hybrid_indices.assert_called_once() controller.move_hybrid_indices.assert_called_once()
self.assertEqual(captured["host_indices"].device.type, "cpu") self.assertEqual([indices.device.type for indices in captured], ["cpu", "cpu"])
self.assertEqual(captured["pool_transfers"][0].host_indices.device.type, "cpu")
def test_write_back_jit_cache_controller_keeps_host_indices_on_cpu(self): def test_write_back_jit_cache_controller_keeps_host_indices_on_cpu(self):
captured = {} captured = {}
@@ -840,7 +995,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
controller = HiCacheController.__new__(HiCacheController) controller = HiCacheController.__new__(HiCacheController)
controller.write_queue = [ controller.write_queue = [
ManagerCacheOperation( CacheOperation(
host_indices=_indices(0, 4), host_indices=_indices(0, 4),
device_indices=_indices(4, 8), device_indices=_indices(4, 8),
node_id=1, node_id=1,
@@ -850,7 +1005,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
controller.mem_pool_host = FakeHostPool() controller.mem_pool_host = FakeHostPool()
controller.mem_pool_device = None controller.mem_pool_device = None
controller.has_draft = False controller.has_draft = False
controller.write_stream = object() controller.device = "cuda"
controller.ack_write_queue = [] controller.ack_write_queue = []
controller.move_indices = mock.Mock( controller.move_indices = mock.Mock(
side_effect=AssertionError( side_effect=AssertionError(
@@ -858,10 +1013,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
) )
) )
with mock.patch.object( self._start_writing(controller)
manager_cache_controller, "device_module", _FakeDeviceModule
):
controller.start_writing()
controller.move_indices.assert_not_called() controller.move_indices.assert_not_called()
self.assertEqual(captured["host_indices"].device.type, "cpu") self.assertEqual(captured["host_indices"].device.type, "cpu")
@@ -879,7 +1031,7 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
): ):
captured["host_indices"] = host_indices captured["host_indices"] = host_indices
op = ManagerCacheOperation( op = CacheOperation(
host_indices=_indices(0, 4), host_indices=_indices(0, 4),
device_indices=_indices(4, 8), device_indices=_indices(4, 8),
node_id=1, node_id=1,
@@ -890,16 +1042,13 @@ class TestHiCacheStagedWriteBackDispatch(unittest.TestCase):
controller.mem_pool_host = FakeHostPool() controller.mem_pool_host = FakeHostPool()
controller.mem_pool_device = None controller.mem_pool_device = None
controller.has_draft = False controller.has_draft = False
controller.write_stream = object() controller.device = "cuda"
controller.ack_write_queue = [] controller.ack_write_queue = []
controller.move_indices = mock.Mock( controller.move_indices = mock.Mock(
return_value=(op.host_indices, op.device_indices) return_value=(op.host_indices, op.device_indices)
) )
with mock.patch.object( self._start_writing(controller)
manager_cache_controller, "device_module", _FakeDeviceModule
):
controller.start_writing()
controller.move_indices.assert_called_once() controller.move_indices.assert_called_once()
self.assertEqual(captured["host_indices"].device.type, "cpu") self.assertEqual(captured["host_indices"].device.type, "cpu")
@@ -148,7 +148,6 @@ _EXPOSED = {
("disaggregation/common/conn.py", "disaggregation_bootstrap_port"), ("disaggregation/common/conn.py", "disaggregation_bootstrap_port"),
("disaggregation/common/conn.py", "pp_size"), ("disaggregation/common/conn.py", "pp_size"),
("disaggregation/decode_kvcache_offload_manager.py", "hicache_io_backend"), ("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/decode_kvcache_offload_manager.py", "served_model_name"),
("disaggregation/encode_receiver.py", "disaggregation_ib_device"), ("disaggregation/encode_receiver.py", "disaggregation_ib_device"),
("disaggregation/encode_receiver.py", "encoder_transfer_backend"), ("disaggregation/encode_receiver.py", "encoder_transfer_backend"),