Keep CP HiCache host visibility behind per-layer target and draft backup

CP shared KV needs HiCache backup to overlap with layer execution without exposing partially copied host state. Split CP backup into reservation, pending radix state, per-layer target/draft D2H submission, and one final ack-driven visibility commit. The all-layer path remains available only as an explicit fallback and now logs a warning when used.

Constraint: CP shared KV owner-lane metadata and draft/MTP KV must stay strongly synchronized with target KV.
Constraint: Local CUDA tests are disallowed; CUDA verification was run only in the g0034 container.
Rejected: Let target layer hooks copy draft KV too | draft may not have stored that layer yet, which can corrupt MTP accept behavior.
Rejected: Silent all-layer fallback | it hides performance regressions and makes ETE logs ambiguous.
Confidence: medium
Scope-risk: broad
Directive: Reserved or partially copied host payloads must remain invisible until final ack commits pending_host_backups.
Tested: g0034 docker /sgl-workspace/sglang-tai PYTHONPATH=/mnt/beegfs/cjy/tai-kernel/python:python python -m pytest test/registered/unit/managers/test_hicache_controller_cp.py -q -> 49 passed.
Tested: g0034 docker /sgl-workspace/sglang-tai PYTHONPATH=/mnt/beegfs/cjy/tai-kernel/python:python python -m pytest test/registered/unit/mem_cache/test_cp_hicache_metadata.py -q -> 58 passed.
Tested: g0034 docker /mnt/beegfs/cjy/tai-kernel PYTHONPATH=python python -m pytest tests/nsa_prefill/test_kvcacheio_lf_pf.py -q -> 7 passed.
Not-tested: Long-running GLM5 CP+HiCache+MTP ETE throughput and host-pressure soak.
This commit is contained in:
laoyao0822
2026-05-27 07:43:10 +08:00
parent 367dff06f3
commit 03529319a1
11 changed files with 2104 additions and 72 deletions

View File

