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:
@@ -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."""
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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 = [
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user