feat: map CP HiCache writes to owned physical slots

This commit is contained in:
2026-05-08 01:14:58 +08:00
parent eb0ec4299d
commit 301420f976
2 changed files with 288 additions and 2 deletions

View File

@@ -16,6 +16,7 @@ limitations under the License.
import logging
import threading
import time
from dataclasses import dataclass
from queue import Empty, Full, Queue
from typing import TYPE_CHECKING, List, NamedTuple, Optional
@@ -40,7 +41,8 @@ from sglang.srt.layers.dp_attention import (
get_attention_tp_size,
is_dp_attention_enabled,
)
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool, NSATokenToKVPool
from sglang.srt.utils import get_device_module
logger = logging.getLogger(__name__)
@@ -244,6 +246,18 @@ class PrefetchOperation(StorageOperation):
return self._terminated_flag
@dataclass
class HiCacheWriteResult:
metadata: object
required_host_slots: int = 0
@dataclass
class HiCacheWriteFailure:
required_host_slots: int
metadata: object = None
class HiCacheController:
def __init__(
@@ -262,6 +276,7 @@ class HiCacheController:
pp_rank: int = 0,
pp_size: int = 1,
enable_storage_metrics: bool = False,
cp_shared_kv_layout: Optional[CpSharedKVLayout] = None,
):
self.tp_group = tp_group
self.mem_pool_device_allocator = token_to_kv_pool_allocator
@@ -271,6 +286,14 @@ class HiCacheController:
if isinstance(mem_pool_device, HybridLinearKVPool):
mem_pool_device = mem_pool_device.full_kv_pool
self.mem_pool_device = mem_pool_device
self.cp_shared_kv_layout = cp_shared_kv_layout
self.uses_cp_hicache = cp_shared_kv_layout is not None
if self.uses_cp_hicache and not isinstance(
self.mem_pool_device, NSATokenToKVPool
):
raise ValueError(
"CP shared KV HiCache host integration requires NSATokenToKVPool."
)
self.mem_pool_host = mem_pool_host
self.write_policy = write_policy
self.page_size = page_size
@@ -656,10 +679,13 @@ class HiCacheController:
device_indices: torch.Tensor,
priority: Optional[int] = None,
node_id: int = -1,
) -> Optional[torch.Tensor]:
) -> Optional[torch.Tensor | HiCacheWriteResult | HiCacheWriteFailure]:
"""
Back up KV caches from device memory to host memory.
"""
if self.uses_cp_hicache:
return self._write_cp(device_indices, priority, node_id)
host_indices = self.mem_pool_host.alloc(len(device_indices))
if host_indices is None:
return None
@@ -697,6 +723,56 @@ class HiCacheController:
self.ack_write_queue.append(HiCacheAck(start_event, finish_event, op.node_ids))
def _append_completed_write_ack(self, node_id: int) -> None:
event = device_module.Event()
event.record()
self.ack_write_queue.append(HiCacheAck(event, event, [node_id]))
def _append_completed_load_ack(self, node_id: int) -> None:
event = device_module.Event()
event.record()
self.ack_load_queue.append(HiCacheAck(event, event, [node_id]))
def _write_cp(
self,
device_indices: torch.Tensor,
priority: Optional[int] = None,
node_id: int = -1,
) -> HiCacheWriteResult | HiCacheWriteFailure:
from sglang.srt.mem_cache.hiradix_cache import CpHiCacheNodeMetadata
layout = self.cp_shared_kv_layout
owned_mask = layout.owned_by_this_rank(device_indices)
owned_positions = owned_mask.nonzero(as_tuple=True)[0].cpu()
logical_len = len(device_indices)
if owned_positions.numel() == 0:
self._append_completed_write_ack(node_id)
return HiCacheWriteResult(
metadata=CpHiCacheNodeMetadata(
logical_len=logical_len,
owned_positions=owned_positions,
host_indices=torch.empty((0,), dtype=torch.int64),
)
)
owned_logical_indices = device_indices[owned_mask]
physical_device_indices = layout.logical_locs_to_physical(owned_logical_indices)
host_indices = self.mem_pool_host.alloc(len(physical_device_indices))
if host_indices is None:
return HiCacheWriteFailure(required_host_slots=len(physical_device_indices))
self.write_queue.append(
CacheOperation(host_indices, physical_device_indices, node_id, priority)
)
self.start_writing()
return HiCacheWriteResult(
metadata=CpHiCacheNodeMetadata(
logical_len=logical_len,
owned_positions=owned_positions,
host_indices=host_indices.cpu(),
)
)
def load(
self,
host_indices: torch.Tensor,