@@ -16,9 +16,9 @@ limitations under the License.
import logging
import threading
import time
from dataclasses import dataclass
from dataclasses import dataclass, field
from queue import Empty, Full, Queue
from typing import TYPE_CHECKING, List, NamedTuple, Optional
from typing import TYPE_CHECKING, Dict, List, NamedTuple, Optional, Set
import torch
@@ -266,6 +266,32 @@ class HiCacheWriteFailure:
metadata: object = None
@dataclass
class HiCacheWriteReservation:
metadata: object
host_indices: torch.Tensor
physical_device_indices: torch.Tensor
node_id: int = -1
priority: Optional[int] = None
draft_host_indices: Optional[torch.Tensor] = None
required_host_slots: int = 0
@dataclass
class HiCacheLayerWriteState:
reservation: HiCacheWriteReservation
total_layers: int
start_event: object
finish_event: object
layer_events: List[object] = field(default_factory=list)
completed_target_layers: Set[int] = field(default_factory=set)
completed_draft_layers: Set[int] = field(default_factory=set)
host_indices: Optional[torch.Tensor] = None
physical_device_indices: Optional[torch.Tensor] = None
draft_host_indices: Optional[torch.Tensor] = None
ack_appended: bool = False
class HiCacheController:
def __init__(
@@ -333,6 +359,10 @@ class HiCacheController:
self.layer_num = self.mem_pool_device.layer_num
self.layer_done_counter = LayerDoneCounter(self.layer_num)
self.mem_pool_device.register_layer_transfer_counter(self.layer_done_counter)
if hasattr(self.mem_pool_device, "register_layer_backup_notifier"):
self.mem_pool_device.register_layer_backup_notifier(
lambda layer_id: self.on_layer_kv_stored(layer_id, source="target")
)
self.draft_mem_pool_host = None
self.draft_mem_pool_device = None
@@ -351,6 +381,7 @@ class HiCacheController:
self.draft_write_queue: List[CacheOperation] = []
self.ack_load_queue: List[HiCacheAck] = []
self.ack_write_queue: List[HiCacheAck] = []
self.pending_layer_writes: Dict[int, HiCacheLayerWriteState] = {}
if draft_mem_pool_host is not None or draft_mem_pool_device is not None:
self.attach_draft_pool(draft_mem_pool_device, draft_mem_pool_host)
@@ -398,6 +429,28 @@ class HiCacheController:
draft_mem_pool_device.register_layer_transfer_counter(
self.layer_done_counter
)
if hasattr(draft_mem_pool_device, "register_layer_backup_notifier"):
draft_mem_pool_device.register_layer_backup_notifier(
lambda layer_id: self.on_layer_kv_stored(layer_id, source="draft")
)
def on_layer_kv_stored(self, layer_id: int, source: str = "target") -> None:
if not self.pending_layer_writes:
return
if source not in ("target", "draft"):
raise ValueError(f"Unknown CP HiCache layer backup source: {source}")
reservations = [
state.reservation for state in list(self.pending_layer_writes.values())
]
for reservation in reservations:
if reservation.node_id not in self.pending_layer_writes:
continue
self.submit_write_cp_layer(
reservation,
layer_id,
submit_target=(source == "target"),
submit_draft=(source == "draft"),
)
def _start_storage_threads(self):
"""Start storage prefetch/backup threads and their queues.
@@ -703,6 +756,7 @@ class HiCacheController:
self.load_buffer.clear()
self.ack_write_queue.clear()
self.ack_load_queue.clear()
self.pending_layer_writes.clear()
if self.enable_storage:
self.prefetch_thread.join()
self.backup_thread.join()
@@ -859,12 +913,12 @@ class HiCacheController:
device_indices, self.page_size, "physical_device_indices"
)
def _write_cp(
def reserve_write_cp(
self,
device_indices: torch.Tensor,
priority: Optional[int] = None,
node_id: int = -1,
) -> HiCacheWriteResult | HiCacheWriteFailure:
) -> HiCacheWriteReservation | HiCacheWriteFailure:
from sglang.srt.mem_cache.hiradix_cache import CpHiCacheNodeMetadata
layout = self.cp_shared_kv_layout
@@ -882,32 +936,31 @@ class HiCacheController:
f"logical_len={logical_len} page_size={page_size}"
)
page_first_locs = device_indices[::page_size]
logical_pages = torch.div(
page_first_locs, page_size, rounding_mode="floor"
)
logical_pages = torch.div(page_first_locs, page_size, rounding_mode="floor")
page_owners = layout.owner_for_logical_pages(logical_pages).to(
dtype=torch.int8, device="cpu"
)
if owned_positions.numel() == 0:
self._append_completed_write_ack(node_id)
logger.info(
"[CacheCtrl-write] _write_cp zero-owned rank: node_id=%d logical_len=%d",
"[CacheCtrl-write] reserve_write_cp zero-owned rank: node_id=%d logical_len=%d",
node_id,
logical_len,
)
return HiCacheWriteResult(
empty = torch.empty((0,), dtype=torch.int64)
return HiCacheWriteReservation(
metadata=CpHiCacheNodeMetadata(
logical_len=logical_len,
owned_positions=owned_positions,
host_indices=torch.empty((0,), dtype=torch.int64),
host_indices=empty,
page_owners=page_owners,
page_size=page_size,
draft_host_indices=(
torch.empty((0,), dtype=torch.int64)
if self.has_draft_hicache
else None
),
)
draft_host_indices=(empty.clone() if self.has_draft_hicache else None),
),
host_indices=empty,
physical_device_indices=empty,
node_id=node_id,
priority=priority,
draft_host_indices=(empty.clone() if self.has_draft_hicache else None),
)
owned_logical_indices = device_indices[owned_mask]
@@ -915,12 +968,13 @@ class HiCacheController:
host_indices = self.mem_pool_host.alloc(len(physical_device_indices))
if host_indices is None:
logger.info(
"[CacheCtrl-write] _write_cp FAILED (host full): node_id=%d logical_len=%d owned=%d",
"[CacheCtrl-write] reserve_write_cp FAILED (host full): node_id=%d logical_len=%d owned=%d",
node_id,
logical_len,
owned_positions.numel(),
)
return HiCacheWriteFailure(required_host_slots=len(physical_device_indices))
draft_host_indices = None
if self.has_draft_hicache:
draft_host_indices = self.draft_mem_pool_host.alloc(
@@ -929,7 +983,7 @@ class HiCacheController:
if draft_host_indices is None:
self.mem_pool_host.free(host_indices)
logger.info(
"[CacheCtrl-write] _write_cp FAILED (draft host full): node_id=%d logical_len=%d owned=%d",
"[CacheCtrl-write] reserve_write_cp FAILED (draft host full): node_id=%d logical_len=%d owned=%d",
node_id,
logical_len,
owned_positions.numel(),
@@ -952,25 +1006,7 @@ class HiCacheController:
self.draft_mem_pool_host.free(draft_host_indices)
raise
self.write_queue.append(
CacheOperation(host_indices, physical_device_indices, node_id, priority)
)
if draft_host_indices is not None:
self.draft_write_queue.append(
CacheOperation(
draft_host_indices, physical_device_indices, node_id, priority
)
)
self.start_writing()
logger.info(
"[CacheCtrl-write] _write_cp submitted: node_id=%d logical_len=%d owned=%d physical=%d draft=%s",
node_id,
logical_len,
owned_positions.numel(),
len(physical_device_indices),
draft_host_indices is not None,
)
return HiCacheWriteResult(
return HiCacheWriteReservation(
metadata=CpHiCacheNodeMetadata(
logical_len=logical_len,
owned_positions=owned_positions,
@@ -980,8 +1016,255 @@ class HiCacheController:
draft_host_indices=(
draft_host_indices.cpu() if draft_host_indices is not None else None
),
),
host_indices=host_indices,
physical_device_indices=physical_device_indices,
node_id=node_id,
priority=priority,
draft_host_indices=draft_host_indices,
)
def submit_write_cp_all_layer(self, reservation: HiCacheWriteReservation) -> None:
logger.warning(
"[CacheCtrl-write] CP HiCache all-layer backup fallback: node_id=%d logical_len=%d owned=%d physical=%d draft=%s",
reservation.node_id,
reservation.metadata.logical_len,
reservation.metadata.owned_positions.numel(),
len(reservation.physical_device_indices),
reservation.draft_host_indices is not None,
)
if len(reservation.physical_device_indices) == 0:
self._append_completed_write_ack(reservation.node_id)
logger.info(
"[CacheCtrl-write] submit_write_cp_all_layer zero-owned ack: node_id=%d logical_len=%d",
reservation.node_id,
reservation.metadata.logical_len,
)
return
self.write_queue.append(
CacheOperation(
reservation.host_indices,
reservation.physical_device_indices,
reservation.node_id,
reservation.priority,
)
)
if reservation.draft_host_indices is not None:
self.draft_write_queue.append(
CacheOperation(
reservation.draft_host_indices,
reservation.physical_device_indices,
reservation.node_id,
reservation.priority,
)
)
self.start_writing()
logger.info(
"[CacheCtrl-write] submit_write_cp_all_layer submitted: node_id=%d logical_len=%d owned=%d physical=%d draft=%s",
reservation.node_id,
reservation.metadata.logical_len,
reservation.metadata.owned_positions.numel(),
len(reservation.physical_device_indices),
reservation.draft_host_indices is not None,
)
def _get_or_create_layer_write_state(
self, reservation: HiCacheWriteReservation
) -> HiCacheLayerWriteState:
state = self.pending_layer_writes.get(reservation.node_id)
if state is not None:
if state.reservation is not reservation:
raise RuntimeError(
f"Conflicting CP HiCache per-layer backup reservation for "
f"node_id={reservation.node_id}"
)
return state
draft_layer_num = (
self.draft_mem_pool_device.layer_num
if reservation.draft_host_indices is not None
else 0
)
total_layers = max(self.layer_num, draft_layer_num)
if total_layers <= 0:
raise RuntimeError(
f"Invalid CP HiCache per-layer backup layer count: {total_layers}"
)
state = HiCacheLayerWriteState(
reservation=reservation,
total_layers=total_layers,
start_event=device_module.Event(),
finish_event=device_module.Event(),
)
if len(reservation.physical_device_indices) > 0:
op = CacheOperation(
reservation.host_indices,
reservation.physical_device_indices,
reservation.node_id,
reservation.priority,
)
state.host_indices, state.physical_device_indices = self.move_indices(
op, self.mem_pool_host
)
if reservation.draft_host_indices is not None:
draft_op = CacheOperation(
reservation.draft_host_indices,
reservation.physical_device_indices,
reservation.node_id,
reservation.priority,
)
state.draft_host_indices, _ = self.move_indices(
draft_op, self.draft_mem_pool_host
)
self.pending_layer_writes[reservation.node_id] = state
state.start_event.record()
return state
def submit_write_cp_layer(
self,
reservation: HiCacheWriteReservation,
layer_id: int,
*,
submit_target: bool = True,
submit_draft: bool = True,
) -> None:
"""Submit one layer of a reserved CP host backup.
Target and draft D2H are still one logical operation: this method only
queues a final write ack when all target/draft layers have been submitted.
"""
state = self._get_or_create_layer_write_state(reservation)
if layer_id < 0 or layer_id >= state.total_layers:
raise ValueError(
f"layer_id={layer_id} is outside CP HiCache backup layer range "
f"[0, {state.total_layers}) for node_id={reservation.node_id}"
)
needs_target = (
submit_target
and layer_id < self.layer_num
and layer_id not in state.completed_target_layers
)
needs_draft = (
submit_draft
and state.draft_host_indices is not None
and layer_id < self.draft_mem_pool_device.layer_num
and layer_id not in state.completed_draft_layers
)
if not needs_target and not needs_draft:
return
layer_event = device_module.Event()
layer_event.record()
state.layer_events.append(layer_event)
if len(reservation.physical_device_indices) > 0:
with device_module.stream(self.write_stream):
state.start_event.wait(self.write_stream)
layer_event.wait(self.write_stream)
if needs_target:
self.mem_pool_host.backup_from_device_per_layer(
self.mem_pool_device,
state.host_indices,
state.physical_device_indices,
layer_id,
self.io_backend,
)
if needs_draft:
self.draft_mem_pool_host.backup_from_device_per_layer(
self.draft_mem_pool_device,
state.draft_host_indices,
state.physical_device_indices,
layer_id,
self.io_backend,
)
if needs_target:
state.completed_target_layers.add(layer_id)
if needs_draft:
state.completed_draft_layers.add(layer_id)
target_done = len(state.completed_target_layers) >= self.layer_num
draft_done = (
state.draft_host_indices is None
or len(state.completed_draft_layers)
>= self.draft_mem_pool_device.layer_num
)
if not target_done or not draft_done:
return
if state.ack_appended:
return
with device_module.stream(self.write_stream):
state.finish_event.record()
if state.host_indices is not None and state.host_indices.is_cuda:
state.host_indices.record_stream(self.write_stream)
if (
state.physical_device_indices is not None
and state.physical_device_indices.is_cuda
):
state.physical_device_indices.record_stream(self.write_stream)
if (
state.draft_host_indices is not None
and state.draft_host_indices.is_cuda
):
state.draft_host_indices.record_stream(self.write_stream)
state.ack_appended = True
self.pending_layer_writes.pop(reservation.node_id, None)
self.ack_write_queue.append(
HiCacheAck(state.start_event, state.finish_event, [reservation.node_id])
)
logger.info(
"[CacheCtrl-write] submit_write_cp_layer final ack: node_id=%d logical_len=%d layers=%d draft=%s",
reservation.node_id,
reservation.metadata.logical_len,
state.total_layers,
reservation.draft_host_indices is not None,
)
def submit_write_cp_per_layer(
self,
reservation: HiCacheWriteReservation,
*,
catch_up_all_layers: bool = True,
) -> None:
"""Register a CP write reservation for per-layer host backup.
Current radix insertion calls write_backup after the request KV has already
been materialized. For that path, catch up by submitting all layers
immediately through the per-layer API. Future early reservations can use
catch_up_all_layers=False and rely on on_layer_kv_stored().
"""
state = self._get_or_create_layer_write_state(reservation)
if not catch_up_all_layers:
logger.info(
"[CacheCtrl-write] submit_write_cp_per_layer registered: node_id=%d logical_len=%d layers=%d draft=%s",
reservation.node_id,
reservation.metadata.logical_len,
state.total_layers,
reservation.draft_host_indices is not None,
)
return
total_layers = state.total_layers
for layer_id in range(total_layers):
self.submit_write_cp_layer(reservation, layer_id)
def _write_cp(
self,
device_indices: torch.Tensor,
priority: Optional[int] = None,
node_id: int = -1,
) -> HiCacheWriteResult | HiCacheWriteFailure:
reservation = self.reserve_write_cp(device_indices, priority, node_id)
if isinstance(reservation, HiCacheWriteFailure):
return reservation
self.submit_write_cp_per_layer(reservation)
return HiCacheWriteResult(metadata=reservation.metadata)
def set_draft_kv_pool(self, draft_device_pool, draft_host_pool) -> None:
"""Register draft KV pools so L2 ops piggyback draft transfers."""

View File

@@ -643,6 +643,7 @@ class Req(ReqDllmMixin):
self.last_host_node: Any = None
self.last_host_backup_node: Any = None
self.host_hit_length = 0
self.prefix_match_deferred_by_pending_backup = False
# Tokens loaded from storage backend (L3) during prefetch for this request
self.storage_hit_length = 0
# The node to lock until for swa radix tree lock ref
@@ -930,6 +931,9 @@ class Req(ReqDllmMixin):
self.cache_protected_len = match_result.cache_protected_len
else:
self.cache_protected_len = len(self.prefix_indices)
self.prefix_match_deferred_by_pending_backup = (
match_result.pending_backup_deferred_node is not None
)
if self.is_dllm():
self._update_block_offset_for_dllm()

View File

@@ -212,6 +212,9 @@ class SchedulePolicy:
match_result.last_host_node,
match_result.host_hit_length,
)
r.prefix_match_deferred_by_pending_backup = (
match_result.pending_backup_deferred_node is not None
)
# NOTE(sang): This logic is for in-batch prefix caching;
# If there are more than 1 request that have small matching prefix from
@@ -748,6 +751,9 @@ class PrefillAdder:
def add_one_req(
self, req: Req, has_chunked_req: bool, truncation_align_size: Optional[int]
):
if getattr(req, "prefix_match_deferred_by_pending_backup", False):
return AddReqResult.OTHER
if (self.prefill_delayer_single_pass is not None) and (
not self.prefill_delayer_single_pass.negotiate_should_allow_prefill(
local_prefillable=True,

View File

@@ -144,6 +144,7 @@ class MatchResult(NamedTuple):
host_hit_length: int = 0
mamba_branching_seqlen: Optional[int] = None
cache_protected_len: Optional[int] = None
pending_backup_deferred_node: Any = None
class BasePrefixCache(ABC, PrefixCacheTrait):

View File

@@ -14,7 +14,11 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Tuple
import torch
from sglang.srt.environ import envs
from sglang.srt.managers.cache_controller import HiCacheController, PrefetchOperation
from sglang.srt.managers.cache_controller import (
HiCacheController,
HiCacheWriteFailure,
PrefetchOperation,
)
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams,
@@ -232,6 +236,24 @@ class CpHiCacheNodeMetadata:
)
@dataclass
class PendingHiCacheBackup:
node: TreeNode
metadata: CpHiCacheNodeMetadata
logical_len: int
submitted: bool = False
local_done: bool = False
locked: bool = True
class HiCachePendingBackupSplit(Exception):
def __init__(self, node: TreeNode):
self.node = node
super().__init__(
f"Cannot split node_id={getattr(node, 'id', None)} while CP HiCache backup is pending"
)
class HiRadixCache(RadixCache):
def _create_token_to_kv_pool_host(
@@ -433,6 +455,8 @@ class HiRadixCache(RadixCache):
# record the nodes with ongoing write through
self.ongoing_write_through = {}
# record CP host reservations/backups that are not request-visible yet.
self.pending_host_backups: Dict[int, PendingHiCacheBackup] = {}
# record the node segments with ongoing load back
self.ongoing_load_back = {}
# record the ongoing prefetch requests
@@ -886,6 +910,8 @@ class HiRadixCache(RadixCache):
self.cache_controller.clear_draft_host_pool()
# Clear per-request tracking dicts
self.prefetch_loaded_tokens_by_reqid.clear()
if hasattr(self, "pending_host_backups"):
self.pending_host_backups.clear()
self.evictable_host_leaves.clear()
self.pinned_size_ = 0
super().reset()
@@ -945,7 +971,10 @@ class HiRadixCache(RadixCache):
return True
def _node_host_write_pending(self, node: TreeNode) -> bool:
return getattr(node, "id", None) in getattr(self, "ongoing_write_through", {})
node_id = getattr(node, "id", None)
return node_id in getattr(self, "ongoing_write_through", {}) or node_id in getattr(
self, "pending_host_backups", {}
)
def _node_host_write_ready(self, node: TreeNode) -> bool:
return not self._node_host_write_pending(node)
@@ -960,6 +989,27 @@ class HiRadixCache(RadixCache):
)
return node.host_value is not None
def _commit_pending_backup(self, node_id: int) -> TreeNode:
pending = self.pending_host_backups.pop(node_id)
node = pending.node
node.host_len = pending.logical_len
node.cp_hicache = pending.metadata
node.host_value = None
if pending.locked:
self.dec_node_lock_ref(node)
return node
def _rollback_pending_backup(self, node_id: int) -> TreeNode:
pending = self.pending_host_backups.pop(node_id)
node = pending.node
self.cache_controller.evict_cp_host(pending.metadata)
if node.cp_hicache is pending.metadata:
node.cp_hicache = None
node.host_len = 0
if pending.locked:
self.dec_node_lock_ref(node)
return node
def _node_host_len(self, node: TreeNode) -> int:
if self._uses_cp_hicache:
return node.host_len
@@ -979,30 +1029,26 @@ class HiRadixCache(RadixCache):
write_back,
)
if self._uses_cp_hicache:
result = self.cache_controller.write(
result = self.cache_controller.reserve_write_cp(
device_indices=node.value,
node_id=node.id,
)
metadata = getattr(result, "metadata", None)
required_host_slots = 0
if metadata is None:
if isinstance(result, HiCacheWriteFailure):
required_host_slots = result.required_host_slots
self._evict_host_for_physical_slots(
required_host_slots,
synchronize_across_ranks=getattr(self, "tp_world_size", 1) > 1,
)
if metadata is None:
self._evict_host_for_physical_slots(
required_host_slots,
synchronize_across_ranks=getattr(self, "tp_world_size", 1) > 1,
)
logger.info(
"[HiCache-write] write_backup CP retry after host eviction: node_id=%d needed_slots=%d",
node.id,
required_host_slots,
)
result = self.cache_controller.write(
result = self.cache_controller.reserve_write_cp(
device_indices=node.value,
node_id=node.id,
)
metadata = getattr(result, "metadata", None)
if metadata is None:
if isinstance(result, HiCacheWriteFailure):
logger.info(
"[HiCache-write] write_backup CP FAILED (host full): node_id=%d len=%d",
node.id,
@@ -1010,18 +1056,29 @@ class HiRadixCache(RadixCache):
)
return 0
node.host_len = len(node.value)
node.cp_hicache = metadata
node.host_value = None
self.ongoing_write_through[node.id] = node
if not write_back:
self.inc_node_lock_ref(node)
self.pending_host_backups[node.id] = PendingHiCacheBackup(
node=node,
metadata=result.metadata,
logical_len=len(node.value),
submitted=True,
locked=not write_back,
)
try:
self.cache_controller.submit_write_cp_per_layer(result)
except Exception:
self.ongoing_write_through.pop(node.id, None)
self._rollback_pending_backup(node.id)
raise
logger.info(
"[HiCache-write] write_backup CP SUCCESS: node_id=%d host_len=%d owned_positions=%d ongoing_writes=%d",
"[HiCache-write] write_backup CP SUBMITTED: node_id=%d logical_len=%d owned_positions=%d ongoing_writes=%d pending_backups=%d",
node.id,
node.host_len,
node.cp_hicache.owned_positions.numel(),
len(node.value),
result.metadata.owned_positions.numel(),
len(self.ongoing_write_through),
len(self.pending_host_backups),
)
return len(node.value)
@@ -1102,13 +1159,22 @@ class HiRadixCache(RadixCache):
self.write_backup(node)
def writing_check(self, write_back=False):
def complete_write_ack_node(ack_id: int) -> TreeNode:
if ack_id in self.pending_host_backups:
backuped_node = self._commit_pending_backup(ack_id)
self.ongoing_write_through.pop(ack_id, None)
return backuped_node
backuped_node = self.ongoing_write_through.pop(ack_id)
self.dec_node_lock_ref(backuped_node)
return backuped_node
if write_back:
# blocking till all write back complete
while len(self.ongoing_write_through) > 0:
for _, finish_event, ack_list in self.cache_controller.ack_write_queue:
finish_event.synchronize()
for ack_id in ack_list:
backuped_node = self.ongoing_write_through.pop(ack_id)
backuped_node = complete_write_ack_node(ack_id)
if self.enable_storage:
self.write_backup_storage(backuped_node)
self.cache_controller.ack_write_queue.clear()
@@ -1148,9 +1214,8 @@ class HiRadixCache(RadixCache):
_, finish_event, ack_list = self.cache_controller.ack_write_queue.pop(0)
finish_event.synchronize()
for ack_id in ack_list:
backuped_node = self.ongoing_write_through.pop(ack_id)
backuped_node = complete_write_ack_node(ack_id)
released_nodes.append(ack_id)
self.dec_node_lock_ref(backuped_node)
if self.enable_storage:
self.write_backup_storage(backuped_node)
finish_count -= 1
@@ -1985,7 +2050,13 @@ class HiRadixCache(RadixCache):
host_hit_length=0,
)
value, last_node = self._match_prefix_helper(self.root_node, key)
deferred_node = None
try:
value, last_node = self._match_prefix_helper(self.root_node, key)
except HiCachePendingBackupSplit as exc:
value = []
last_node = exc.node.parent if exc.node.parent is not None else self.root_node
deferred_node = exc.node
if value:
value = torch.cat(value)
else:
@@ -2018,6 +2089,7 @@ class HiRadixCache(RadixCache):
last_device_node=last_node,
last_host_node=last_host_node,
host_hit_length=host_hit_length,
pending_backup_deferred_node=deferred_node,
)
def prefetch_from_storage(
@@ -2121,6 +2193,11 @@ class HiRadixCache(RadixCache):
child.pin_expiry = time.monotonic() + child.pin_ttl
prefix_len = self.key_match_fn(child.key, key)
if prefix_len < len(child.key):
if (
self._uses_cp_hicache
and child.id in getattr(self, "pending_host_backups", {})
):
raise HiCachePendingBackupSplit(child)
new_node = self._split_node(child.key, child, prefix_len)
if not new_node.evicted:
value.append(new_node.value)
@@ -2138,6 +2215,11 @@ class HiRadixCache(RadixCache):
return value, node
def _split_node(self, key: RadixKey, child: TreeNode, split_len: int):
if (
self._uses_cp_hicache
and child.id in getattr(self, "pending_host_backups", {})
):
raise HiCachePendingBackupSplit(child)
# child node split into new_node -> child
new_node = TreeNode(priority=child.priority)
new_node.children = {self.get_child_key_fn(key[split_len:]): child}

View File

@@ -694,6 +694,8 @@ class KVCache(abc.ABC):
# default state for optional layer-wise transfer control
self.layer_transfer_counter = None
self.layer_backup_notifiers = []
self.layer_backup_notify_after_indexer = False
# for disagg with nvlink
self.enable_custom_mem_pool, self.custom_mem_pool, _ = (
@@ -745,6 +747,15 @@ class KVCache(abc.ABC):
def register_layer_transfer_counter(self, layer_transfer_counter: LayerDoneCounter):
self.layer_transfer_counter = layer_transfer_counter
def register_layer_backup_notifier(self, notifier):
self.layer_backup_notifiers.append(notifier)
def _notify_layer_kv_stored(self, layer_id: int, source: str = "kv"):
if self.layer_backup_notify_after_indexer and source != "indexer":
return
for notifier in self.layer_backup_notifiers:
notifier(layer_id)
def get_cpu_copy(self, indices):
raise NotImplementedError()
@@ -1036,6 +1047,7 @@ class MHATokenToKVPool(KVCache):
alt_stream=self.alt_stream,
same_kv_dim=self.same_kv_dim,
)
self._notify_layer_kv_stored(layer_id - self.start_layer)
def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor):
if envs.SGLANG_NATIVE_MOVE_KV_CACHE.get():
@@ -1583,6 +1595,7 @@ class MLATokenToKVPool(KVCache):
)
else:
self.kv_buffer[layer_id - self.start_layer][loc] = cache_k
self._notify_layer_kv_stored(layer_id - self.start_layer)
def set_mla_kv_buffer(
self,
@@ -1624,6 +1637,7 @@ class MLATokenToKVPool(KVCache):
cache_k_nope,
cache_k_rope,
)
self._notify_layer_kv_stored(layer_id - self.start_layer)
def get_mla_kv_buffer(
self,
@@ -1845,6 +1859,7 @@ class NSATokenToKVPool(MLATokenToKVPool):
use_nsa=True,
override_kv_cache_dim=override_dim,
)
self.layer_backup_notify_after_indexer = True
# self.index_k_dtype = torch.float8_e4m3fn
# self.index_k_scale_dtype = torch.float32
self.index_head_dim = index_head_dim
@@ -1951,6 +1966,7 @@ class NSATokenToKVPool(MLATokenToKVPool):
index_buf_accessor.SetKAndS.execute(
pool=self, buf=buf, loc=loc, index_k=index_k, index_k_scale=index_k_scale
)
self._notify_layer_kv_stored(layer_id - self.start_layer, source="indexer")
def get_state_buf_infos(self):
data_ptrs = [

View File

@@ -2,7 +2,7 @@ import abc
import logging
import threading
from collections import defaultdict
from functools import wraps
from functools import lru_cache, wraps
from typing import Optional
import psutil
@@ -52,6 +52,20 @@ if _is_npu:
logger = logging.getLogger(__name__)
@lru_cache(maxsize=1)
def _load_tai_transfer_kv_per_layer_mla_lf_pf():
try:
from tai_kernel.nsa_prefill import transfer_kv_per_layer_mla_lf_pf
return transfer_kv_per_layer_mla_lf_pf
except Exception as exc:
raise RuntimeError(
"Per-layer D2H backup for page_first MLA/NSA HiCache requires "
"tai_kernel.nsa_prefill.transfer_kv_per_layer_mla_lf_pf. "
"Build/sync tai-kernel with the LF->PF per-layer kvcacheio op."
) from exc
def synchronized(func):
@wraps(func)
def wrapper(self, *args, **kwargs):
@@ -221,6 +235,15 @@ class HostKVCache(abc.ABC):
"""
raise NotImplementedError()
@abc.abstractmethod
def backup_from_device_per_layer(
self, device_pool, host_indices, device_indices, layer_id, io_backend
) -> None:
"""
Backup KV data from the device memory pool to the host memory pool for a specific layer.
"""
raise NotImplementedError()
@abc.abstractmethod
def backup_from_device_all_layer(
self, device_pool, host_indices, device_indices, io_backend
@@ -577,6 +600,77 @@ class MHATokenToKVPoolHost(HostKVCache):
else:
raise ValueError(f"Unsupported IO backend: {io_backend}")
def backup_from_device_per_layer(
self,
device_pool,
host_indices,
device_indices,
layer_id,
io_backend,
):
if io_backend == "kernel":
if self.layout == "layer_first":
transfer_kv_per_layer(
src_k=device_pool.k_buffer[layer_id],
dst_k=self.k_buffer[layer_id],
src_v=device_pool.v_buffer[layer_id],
dst_v=self.v_buffer[layer_id],
src_indices=device_indices,
dst_indices=host_indices,
item_size=self.token_stride_size,
)
elif self.layout == "page_first":
raise NotImplementedError(
"Per-layer D2H backup for page_first MHA kernel layout requires "
"a dedicated LF->PF per-layer kernel; use all-layer backup until "
"that kernel exists."
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
elif io_backend == "direct":
if self.layout == "layer_first":
transfer_kv_direct(
src_layers=[
device_pool.k_buffer[layer_id],
device_pool.v_buffer[layer_id],
],
dst_layers=[self.k_buffer[layer_id], self.v_buffer[layer_id]],
src_indices=device_indices,
dst_indices=host_indices,
page_size=self.page_size,
)
elif self.layout == "page_first_direct":
for host_index, device_index in zip(
host_indices.cpu().tolist(), device_indices.cpu().tolist()
):
host_page = host_index // self.page_size
host_offset = host_index % self.page_size
self.k_buffer[host_page, layer_id, host_offset] = (
device_pool.k_buffer[layer_id][device_index]
)
self.v_buffer[host_page, layer_id, host_offset] = (
device_pool.v_buffer[layer_id][device_index]
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
elif io_backend == "kernel_ascend":
if self.layout == "page_first_direct":
if layer_id == 0:
transfer_kv_dim_exchange(
device_indices=device_indices,
host_indices=host_indices,
device_k=device_pool.k_buffer,
host_k=self.k_buffer,
device_v=device_pool.v_buffer,
host_v=self.v_buffer,
page_size=self.page_size,
direction=TransferDirection.D2H,
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
else:
raise ValueError(f"Unsupported IO backend: {io_backend}")
def get_data_page(self, index, flat: bool = True) -> torch.Tensor:
if self.layout == "layer_first":
data_page = self.kv_buffer[:, :, index : index + self.page_size, :, :]
@@ -993,6 +1087,75 @@ class MLATokenToKVPoolHost(HostKVCache):
else:
raise ValueError(f"Unsupported IO backend: {io_backend}")
def backup_from_device_per_layer(
self,
device_pool,
host_indices,
device_indices,
layer_id,
io_backend,
):
if io_backend == "kernel":
if self.layout == "layer_first":
transfer_kv_per_layer_mla(
src=device_pool.kv_buffer[layer_id],
dst=self.kv_buffer[layer_id],
src_indices=device_indices,
dst_indices=host_indices,
item_size=self.token_stride_size,
)
elif self.layout == "page_first":
_load_tai_transfer_kv_per_layer_mla_lf_pf()(
device_pool.kv_buffer[layer_id],
self.kv_buffer,
device_indices,
host_indices,
layer_id=layer_id,
item_size=self.token_stride_size,
dst_layout_dim=self.layout_dim,
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
elif io_backend == "direct":
if self.layout == "layer_first":
transfer_kv_direct(
src_layers=[device_pool.kv_buffer[layer_id]],
dst_layers=[self.kv_buffer[layer_id]],
src_indices=device_indices,
dst_indices=host_indices,
page_size=self.page_size,
)
elif self.layout == "page_first_direct":
for host_index, device_index in zip(
host_indices.cpu().tolist(), device_indices.cpu().tolist()
):
host_page = host_index // self.page_size
host_offset = host_index % self.page_size
self.kv_buffer[host_page, layer_id, host_offset] = (
device_pool.kv_buffer[layer_id][device_index]
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
elif io_backend == "kernel_ascend":
if self.layout == "page_first_kv_split":
if layer_id == 0:
transfer_kv_dim_exchange(
device_indices=device_indices,
host_indices=host_indices,
device_k=device_pool.k_buffer,
host_k=self.k_buffer,
device_v=device_pool.v_buffer,
host_v=self.v_buffer,
device_index_k=device_pool.index_k_buffer,
host_index_k=self.index_k_buffer,
page_size=self.page_size,
direction=TransferDirection.D2H,
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
else:
raise ValueError(f"Unsupported IO backend: {io_backend}")
def get_data_page(self, index, flat: bool = True) -> torch.Tensor:
if self.layout == "layer_first":
data_page = self.kv_buffer[:, index : index + self.page_size, :, :]
@@ -1299,6 +1462,56 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
else:
raise ValueError(f"Unsupported IO backend: {io_backend}")
def _backup_indexer_from_device_per_layer(
self, device_pool, host_indices, device_indices, layer_id, io_backend
):
host_page_indices, device_page_indices = self._get_indexer_page_indices(
host_indices, device_indices
)
use_kernel = io_backend == "kernel" and self.indexer_page_stride_size % 8 == 0
if use_kernel:
if self.layout == "layer_first":
transfer_kv_per_layer_mla(
src=device_pool.index_k_with_scale_buffer[layer_id],
dst=self.index_k_with_scale_buffer[layer_id],
src_indices=device_page_indices,
dst_indices=host_page_indices,
item_size=self.indexer_page_stride_size,
)
elif self.layout == "page_first":
_load_tai_transfer_kv_per_layer_mla_lf_pf()(
device_pool.index_k_with_scale_buffer[layer_id],
self.index_k_with_scale_buffer,
device_page_indices,
host_page_indices,
layer_id=layer_id,
item_size=self.indexer_page_stride_size,
dst_layout_dim=self.indexer_layout_dim,
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
elif io_backend == "direct":
if self.layout == "layer_first":
transfer_kv_direct(
src_layers=[device_pool.index_k_with_scale_buffer[layer_id]],
dst_layers=[self.index_k_with_scale_buffer[layer_id]],
src_indices=device_page_indices,
dst_indices=host_page_indices,
page_size=1,
)
elif self.layout == "page_first_direct":
for host_page, device_page in zip(
host_page_indices.cpu().tolist(),
device_page_indices.cpu().tolist(),
):
self.index_k_with_scale_buffer[host_page, layer_id, 0] = (
device_pool.index_k_with_scale_buffer[layer_id][device_page]
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
else:
raise ValueError(f"Unsupported IO backend: {io_backend}")
def load_to_device_per_layer(
self,
device_pool,
@@ -1323,3 +1536,13 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
self._backup_indexer_from_device_all_layer(
device_pool, host_indices, device_indices, io_backend
)
def backup_from_device_per_layer(
self, device_pool, host_indices, device_indices, layer_id, io_backend
):
super().backup_from_device_per_layer(
device_pool, host_indices, device_indices, layer_id, io_backend
)
self._backup_indexer_from_device_per_layer(
device_pool, host_indices, device_indices, layer_id, io_backend
)