feat: map CP HiCache writes to owned physical slots
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user