Make CP HiCache residency owner-lane deterministic
CP shared KV cannot treat capacity as a scalar token count: cache-hit load-back and fresh extend allocation both have to preserve the logical page owner pattern or later direct writes, HiCache reload, and prefix materialization can read the wrong lane. This change moves the critical paths to owner-lane plans, makes owner-lane exhaustion recoverable during prefill scheduling, and routes shared-KV prefix prefetch through prefetch-stream-safe KV getters so HiCache layer-load waits do not attach to the forward stream. Constraint: CP shared KV correctness depends on page owner lane preservation across allocation, backup, load, eviction, and prefix materialization. Constraint: Avoid adding CP/global collectives for capacity agreement; derive capacity from deterministic local owner-lane state. Rejected: Keep SGLANG_DISABLE_TAI_OWNER_SELECT fallback | legacy allocation can silently break owner-lane invariants. Rejected: Scalar total-token eviction for CP HiCache load-back | total capacity can be sufficient while the required owner lane is exhausted. Confidence: medium Scope-risk: broad Directive: Do not reintroduce silent legacy fallback in owner-lane paths; unexpected owner-lane failure must be warning-level fail-closed or recoverable capacity wait. Tested: Remote g0034 container PYTHONPATH=python python -m pytest test/registered/unit/mem_cache/test_alloc_pages_with_owners.py test/registered/unit/mem_cache/test_cp_shared_kv_layout.py test/registered/unit/mem_cache/test_cp_hicache_load_back_owner_lanes.py test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py -q -> 95 passed. Tested: Local py_compile for modified runtime/cache/scheduler modules. Not-tested: Full CUDA ETE performance trace for cache-hit overlap and MTP accept-rate impact. Co-authored-by: OmX <omx@oh-my-codex.dev>
This commit is contained in:
@@ -45,6 +45,50 @@ def _index_prefetch_fallback_log(reason: str, message: str, *args) -> None:
|
||||
)
|
||||
|
||||
|
||||
def _prefetch_pool_get_key_buffer(
|
||||
*,
|
||||
token_to_kv_pool: Any,
|
||||
layer_id: int,
|
||||
stream: torch.cuda.Stream,
|
||||
path: str,
|
||||
) -> torch.Tensor:
|
||||
getter = getattr(token_to_kv_pool, "get_key_buffer_for_prefetch", None)
|
||||
if getter is not None:
|
||||
return getter(layer_id, stream)
|
||||
if getattr(token_to_kv_pool, "layer_transfer_counter", None) is not None:
|
||||
logger.warning(
|
||||
"[CP_SHARED_KV_FALLBACK][%s_prefetch] "
|
||||
"reason=prefetch_safe_getter_unavailable layer_id=%s pool=%s "
|
||||
"has_layer_transfer_counter=True",
|
||||
path,
|
||||
layer_id,
|
||||
type(token_to_kv_pool).__name__,
|
||||
)
|
||||
return token_to_kv_pool.get_key_buffer(layer_id)
|
||||
|
||||
|
||||
def _prefetch_pool_get_index_buffer(
|
||||
*,
|
||||
token_to_kv_pool: Any,
|
||||
layer_id: int,
|
||||
stream: torch.cuda.Stream,
|
||||
) -> torch.Tensor:
|
||||
getter = getattr(
|
||||
token_to_kv_pool, "get_index_k_with_scale_buffer_for_prefetch", None
|
||||
)
|
||||
if getter is not None:
|
||||
return getter(layer_id, stream)
|
||||
if getattr(token_to_kv_pool, "layer_transfer_counter", None) is not None:
|
||||
_index_prefetch_fallback_log(
|
||||
"prefetch_safe_getter_unavailable",
|
||||
"pool has active layer transfer counter but no prefetch-safe index "
|
||||
"getter. layer_id=%s pool=%s",
|
||||
layer_id,
|
||||
type(token_to_kv_pool).__name__,
|
||||
)
|
||||
return token_to_kv_pool.get_index_k_with_scale_buffer(layer_id=layer_id)
|
||||
|
||||
|
||||
def _is_cuda_stream_capturing() -> bool:
|
||||
if not torch.cuda.is_available():
|
||||
return False
|
||||
@@ -235,9 +279,15 @@ class CpSharedKVMlaPrefetcher:
|
||||
)
|
||||
return None
|
||||
|
||||
prefetch_stream = stream if stream is not None else torch.cuda.Stream()
|
||||
try:
|
||||
first_layer_id = int(getattr(token_to_kv_pool, "start_layer", 0))
|
||||
kv_cache = token_to_kv_pool.get_key_buffer(first_layer_id)
|
||||
kv_cache = _prefetch_pool_get_key_buffer(
|
||||
token_to_kv_pool=token_to_kv_pool,
|
||||
layer_id=first_layer_id,
|
||||
stream=prefetch_stream,
|
||||
path="mla",
|
||||
)
|
||||
remap = get_or_build_shared_token_kv_slot_remap(
|
||||
forward_batch,
|
||||
kv_cache=kv_cache,
|
||||
@@ -276,7 +326,7 @@ class CpSharedKVMlaPrefetcher:
|
||||
dense_num_pages=remap.dense_num_pages,
|
||||
owned_prefix_pages=owned_prefix_pages,
|
||||
owned_total_pages=owned_total_pages,
|
||||
stream=stream,
|
||||
stream=prefetch_stream,
|
||||
)
|
||||
|
||||
def _layer_in_pool(self, token_to_kv_pool: Any, layer_id: int) -> bool:
|
||||
@@ -447,7 +497,12 @@ class CpSharedKVMlaPrefetcher:
|
||||
return
|
||||
|
||||
try:
|
||||
kv_cache = token_to_kv_pool.get_key_buffer(next_layer_id)
|
||||
kv_cache = _prefetch_pool_get_key_buffer(
|
||||
token_to_kv_pool=token_to_kv_pool,
|
||||
layer_id=next_layer_id,
|
||||
stream=self.stream,
|
||||
path="mla",
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Failed to get next-layer KV cache for CP shared KV MLA prefetch."
|
||||
@@ -461,31 +516,61 @@ class CpSharedKVMlaPrefetcher:
|
||||
return
|
||||
|
||||
try:
|
||||
dense_kv_cache = kv_cache.new_zeros(
|
||||
(self.dense_num_pages * self.page_size, *kv_cache.shape[1:])
|
||||
)
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"start_prefix_begin next_layer=%s start_slot=0 end_slot=%s "
|
||||
"dense_rows=%s",
|
||||
next_layer_id,
|
||||
self.prefix_pages,
|
||||
int(dense_kv_cache.shape[0]),
|
||||
)
|
||||
materialize_local_token_kv_page_slots_into(
|
||||
kv_cache=kv_cache,
|
||||
dense_kv_cache=dense_kv_cache,
|
||||
slot_logical_pages=self.slot_logical_pages,
|
||||
layout=self.layout,
|
||||
page_size=self.page_size,
|
||||
start_slot=0,
|
||||
end_slot=self.prefix_pages,
|
||||
)
|
||||
current_stream = torch.cuda.current_stream()
|
||||
self.stream.wait_stream(current_stream)
|
||||
prefix_rows = slot_range_to_token_slice(
|
||||
self.page_size,
|
||||
0,
|
||||
self.prefix_pages,
|
||||
)
|
||||
with torch.cuda.stream(self.stream):
|
||||
dense_kv_cache = kv_cache.new_zeros(
|
||||
(self.dense_num_pages * self.page_size, *kv_cache.shape[1:])
|
||||
)
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"start_prefix_begin next_layer=%s start_slot=0 end_slot=%s "
|
||||
"dense_rows=%s",
|
||||
next_layer_id,
|
||||
self.prefix_pages,
|
||||
int(dense_kv_cache.shape[0]),
|
||||
)
|
||||
materialize_local_token_kv_page_slots_into(
|
||||
kv_cache=kv_cache,
|
||||
dense_kv_cache=dense_kv_cache,
|
||||
slot_logical_pages=self.slot_logical_pages,
|
||||
layout=self.layout,
|
||||
page_size=self.page_size,
|
||||
start_slot=0,
|
||||
end_slot=self.prefix_pages,
|
||||
)
|
||||
event = _all_reduce_materialized_buffer_async(
|
||||
dense_kv_cache[prefix_rows],
|
||||
cp_size=self.layout.cp_size,
|
||||
stream=self.stream,
|
||||
nvtx_source="mla.prefetch_prefix",
|
||||
nvtx_layer_id=next_layer_id,
|
||||
nvtx_cp_rank=self.layout.cp_rank,
|
||||
nvtx_rows=(prefix_rows.start, prefix_rows.stop),
|
||||
)
|
||||
if event is None:
|
||||
self.disabled = True
|
||||
logger.warning(
|
||||
"[CP_SHARED_KV_FALLBACK][mla_prefetch] "
|
||||
"reason=async_reduce_unavailable layer_id=%s cp_rank=%s "
|
||||
"cp_size=%s prefix_pages=%s total_slots=%s",
|
||||
next_layer_id,
|
||||
self.layout.cp_rank,
|
||||
self.layout.cp_size,
|
||||
self.prefix_pages,
|
||||
self.total_slots,
|
||||
)
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"start_disable reason=async_reduce_unavailable next_layer=%s",
|
||||
next_layer_id,
|
||||
)
|
||||
return
|
||||
except Exception:
|
||||
logger.exception("Failed to start CP shared KV MLA prefix prefetch.")
|
||||
self.disabled = True
|
||||
@@ -500,12 +585,14 @@ class CpSharedKVMlaPrefetcher:
|
||||
layer_id=next_layer_id,
|
||||
dense_kv_cache=dense_kv_cache,
|
||||
prefix_rows=prefix_rows,
|
||||
event=event,
|
||||
)
|
||||
self.handles[next_layer_id] = handle
|
||||
self.pending_attention_handle = handle
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"start next_layer=%s prefix_pages=%s prefix_rows=%s dense_rows=%s",
|
||||
"start next_layer=%s prefix_pages=%s prefix_rows=%s dense_rows=%s "
|
||||
"reduce_enqueued=True",
|
||||
next_layer_id,
|
||||
self.prefix_pages,
|
||||
prefix_rows.stop - prefix_rows.start,
|
||||
@@ -776,10 +863,13 @@ class CpSharedKVIndexPrefetcher:
|
||||
)
|
||||
return None
|
||||
|
||||
prefetch_stream = stream if stream is not None else torch.cuda.Stream()
|
||||
try:
|
||||
first_layer_id = int(getattr(token_to_kv_pool, "start_layer", 0))
|
||||
page_buffer = token_to_kv_pool.get_index_k_with_scale_buffer(
|
||||
layer_id=first_layer_id
|
||||
page_buffer = _prefetch_pool_get_index_buffer(
|
||||
token_to_kv_pool=token_to_kv_pool,
|
||||
layer_id=first_layer_id,
|
||||
stream=prefetch_stream,
|
||||
)
|
||||
remap = get_or_build_shared_paged_buffer_slot_remap(
|
||||
forward_batch,
|
||||
@@ -822,7 +912,7 @@ class CpSharedKVIndexPrefetcher:
|
||||
dense_num_pages=remap.dense_num_pages,
|
||||
owned_prefix_pages=owned_prefix_pages,
|
||||
owned_total_pages=owned_total_pages,
|
||||
stream=stream,
|
||||
stream=prefetch_stream,
|
||||
)
|
||||
|
||||
def _layer_in_pool(self, token_to_kv_pool: Any, layer_id: int) -> bool:
|
||||
@@ -990,8 +1080,10 @@ class CpSharedKVIndexPrefetcher:
|
||||
return
|
||||
|
||||
try:
|
||||
page_buffer = token_to_kv_pool.get_index_k_with_scale_buffer(
|
||||
layer_id=next_layer_id
|
||||
page_buffer = _prefetch_pool_get_index_buffer(
|
||||
token_to_kv_pool=token_to_kv_pool,
|
||||
layer_id=next_layer_id,
|
||||
stream=self.stream,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
@@ -1006,26 +1098,56 @@ class CpSharedKVIndexPrefetcher:
|
||||
return
|
||||
|
||||
try:
|
||||
dense_page_buffer = page_buffer.new_zeros(
|
||||
(self.dense_num_pages, *page_buffer.shape[1:])
|
||||
)
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"index_start_prefix_begin next_layer=%s start_slot=0 "
|
||||
"end_slot=%s dense_pages=%s",
|
||||
next_layer_id,
|
||||
self.prefix_pages,
|
||||
int(dense_page_buffer.shape[0]),
|
||||
)
|
||||
materialize_local_paged_buffer_page_slots_into(
|
||||
page_buffer=page_buffer,
|
||||
dense_page_buffer=dense_page_buffer,
|
||||
slot_logical_pages=self.slot_logical_pages,
|
||||
layout=self.layout,
|
||||
start_slot=0,
|
||||
end_slot=self.prefix_pages,
|
||||
)
|
||||
current_stream = torch.cuda.current_stream()
|
||||
self.stream.wait_stream(current_stream)
|
||||
prefix_rows = slot_range_to_page_slice(0, self.prefix_pages)
|
||||
with torch.cuda.stream(self.stream):
|
||||
dense_page_buffer = page_buffer.new_zeros(
|
||||
(self.dense_num_pages, *page_buffer.shape[1:])
|
||||
)
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"index_start_prefix_begin next_layer=%s start_slot=0 "
|
||||
"end_slot=%s dense_pages=%s",
|
||||
next_layer_id,
|
||||
self.prefix_pages,
|
||||
int(dense_page_buffer.shape[0]),
|
||||
)
|
||||
materialize_local_paged_buffer_page_slots_into(
|
||||
page_buffer=page_buffer,
|
||||
dense_page_buffer=dense_page_buffer,
|
||||
slot_logical_pages=self.slot_logical_pages,
|
||||
layout=self.layout,
|
||||
start_slot=0,
|
||||
end_slot=self.prefix_pages,
|
||||
)
|
||||
event = _all_reduce_materialized_buffer_async(
|
||||
dense_page_buffer[prefix_rows],
|
||||
cp_size=self.layout.cp_size,
|
||||
stream=self.stream,
|
||||
nvtx_source="index.prefetch_prefix",
|
||||
nvtx_layer_id=next_layer_id,
|
||||
nvtx_cp_rank=self.layout.cp_rank,
|
||||
nvtx_rows=(prefix_rows.start, prefix_rows.stop),
|
||||
)
|
||||
if event is None:
|
||||
self.disabled = True
|
||||
_index_prefetch_fallback_log(
|
||||
"async_reduce_unavailable",
|
||||
"async reduce unavailable. layer_id=%s cp_rank=%s cp_size=%s "
|
||||
"prefix_pages=%s total_slots=%s",
|
||||
next_layer_id,
|
||||
self.layout.cp_rank,
|
||||
self.layout.cp_size,
|
||||
self.prefix_pages,
|
||||
self.total_slots,
|
||||
)
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"index_start_disable reason=async_reduce_unavailable next_layer=%s",
|
||||
next_layer_id,
|
||||
)
|
||||
return
|
||||
except Exception:
|
||||
logger.exception("Failed to start CP shared KV index prefix prefetch.")
|
||||
self.disabled = True
|
||||
@@ -1040,12 +1162,14 @@ class CpSharedKVIndexPrefetcher:
|
||||
layer_id=next_layer_id,
|
||||
dense_page_buffer=dense_page_buffer,
|
||||
prefix_rows=prefix_rows,
|
||||
event=event,
|
||||
)
|
||||
self.handles[next_layer_id] = handle
|
||||
self.pending_attention_handle = handle
|
||||
self._log_next_layer(
|
||||
next_layer_id,
|
||||
"index_start next_layer=%s prefix_pages=%s dense_pages=%s",
|
||||
"index_start next_layer=%s prefix_pages=%s dense_pages=%s "
|
||||
"reduce_enqueued=True",
|
||||
next_layer_id,
|
||||
self.prefix_pages,
|
||||
int(dense_page_buffer.shape[0]),
|
||||
|
||||
@@ -61,6 +61,10 @@ class LayerLoadingEvent:
|
||||
def wait(self, layer_index: int):
|
||||
device_module.current_stream().wait_event(self.load_events[layer_index])
|
||||
|
||||
def wait_on_stream(self, layer_index: int, stream):
|
||||
assert 0 <= layer_index < self._num_layers
|
||||
stream.wait_event(self.load_events[layer_index])
|
||||
|
||||
@property
|
||||
def finish_event(self):
|
||||
return self.load_events[-1]
|
||||
@@ -101,6 +105,12 @@ class LayerDoneCounter:
|
||||
for consumer_index in self.consumer_indices:
|
||||
self.events[consumer_index].wait(threshold)
|
||||
|
||||
def wait_until_on_stream(self, threshold: int, stream) -> None:
|
||||
if not self.consumer_indices:
|
||||
return
|
||||
for consumer_index in self.consumer_indices:
|
||||
self.events[consumer_index].wait_on_stream(threshold, stream)
|
||||
|
||||
def reset(self):
|
||||
self.producer_index = -1
|
||||
self.consumer_index = -1
|
||||
|
||||
@@ -179,8 +179,9 @@ from sglang.srt.managers.scheduler_update_weights_mixin import (
|
||||
)
|
||||
from sglang.srt.managers.session_controller import SessionController
|
||||
from sglang.srt.managers.utils import GenerationBatchResult, validate_input_length
|
||||
from sglang.srt.mem_cache.base_prefix_cache import DecLockRefParams
|
||||
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
||||
from sglang.srt.mem_cache.common import release_kv_cache
|
||||
from sglang.srt.mem_cache.common import KVCapacityWaitError, release_kv_cache
|
||||
from sglang.srt.mem_cache.radix_cache import RadixCache
|
||||
from sglang.srt.mem_cache.session_aware_cache import SessionAwareCache
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode, PPProxyTensors
|
||||
@@ -2417,6 +2418,8 @@ class Scheduler(
|
||||
if len(can_run_list) == 0:
|
||||
return None
|
||||
|
||||
waiting_queue_before_prepare = list(self.waiting_queue)
|
||||
chunked_req_before_prepare = self.chunked_req
|
||||
can_run_set = set(can_run_list)
|
||||
self.waiting_queue = [x for x in self.waiting_queue if x not in can_run_set]
|
||||
if adder.preempt_list:
|
||||
@@ -2456,7 +2459,20 @@ class Scheduler(
|
||||
self.tree_cache.ready_to_load_host_cache()
|
||||
)
|
||||
|
||||
new_batch.prepare_for_extend()
|
||||
try:
|
||||
new_batch.prepare_for_extend()
|
||||
except KVCapacityWaitError as exc:
|
||||
self._release_prefill_adder_locks(
|
||||
can_run_list, skip_req=chunked_req_before_prepare
|
||||
)
|
||||
self.waiting_queue = waiting_queue_before_prepare
|
||||
self.chunked_req = chunked_req_before_prepare
|
||||
self.running_batch.batch_is_full = True
|
||||
logger.warning(
|
||||
"[CP_SHARED_KV_CAPACITY_WAIT] prefill allocation deferred: %s",
|
||||
exc,
|
||||
)
|
||||
return None
|
||||
|
||||
# Record prefill stats for logging after forward
|
||||
new_batch.prefill_stats = PrefillStats.from_adder(
|
||||
@@ -2485,6 +2501,38 @@ class Scheduler(
|
||||
|
||||
return new_batch
|
||||
|
||||
def _release_prefill_adder_locks(
|
||||
self, reqs: List[Any], *, skip_req: Optional[Any] = None
|
||||
) -> None:
|
||||
"""Release lock refs acquired by PrefillAdder for a failed prefill batch.
|
||||
|
||||
`PrefillAdder.add_one_req` persists one `inc_lock_ref` per staged
|
||||
request after its short-lived scheduling lock context exits. If KV
|
||||
allocation later reports a recoverable capacity wait, the batch never
|
||||
runs, so those persistent refs must be dropped before the requests are
|
||||
returned to the waiting queue.
|
||||
"""
|
||||
|
||||
if getattr(self.tree_cache, "disable", False):
|
||||
return
|
||||
|
||||
use_swa_params = self.tree_cache.supports_swa() and self.tree_cache.is_tree_cache()
|
||||
for req in reqs:
|
||||
if req is skip_req:
|
||||
continue
|
||||
last_node = getattr(req, "last_node", None)
|
||||
if last_node is None:
|
||||
continue
|
||||
if use_swa_params:
|
||||
params = DecLockRefParams(
|
||||
swa_uuid_for_lock=getattr(req, "swa_uuid_for_lock", None)
|
||||
)
|
||||
self.tree_cache.dec_lock_ref(last_node, params)
|
||||
if hasattr(req, "swa_uuid_for_lock"):
|
||||
req.swa_uuid_for_lock = None
|
||||
else:
|
||||
self.tree_cache.dec_lock_ref(last_node)
|
||||
|
||||
def update_running_batch(self, batch: ScheduleBatch) -> Optional[ScheduleBatch]:
|
||||
"""Update the current running decoding batch."""
|
||||
initial_bs = batch.batch_size()
|
||||
|
||||
@@ -78,6 +78,7 @@ class EvictParams:
|
||||
num_tokens: int
|
||||
swa_num_tokens: int = 0
|
||||
mamba_num: int = 0
|
||||
owner_lane_deficits: Optional[list[int]] = None
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
|
||||
@@ -32,6 +32,27 @@ MAMBA_STATE_PER_REQ_NO_CACHE = 1
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class KVCapacityWaitError(RuntimeError):
|
||||
"""Recoverable KV-capacity pressure with CP owner-lane diagnostics."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
*,
|
||||
required_by_owner: list[int] | None = None,
|
||||
available_by_owner: list[int] | None = None,
|
||||
deficit_by_owner: list[int] | None = None,
|
||||
):
|
||||
self.required_by_owner = required_by_owner
|
||||
self.available_by_owner = available_by_owner
|
||||
self.deficit_by_owner = deficit_by_owner
|
||||
super().__init__(
|
||||
f"{message}; required_by_owner={required_by_owner} "
|
||||
f"available_by_owner={available_by_owner} "
|
||||
f"deficit_by_owner={deficit_by_owner}"
|
||||
)
|
||||
|
||||
|
||||
def _log_cp_shared_kv_alloc_fallback(
|
||||
reason: str,
|
||||
message: str,
|
||||
@@ -341,7 +362,12 @@ def _evict_for_compute_owner_lanes(
|
||||
before_available,
|
||||
evictable_size,
|
||||
)
|
||||
evict_result = tree_cache.evict(EvictParams(num_tokens=evict_tokens))
|
||||
evict_result = tree_cache.evict(
|
||||
EvictParams(
|
||||
num_tokens=evict_tokens,
|
||||
owner_lane_deficits=[int(v) for v in deficits],
|
||||
)
|
||||
)
|
||||
after_available = allocator.available_size()
|
||||
evicted_tokens = getattr(evict_result, "num_tokens_evicted", 0)
|
||||
logger.info(
|
||||
@@ -367,7 +393,7 @@ def alloc_paged_token_slots_extend(
|
||||
# Over estimate the number of tokens: assume each request needs a new page.
|
||||
allocator = tree_cache.token_to_kv_pool_allocator
|
||||
num_tokens = extend_num_tokens + len(seq_lens_cpu) * allocator.page_size
|
||||
evict_result = evict_from_tree_cache(tree_cache, num_tokens)
|
||||
evict_result = None
|
||||
|
||||
# logger.info(
|
||||
# "[MemCache-alloc] alloc_paged_token_slots_extend: extend_num_tokens=%d batch_size=%d num_tokens=%d page_size=%d "
|
||||
@@ -420,11 +446,11 @@ def alloc_paged_token_slots_extend(
|
||||
"multi_batch" if len(prefix_lens_cpu) != 1 else "unknown"
|
||||
)
|
||||
|
||||
state = None
|
||||
if backup_state:
|
||||
state = allocator.backup_state()
|
||||
|
||||
if page_compute_owners is not None:
|
||||
state = None
|
||||
if backup_state:
|
||||
state = allocator.backup_state()
|
||||
|
||||
out_cache_loc = alloc_extend_compute_owner(
|
||||
prefix_lens,
|
||||
prefix_lens_cpu,
|
||||
@@ -460,10 +486,10 @@ def alloc_paged_token_slots_extend(
|
||||
required, available, deficits = compute_owner_lane_stats(
|
||||
page_compute_owners
|
||||
)
|
||||
_log_cp_shared_kv_alloc_fallback(
|
||||
"owner_lane_exhausted",
|
||||
"failed to allocate pages from compute-owner lanes; "
|
||||
"falling back to legacy page allocation. extend_num_tokens=%s page_size=%s "
|
||||
logger.warning(
|
||||
"[CP_SHARED_KV_FAIL_FAST][owner_lane_exhausted] "
|
||||
"failed to allocate pages from compute-owner lanes after eviction; "
|
||||
"raising recoverable capacity wait. extend_num_tokens=%s page_size=%s "
|
||||
"required_by_owner=%s available_by_owner=%s deficit_by_owner=%s",
|
||||
extend_num_tokens,
|
||||
allocator.page_size,
|
||||
@@ -471,15 +497,19 @@ def alloc_paged_token_slots_extend(
|
||||
available,
|
||||
deficits,
|
||||
)
|
||||
out_cache_loc = allocator.alloc_extend(
|
||||
prefix_lens,
|
||||
prefix_lens_cpu,
|
||||
seq_lens,
|
||||
seq_lens_cpu,
|
||||
last_loc,
|
||||
extend_num_tokens,
|
||||
raise KVCapacityWaitError(
|
||||
"owner_lane_exhausted: failed to allocate pages from CP "
|
||||
"compute-owner lanes after owner-lane eviction",
|
||||
required_by_owner=required,
|
||||
available_by_owner=available,
|
||||
deficit_by_owner=deficits,
|
||||
)
|
||||
else:
|
||||
evict_result = evict_from_tree_cache(tree_cache, num_tokens)
|
||||
state = None
|
||||
if backup_state:
|
||||
state = allocator.backup_state()
|
||||
|
||||
if alloc_extend_compute_owner is not None:
|
||||
_log_cp_shared_kv_alloc_fallback(
|
||||
compute_owner_unavailable_reason or "compute_owner_not_available",
|
||||
@@ -568,6 +598,7 @@ def alloc_for_extend(
|
||||
extend_lens_device = extend_lens_cpu.to(batch.device, non_blocking=True)
|
||||
|
||||
# Allocate req slots
|
||||
newly_allocated_reqs = [req for req in batch.reqs if req.req_pool_idx is None]
|
||||
req_pool_indices = alloc_req_slots(
|
||||
batch.req_to_token_pool, batch.reqs, batch.tree_cache
|
||||
)
|
||||
@@ -575,23 +606,29 @@ def alloc_for_extend(
|
||||
req_pool_indices_device = req_pool_indices_cpu.to(batch.device, non_blocking=True)
|
||||
|
||||
# Allocate KV cache (throws exception on failure)
|
||||
if batch.tree_cache.page_size == 1:
|
||||
out_cache_loc = alloc_token_slots(batch.tree_cache, batch.extend_num_tokens)
|
||||
else:
|
||||
# Paged allocation - build last_loc
|
||||
last_loc = [
|
||||
(t[-1:] if len(t) > 0 else torch.tensor([-1], device=batch.device))
|
||||
for t in prefix_tensors
|
||||
]
|
||||
out_cache_loc = alloc_paged_token_slots_extend(
|
||||
tree_cache=batch.tree_cache,
|
||||
prefix_lens=prefix_lens_device,
|
||||
prefix_lens_cpu=prefix_lens_cpu,
|
||||
seq_lens=batch.seq_lens,
|
||||
seq_lens_cpu=batch.seq_lens_cpu,
|
||||
last_loc=torch.cat(last_loc),
|
||||
extend_num_tokens=batch.extend_num_tokens,
|
||||
)
|
||||
try:
|
||||
if batch.tree_cache.page_size == 1:
|
||||
out_cache_loc = alloc_token_slots(batch.tree_cache, batch.extend_num_tokens)
|
||||
else:
|
||||
# Paged allocation - build last_loc
|
||||
last_loc = [
|
||||
(t[-1:] if len(t) > 0 else torch.tensor([-1], device=batch.device))
|
||||
for t in prefix_tensors
|
||||
]
|
||||
out_cache_loc = alloc_paged_token_slots_extend(
|
||||
tree_cache=batch.tree_cache,
|
||||
prefix_lens=prefix_lens_device,
|
||||
prefix_lens_cpu=prefix_lens_cpu,
|
||||
seq_lens=batch.seq_lens,
|
||||
seq_lens_cpu=batch.seq_lens_cpu,
|
||||
last_loc=torch.cat(last_loc),
|
||||
extend_num_tokens=batch.extend_num_tokens,
|
||||
)
|
||||
except KVCapacityWaitError:
|
||||
for req in newly_allocated_reqs:
|
||||
if req.req_pool_idx is not None:
|
||||
batch.req_to_token_pool.free(req)
|
||||
raise
|
||||
|
||||
# Write to req_to_token_pool
|
||||
write_cache_indices(
|
||||
|
||||
@@ -33,7 +33,6 @@ from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
MatchPrefixParams,
|
||||
MatchResult,
|
||||
)
|
||||
from sglang.srt.mem_cache.common import _evict_for_compute_owner_lanes
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
from sglang.srt.mem_cache.memory_pool import (
|
||||
MHATokenToKVPool,
|
||||
@@ -288,6 +287,22 @@ class CpHiCacheEvictionPlan:
|
||||
remaining_deficit: Tuple[int, ...]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CpLoadBackPlan:
|
||||
page_owners: List[int]
|
||||
required_by_owner: List[int]
|
||||
available_by_owner: List[int]
|
||||
deficit_by_owner: List[int]
|
||||
host_hit_len: int
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CpLoadBackEvictionPlan:
|
||||
victims: Tuple[TreeNode, ...]
|
||||
planned_freed_by_owner: Tuple[int, ...]
|
||||
remaining_deficit_by_owner: Tuple[int, ...]
|
||||
|
||||
|
||||
class HiCachePendingBackupSplit(Exception):
|
||||
def __init__(self, node: TreeNode):
|
||||
self.node = node
|
||||
@@ -869,6 +884,249 @@ class HiRadixCache(RadixCache):
|
||||
remaining_deficit=tuple(deficits),
|
||||
)
|
||||
|
||||
def _build_cp_load_back_plan(
|
||||
self, nodes_to_load: List[TreeNode], *, node_id: int
|
||||
) -> CpLoadBackPlan:
|
||||
allocator = self.token_to_kv_pool_allocator
|
||||
compute_owner_lane_stats = getattr(allocator, "compute_owner_lane_stats", None)
|
||||
if compute_owner_lane_stats is None:
|
||||
raise RuntimeError(
|
||||
"CP HiCache load-back requires an allocator with "
|
||||
f"compute_owner_lane_stats: node_id={node_id}"
|
||||
)
|
||||
|
||||
page_owners: List[int] = []
|
||||
host_hit_len = 0
|
||||
for node in nodes_to_load:
|
||||
metadata = getattr(node, "cp_hicache", None)
|
||||
if metadata is None:
|
||||
raise RuntimeError(
|
||||
"CP HiCache load-back missing cp_hicache metadata: "
|
||||
f"load_node_id={getattr(node, 'id', '?')} root_node_id={node_id}"
|
||||
)
|
||||
node_host_len = int(self._node_host_len(node))
|
||||
if node_host_len % self.page_size != 0:
|
||||
raise RuntimeError(
|
||||
"CP HiCache load-back requires page-aligned host lengths: "
|
||||
f"load_node_id={getattr(node, 'id', '?')} "
|
||||
f"host_len={node_host_len} page_size={self.page_size}"
|
||||
)
|
||||
if int(getattr(metadata, "logical_len", node_host_len)) != node_host_len:
|
||||
raise RuntimeError(
|
||||
"CP HiCache load-back metadata length mismatch: "
|
||||
f"load_node_id={getattr(node, 'id', '?')} "
|
||||
f"host_len={node_host_len} metadata_len={metadata.logical_len}"
|
||||
)
|
||||
page_owners.extend(int(owner) for owner in metadata.page_owners.tolist())
|
||||
host_hit_len += node_host_len
|
||||
|
||||
expected_pages = host_hit_len // self.page_size
|
||||
if len(page_owners) != expected_pages:
|
||||
raise RuntimeError(
|
||||
"CP HiCache load-back page owner count mismatch: "
|
||||
f"node_id={node_id} host_hit_len={host_hit_len} "
|
||||
f"page_size={self.page_size} page_owners={len(page_owners)}"
|
||||
)
|
||||
|
||||
required, available, deficits = compute_owner_lane_stats(page_owners)
|
||||
return CpLoadBackPlan(
|
||||
page_owners=list(page_owners),
|
||||
required_by_owner=[int(v) for v in required],
|
||||
available_by_owner=[int(v) for v in available],
|
||||
deficit_by_owner=[int(v) for v in deficits],
|
||||
host_hit_len=host_hit_len,
|
||||
)
|
||||
|
||||
def _refresh_cp_load_back_plan(self, plan: CpLoadBackPlan) -> CpLoadBackPlan:
|
||||
required, available, deficits = (
|
||||
self.token_to_kv_pool_allocator.compute_owner_lane_stats(plan.page_owners)
|
||||
)
|
||||
return CpLoadBackPlan(
|
||||
page_owners=plan.page_owners,
|
||||
required_by_owner=[int(v) for v in required],
|
||||
available_by_owner=[int(v) for v in available],
|
||||
deficit_by_owner=[int(v) for v in deficits],
|
||||
host_hit_len=plan.host_hit_len,
|
||||
)
|
||||
|
||||
def _cp_device_leaf_is_load_back_victim(self, node: TreeNode) -> bool:
|
||||
if node == getattr(self, "root_node", None):
|
||||
return False
|
||||
if getattr(node, "lock_ref", 0) > 0:
|
||||
return False
|
||||
if getattr(node, "value", None) is None:
|
||||
return False
|
||||
parent = getattr(node, "parent", None)
|
||||
if parent is None:
|
||||
return False
|
||||
try:
|
||||
parent_key = self.get_child_key_fn(node.key)
|
||||
except Exception:
|
||||
return False
|
||||
if getattr(parent, "children", {}).get(parent_key) is not node:
|
||||
return False
|
||||
if self._is_pinned(node):
|
||||
return False
|
||||
return True
|
||||
|
||||
def _cp_load_back_node_owner_page_counts(
|
||||
self, node: TreeNode, cp_size: int
|
||||
) -> Tuple[int, ...]:
|
||||
value = getattr(node, "value", None)
|
||||
if value is None or int(value.numel()) == 0:
|
||||
return tuple(0 for _ in range(cp_size))
|
||||
if int(value.numel()) % self.page_size != 0:
|
||||
raise RuntimeError(
|
||||
"CP HiCache load-back eviction requires page-aligned device values: "
|
||||
f"node_id={getattr(node, 'id', '?')} value_len={value.numel()} "
|
||||
f"page_size={self.page_size}"
|
||||
)
|
||||
|
||||
first_locs = value[:: self.page_size]
|
||||
logical_pages = torch.div(first_locs, self.page_size, rounding_mode="floor")
|
||||
owners = torch.remainder(logical_pages - 1, cp_size)
|
||||
return tuple(int((owners == owner).sum().item()) for owner in range(cp_size))
|
||||
|
||||
def _plan_cp_load_back_owner_lane_evictions(
|
||||
self, plan: CpLoadBackPlan
|
||||
) -> CpLoadBackEvictionPlan:
|
||||
deficits = [max(0, int(v)) for v in plan.deficit_by_owner]
|
||||
cp_size = len(deficits)
|
||||
planned_freed = [0 for _ in range(cp_size)]
|
||||
victims: List[TreeNode] = []
|
||||
used_node_ids = set()
|
||||
|
||||
while any(v > 0 for v in deficits):
|
||||
best_node = None
|
||||
best_counts = None
|
||||
best_score = None
|
||||
for node in list(getattr(self, "evictable_leaves", set())):
|
||||
node_id = getattr(node, "id", None)
|
||||
if node_id in used_node_ids:
|
||||
continue
|
||||
if not self._cp_device_leaf_is_load_back_victim(node):
|
||||
continue
|
||||
counts = self._cp_load_back_node_owner_page_counts(node, cp_size)
|
||||
contribution = sum(
|
||||
min(int(count), int(deficit))
|
||||
for count, deficit in zip(counts, deficits)
|
||||
)
|
||||
if contribution <= 0:
|
||||
continue
|
||||
score = (
|
||||
-int(contribution),
|
||||
self.eviction_strategy.get_priority(node),
|
||||
int(node_id or 0),
|
||||
)
|
||||
if best_score is None or score < best_score:
|
||||
best_score = score
|
||||
best_node = node
|
||||
best_counts = counts
|
||||
|
||||
if best_node is None or best_counts is None:
|
||||
break
|
||||
|
||||
victims.append(best_node)
|
||||
used_node_ids.add(getattr(best_node, "id", None))
|
||||
for owner, count in enumerate(best_counts):
|
||||
planned_freed[owner] += int(count)
|
||||
deficits[owner] = max(0, deficits[owner] - int(count))
|
||||
|
||||
return CpLoadBackEvictionPlan(
|
||||
victims=tuple(victims),
|
||||
planned_freed_by_owner=tuple(planned_freed),
|
||||
remaining_deficit_by_owner=tuple(deficits),
|
||||
)
|
||||
|
||||
def _evict_cp_load_back_owner_lanes(
|
||||
self, plan: CpLoadBackPlan, *, node_id: int
|
||||
) -> CpLoadBackPlan:
|
||||
eviction_plan = self._plan_cp_load_back_owner_lane_evictions(plan)
|
||||
if any(v > 0 for v in eviction_plan.remaining_deficit_by_owner):
|
||||
logger.warning(
|
||||
"[CP_HICACHE_FALLBACK][cp_load_back_owner_lane_eviction_insufficient] "
|
||||
"node_id=%d required_by_owner=%s available_by_owner=%s "
|
||||
"deficit_by_owner=%s planned_freed_by_owner=%s "
|
||||
"remaining_deficit_by_owner=%s evictable_size=%d protected_size=%d "
|
||||
"victims=%s allocator_state=%s",
|
||||
node_id,
|
||||
plan.required_by_owner,
|
||||
plan.available_by_owner,
|
||||
plan.deficit_by_owner,
|
||||
eviction_plan.planned_freed_by_owner,
|
||||
eviction_plan.remaining_deficit_by_owner,
|
||||
int(getattr(self, "evictable_size_", 0)),
|
||||
int(getattr(self, "protected_size_", 0)),
|
||||
[getattr(node, "id", None) for node in eviction_plan.victims],
|
||||
self.token_to_kv_pool_allocator.allocator_state_str(),
|
||||
)
|
||||
return self._refresh_cp_load_back_plan(plan)
|
||||
|
||||
if len(eviction_plan.victims) == 0:
|
||||
return plan
|
||||
|
||||
num_evicted = 0
|
||||
write_back_nodes: List[TreeNode] = []
|
||||
for victim in eviction_plan.victims:
|
||||
victim_id = getattr(victim, "id", None)
|
||||
if not self._cp_device_leaf_is_load_back_victim(victim):
|
||||
raise RuntimeError(
|
||||
"CP HiCache load-back owner-lane eviction plan diverged "
|
||||
f"before reservation: node_id={node_id} victim_id={victim_id}"
|
||||
)
|
||||
|
||||
if victim.pin_expiry > 0 and time.monotonic() > victim.pin_expiry:
|
||||
self._clear_pin(victim)
|
||||
|
||||
if not self._node_backuped(victim):
|
||||
if self.cache_controller.write_policy == "write_back":
|
||||
logger.warning(
|
||||
"[CP_HICACHE_FALLBACK][cp_load_back_owner_lane_write_back] "
|
||||
"node_id=%d victim_id=%s reason=unbacked_write_back_victim",
|
||||
node_id,
|
||||
victim_id,
|
||||
)
|
||||
written = self.write_backup(victim, write_back=True)
|
||||
if written > 0:
|
||||
write_back_nodes.append(victim)
|
||||
continue
|
||||
num_evicted += self._evict_regular(victim)
|
||||
else:
|
||||
num_evicted += self._evict_backuped(victim)
|
||||
|
||||
for child in victim.parent.children.values():
|
||||
if child in write_back_nodes:
|
||||
continue
|
||||
if not child.evicted:
|
||||
break
|
||||
else:
|
||||
self._update_leaf_status(victim.parent)
|
||||
|
||||
if write_back_nodes:
|
||||
self.writing_check(write_back=True)
|
||||
for victim in write_back_nodes:
|
||||
if self._node_backuped(victim):
|
||||
num_evicted += self._evict_backuped(victim)
|
||||
|
||||
refreshed = self._refresh_cp_load_back_plan(plan)
|
||||
logger.info(
|
||||
"[HiCache-load] owner-lane device eviction before CP load-back: "
|
||||
"node_id=%d victims=%s num_evicted=%d required_by_owner=%s "
|
||||
"before_available_by_owner=%s before_deficit_by_owner=%s "
|
||||
"after_available_by_owner=%s after_deficit_by_owner=%s "
|
||||
"planned_freed_by_owner=%s",
|
||||
node_id,
|
||||
[getattr(node, "id", None) for node in eviction_plan.victims],
|
||||
num_evicted,
|
||||
plan.required_by_owner,
|
||||
plan.available_by_owner,
|
||||
plan.deficit_by_owner,
|
||||
refreshed.available_by_owner,
|
||||
refreshed.deficit_by_owner,
|
||||
eviction_plan.planned_freed_by_owner,
|
||||
)
|
||||
return refreshed
|
||||
|
||||
def _cp_capacity_debug_enabled(self) -> bool:
|
||||
return bool(
|
||||
getattr(self, "_cp_hicache_capacity_debug", False)
|
||||
@@ -2470,14 +2728,28 @@ class HiRadixCache(RadixCache):
|
||||
result = self.inc_lock_ref(ancester_node)
|
||||
delta = result.delta
|
||||
|
||||
# load it all or not at all
|
||||
host_hit_len = sum(self._node_host_len(n) for n in nodes_to_load)
|
||||
# load it all or not at all. The scalar length remains only a
|
||||
# coarse upper bound; final CP admission is the owner-lane vector
|
||||
# check below.
|
||||
try:
|
||||
load_back_plan = self._build_cp_load_back_plan(
|
||||
nodes_to_load, node_id=last_hit_node.id
|
||||
)
|
||||
except Exception:
|
||||
self.dec_lock_ref(ancester_node)
|
||||
raise
|
||||
host_hit_len = load_back_plan.host_hit_len
|
||||
logger.info(
|
||||
"[HiCache-load] load_back CP: node_id=%d nodes_to_load=%d host_hit_len=%d threshold=%d",
|
||||
"[HiCache-load] load_back CP: node_id=%d nodes_to_load=%d "
|
||||
"host_hit_len=%d threshold=%d required_by_owner=%s "
|
||||
"available_by_owner=%s deficit_by_owner=%s",
|
||||
last_hit_node.id,
|
||||
len(nodes_to_load),
|
||||
host_hit_len,
|
||||
self.load_back_threshold,
|
||||
load_back_plan.required_by_owner,
|
||||
load_back_plan.available_by_owner,
|
||||
load_back_plan.deficit_by_owner,
|
||||
)
|
||||
if host_hit_len < self.load_back_threshold or (
|
||||
host_hit_len > mem_quota + delta if mem_quota is not None else False
|
||||
@@ -2491,32 +2763,48 @@ class HiRadixCache(RadixCache):
|
||||
)
|
||||
return None
|
||||
|
||||
if any(v > 0 for v in load_back_plan.deficit_by_owner):
|
||||
load_back_plan = self._evict_cp_load_back_owner_lanes(
|
||||
load_back_plan, node_id=last_hit_node.id
|
||||
)
|
||||
|
||||
if any(v > 0 for v in load_back_plan.deficit_by_owner):
|
||||
self.dec_lock_ref(ancester_node)
|
||||
logger.warning(
|
||||
"[CP_HICACHE_FALLBACK][cp_load_back_owner_lane_capacity_failed] "
|
||||
"node_id=%d host_hit_len=%d nodes_to_load=%d "
|
||||
"required_by_owner=%s available_by_owner=%s "
|
||||
"deficit_by_owner=%s evictable_size=%d protected_size=%d "
|
||||
"ongoing_load_back_count=%d allocator_state=%s",
|
||||
last_hit_node.id,
|
||||
host_hit_len,
|
||||
len(nodes_to_load),
|
||||
load_back_plan.required_by_owner,
|
||||
load_back_plan.available_by_owner,
|
||||
load_back_plan.deficit_by_owner,
|
||||
int(getattr(self, "evictable_size_", 0)),
|
||||
int(getattr(self, "protected_size_", 0)),
|
||||
len(getattr(self, "ongoing_load_back", {})),
|
||||
self.token_to_kv_pool_allocator.allocator_state_str(),
|
||||
)
|
||||
return None
|
||||
|
||||
device_indices = self.cache_controller.load_cp(
|
||||
nodes_to_load, node_id=last_hit_node.id
|
||||
)
|
||||
if device_indices is None:
|
||||
logger.info(
|
||||
"[HiCache-load] load_back CP retry with lane-aware eviction: "
|
||||
"node_id=%d tokens_needed=%d",
|
||||
failed_plan = self._refresh_cp_load_back_plan(load_back_plan)
|
||||
logger.warning(
|
||||
"[CP_HICACHE_FALLBACK][cp_load_back_preflight_mismatch] "
|
||||
"node_id=%d host_hit_len=%d required_by_owner=%s "
|
||||
"available_by_owner=%s deficit_by_owner=%s "
|
||||
"allocator_state=%s",
|
||||
last_hit_node.id,
|
||||
host_hit_len,
|
||||
)
|
||||
# Lane-aware eviction: alloc_pages_with_owners failed because
|
||||
# at least one owner lane is short of free pages. Targeted
|
||||
# eviction frees pages from THOSE specific lanes, leaving
|
||||
# other lanes untouched. Mirrors cold prefill's retry path
|
||||
# in common.py:alloc_paged_token_slots_extend.
|
||||
_retry_page_owners: List[int] = []
|
||||
for _node in nodes_to_load:
|
||||
# page_owners is CPU int8; tolist() returns list[int].
|
||||
_retry_page_owners.extend(_node.cp_hicache.page_owners.tolist())
|
||||
_evict_for_compute_owner_lanes(
|
||||
tree_cache=self,
|
||||
allocator=self.token_to_kv_pool_allocator,
|
||||
page_compute_owners=_retry_page_owners,
|
||||
)
|
||||
device_indices = self.cache_controller.load_cp(
|
||||
nodes_to_load, node_id=last_hit_node.id
|
||||
failed_plan.required_by_owner,
|
||||
failed_plan.available_by_owner,
|
||||
failed_plan.deficit_by_owner,
|
||||
self.token_to_kv_pool_allocator.allocator_state_str(),
|
||||
)
|
||||
self.dec_lock_ref(ancester_node)
|
||||
if device_indices is None:
|
||||
|
||||
@@ -746,6 +746,32 @@ class KVCache(abc.ABC):
|
||||
def register_layer_transfer_counter(self, layer_transfer_counter: LayerDoneCounter):
|
||||
self.layer_transfer_counter = layer_transfer_counter
|
||||
|
||||
def wait_layer_transfer_on_stream(self, layer_id: int, stream) -> None:
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until_on_stream(
|
||||
layer_id - self.start_layer, stream
|
||||
)
|
||||
|
||||
def get_key_buffer_for_prefetch(self, layer_id: int, stream) -> torch.Tensor:
|
||||
self.wait_layer_transfer_on_stream(layer_id, stream)
|
||||
raw_getter = getattr(self, "_get_key_buffer", None)
|
||||
if raw_getter is None:
|
||||
raise NotImplementedError(
|
||||
f"{type(self).__name__} does not expose a non-blocking key-buffer getter"
|
||||
)
|
||||
return raw_getter(layer_id)
|
||||
|
||||
def get_index_k_with_scale_buffer_for_prefetch(
|
||||
self, layer_id: int, stream
|
||||
) -> torch.Tensor:
|
||||
self.wait_layer_transfer_on_stream(layer_id, stream)
|
||||
raw_getter = getattr(self, "_get_index_k_with_scale_buffer", None)
|
||||
if raw_getter is None:
|
||||
raise NotImplementedError(
|
||||
f"{type(self).__name__} does not expose a non-blocking index-buffer getter"
|
||||
)
|
||||
return raw_getter(layer_id)
|
||||
|
||||
def register_layer_backup_notifier(self, notifier):
|
||||
self.layer_backup_notifiers.append(notifier)
|
||||
|
||||
@@ -1379,6 +1405,10 @@ class HybridLinearKVPool(KVCache):
|
||||
layer_id = self._transfer_full_attention_id(layer_id)
|
||||
return self.full_kv_pool.get_key_buffer(layer_id)
|
||||
|
||||
def get_key_buffer_for_prefetch(self, layer_id: int, stream):
|
||||
layer_id = self._transfer_full_attention_id(layer_id)
|
||||
return self.full_kv_pool.get_key_buffer_for_prefetch(layer_id, stream)
|
||||
|
||||
def get_value_buffer(self, layer_id: int):
|
||||
layer_id = self._transfer_full_attention_id(layer_id)
|
||||
return self.full_kv_pool.get_value_buffer(layer_id)
|
||||
@@ -1552,25 +1582,29 @@ class MLATokenToKVPool(KVCache):
|
||||
|
||||
return _copy_buf_infos(self._contiguous_buf_infos)
|
||||
|
||||
def get_key_buffer(self, layer_id: int):
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
|
||||
def _get_key_buffer(self, layer_id: int):
|
||||
if self.store_dtype != self.dtype:
|
||||
return self.kv_buffer[layer_id - self.start_layer].view(self.dtype)
|
||||
|
||||
return self.kv_buffer[layer_id - self.start_layer]
|
||||
|
||||
def get_value_buffer(self, layer_id: int):
|
||||
def get_key_buffer(self, layer_id: int):
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
return self._get_key_buffer(layer_id)
|
||||
|
||||
def _get_value_buffer(self, layer_id: int):
|
||||
if self.store_dtype != self.dtype:
|
||||
return self.kv_buffer[layer_id - self.start_layer][
|
||||
..., : self.kv_lora_rank
|
||||
].view(self.dtype)
|
||||
return self.kv_buffer[layer_id - self.start_layer][..., : self.kv_lora_rank]
|
||||
|
||||
def get_value_buffer(self, layer_id: int):
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
return self._get_value_buffer(layer_id)
|
||||
|
||||
def get_kv_buffer(self, layer_id: int):
|
||||
return self.get_key_buffer(layer_id), self.get_value_buffer(layer_id)
|
||||
|
||||
@@ -1724,10 +1758,7 @@ class MLATokenToKVPoolFP4(MLATokenToKVPool):
|
||||
del self.kv_scale_buffer
|
||||
self._contiguous_buf_infos = None
|
||||
|
||||
def get_key_buffer(self, layer_id: int):
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
|
||||
def _get_key_buffer(self, layer_id: int):
|
||||
if self.store_dtype != self.dtype:
|
||||
cache_k_nope_fp4 = self.kv_buffer[layer_id - self.start_layer].view(
|
||||
torch.uint8
|
||||
@@ -1743,6 +1774,11 @@ class MLATokenToKVPoolFP4(MLATokenToKVPool):
|
||||
|
||||
return self.kv_buffer[layer_id - self.start_layer]
|
||||
|
||||
def get_key_buffer(self, layer_id: int):
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
return self._get_key_buffer(layer_id)
|
||||
|
||||
def set_kv_buffer(
|
||||
self,
|
||||
layer: RadixAttention,
|
||||
@@ -1893,10 +1929,13 @@ class NSATokenToKVPool(MLATokenToKVPool):
|
||||
]
|
||||
self._finalize_allocation_log(size)
|
||||
|
||||
def _get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
|
||||
return self.index_k_with_scale_buffer[layer_id - self.start_layer]
|
||||
|
||||
def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor:
|
||||
if self.layer_transfer_counter is not None:
|
||||
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
|
||||
return self.index_k_with_scale_buffer[layer_id - self.start_layer]
|
||||
return self._get_index_k_with_scale_buffer(layer_id)
|
||||
|
||||
def get_index_k_continuous(
|
||||
self,
|
||||
|
||||
Reference in New Issue
Block a user