refactor(hicache): flatten L2 transfer execution (#34793)
GB300 test fails unrelated
This commit is contained in:
@@ -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)
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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"),
|
||||||
|
|||||||
Reference in New Issue
Block a user