Reduce CP shared-KV batch overhead without reverting bs1 planning

CP shared-KV bs>1 exposed two separate overhead sources: HiCache load-back could thrash near capacity, and partial-current sync compose could all-reduce row-major page-table gaps between request prefixes. Keep the intended batch-plan path for bs=1, but make host/L1 free-room handling less reactive and teach the MLA/index sync compose path to use exact per-request prefix slot spans instead of one bounding span.\n\nThe exact-span path preserves the single-span IPC fast path for bs=1/single-span cases, while avoiding over-communication for heterogeneous cache-hit batches. The HiCache metadata tests cover host/L1 free-room propagation and load-back batching behavior.\n\nConstraint: bs=1 using CPSharedKVBatchPlan is the expected steady-state path and must not be treated as a regression.\nConstraint: Remote production-like validation runs inside g0034 container /sgl-workspace/sglang-tai.\nRejected: Disable batch-plan for bs=1 | user confirmed this is intended behavior and it would hide the actual bs>1 overhead.\nRejected: Keep one bounding prefix span for bs>1 | row-major page tables can include current/gap slots and inflate per-layer all-reduce work.\nConfidence: medium\nScope-risk: moderate\nDirective: Do not replace exact prefix spans with a single row-major bounding span unless ETE data proves collective launch count dominates gap over-communication.\nTested: g0034 docker py_compile for changed Python/test files.\nTested: g0034 docker PYTHONPATH=python python -m pytest -q test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py test/registered/unit/mem_cache/test_cp_hicache_load_back_owner_lanes.py test/registered/unit/mem_cache/test_cp_hicache_metadata.py — 241 passed, 5 warnings, 2 subtests passed.\nNot-tested: Full ETE bs>1 throughput after the exact-prefix-span change.\nNot-tested: CUDA kernel benchmark and live traffic replay.
This commit is contained in:
laoyao0822
2026-06-06 02:03:16 +08:00
parent 1b99de7459
commit 81bfd4bfd3
10 changed files with 544 additions and 92 deletions

View File

@@ -2310,6 +2310,72 @@ def build_batch_prefix_slot_span(
return (start_slot, end_slot)
def build_batch_prefix_slot_spans(
*,
logical_pages: torch.Tensor,
prefix_lens_cpu,
page_size: int,
) -> list[tuple[int, int]]:
"""Return exact flattened page-table slot spans for batched prefix pages.
Unlike :func:`build_batch_prefix_slot_span`, this does not collapse multiple
request rows into one bounding span. Row-major page tables can have large
current/empty gaps between request prefixes; reducing those gaps is a
performance bug for cache-hit-heavy bs>1 prefill.
"""
if page_size <= 0:
raise ValueError(f"page_size must be positive, got {page_size}")
if prefix_lens_cpu is None:
raise ValueError("prefix_lens_cpu is required for batch prefix slot spans")
prefix_lens = [int(x) for x in prefix_lens_cpu]
if not prefix_lens:
return []
if logical_pages.dim() == 1:
if len(prefix_lens) != 1:
raise ValueError(
"1D logical_pages can only describe one request for prefix slot spans: "
f"batch_size={len(prefix_lens)} logical_pages_shape={tuple(logical_pages.shape)}"
)
pages_per_request = int(logical_pages.numel())
else:
if int(logical_pages.shape[0]) < len(prefix_lens):
raise ValueError(
"logical_pages has fewer rows than prefix_lens_cpu: "
f"rows={int(logical_pages.shape[0])} batch_size={len(prefix_lens)}"
)
pages_per_request = int(
logical_pages.reshape(logical_pages.shape[0], -1).shape[1]
)
spans: list[tuple[int, int]] = []
for req_id, prefix_len in enumerate(prefix_lens):
if prefix_len < 0:
raise ValueError(
f"prefix_lens_cpu contains negative prefix length: req_id={req_id} "
f"prefix_len={prefix_len}"
)
if prefix_len % page_size != 0:
raise ValueError(
"CP shared KV batch prefix slot spans require page-aligned prefixes: "
f"req_id={req_id} prefix_len={prefix_len} page_size={page_size}"
)
prefix_pages = prefix_len // page_size
if prefix_pages == 0:
continue
if prefix_pages > pages_per_request:
raise ValueError(
"prefix pages exceed per-request logical page-table width: "
f"req_id={req_id} prefix_pages={prefix_pages} "
f"pages_per_request={pages_per_request}"
)
req_start = req_id * pages_per_request
spans.append((req_start, req_start + prefix_pages))
return _merge_slot_spans(spans)
def _merge_slot_spans(spans: list[tuple[int, int]]) -> list[tuple[int, int]]:
normalized = sorted(
(int(start), int(end)) for start, end in spans if int(end) > int(start)
@@ -3402,6 +3468,7 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
page_size: int,
prefix_pages: int,
prefix_slot_span: tuple[int, int] | None = None,
prefix_slot_spans: list[tuple[int, int]] | None = None,
current_slot_spans: list[tuple[int, int]] | None = None,
layer_id: int | None = None,
nvtx_source: str = "mla.partial_current_sync",
@@ -3416,14 +3483,30 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
"""
total_slots = int(slot_remap.slot_logical_pages.numel())
if prefix_slot_span is None:
if prefix_slot_spans is not None and prefix_slot_span is not None:
raise ValueError(
"Specify either prefix_slot_span or prefix_slot_spans, not both."
)
if prefix_slot_spans is not None:
prefix_spans = _merge_slot_spans(prefix_slot_spans)
for prefix_start_slot, prefix_end_slot in prefix_spans:
if (
prefix_start_slot < 0
or prefix_end_slot < prefix_start_slot
or prefix_end_slot > total_slots
):
raise ValueError(
"Invalid CP shared KV partial-current prefix slot span: "
f"prefix_slot_span={(prefix_start_slot, prefix_end_slot)} "
f"total_slots={total_slots}"
)
elif prefix_slot_span is None:
if prefix_pages < 0 or prefix_pages > total_slots:
raise ValueError(
"Invalid CP shared KV partial-current prefix range: "
f"prefix_pages={prefix_pages} total_slots={total_slots}"
)
prefix_start_slot = 0
prefix_end_slot = int(prefix_pages)
prefix_spans = [(0, int(prefix_pages))] if int(prefix_pages) > 0 else []
else:
prefix_start_slot, prefix_end_slot = (
int(prefix_slot_span[0]),
@@ -3438,21 +3521,19 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
"Invalid CP shared KV partial-current prefix slot span: "
f"prefix_slot_span={prefix_slot_span} total_slots={total_slots}"
)
prefix_spans = (
[(prefix_start_slot, prefix_end_slot)]
if prefix_end_slot > prefix_start_slot
else []
)
dense_kv_cache = kv_cache.new_zeros(
(slot_remap.dense_num_pages * page_size, *kv_cache.shape[1:])
)
materialized_by_ipc = _try_tai_ipc_materialize_token_kv_page_slots_into(
kv_cache=kv_cache,
dense_kv_cache=dense_kv_cache,
slot_logical_pages=slot_remap.slot_logical_pages,
layout=layout,
page_size=page_size,
start_slot=prefix_start_slot,
end_slot=prefix_end_slot,
)
if not materialized_by_ipc:
materialize_local_token_kv_page_slots_into(
materialized_by_ipc = False
if len(prefix_spans) == 1:
prefix_start_slot, prefix_end_slot = prefix_spans[0]
materialized_by_ipc = _try_tai_ipc_materialize_token_kv_page_slots_into(
kv_cache=kv_cache,
dense_kv_cache=dense_kv_cache,
slot_logical_pages=slot_remap.slot_logical_pages,
@@ -3461,21 +3542,32 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
start_slot=prefix_start_slot,
end_slot=prefix_end_slot,
)
if not materialized_by_ipc:
for prefix_start_slot, prefix_end_slot in prefix_spans:
materialize_local_token_kv_page_slots_into(
kv_cache=kv_cache,
dense_kv_cache=dense_kv_cache,
slot_logical_pages=slot_remap.slot_logical_pages,
layout=layout,
page_size=page_size,
start_slot=prefix_start_slot,
end_slot=prefix_end_slot,
)
prefix_rows = slot_range_to_token_slice(
page_size,
prefix_start_slot,
prefix_end_slot,
)
_all_reduce_materialized_buffer_range(
dense_kv_cache,
layout.cp_size,
prefix_rows.start,
prefix_rows.stop,
nvtx_source=nvtx_source,
nvtx_layer_id=layer_id,
nvtx_cp_rank=layout.cp_rank,
)
prefix_rows = slot_range_to_token_slice(
page_size,
prefix_start_slot,
prefix_end_slot,
)
_all_reduce_materialized_buffer_range(
dense_kv_cache,
layout.cp_size,
prefix_rows.start,
prefix_rows.stop,
nvtx_source=nvtx_source,
nvtx_layer_id=layer_id,
nvtx_cp_rank=layout.cp_rank,
)
logical_locs = filter_locs_mappable_to_physical_pool(
logical_locs=logical_locs,
@@ -3499,7 +3591,7 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
)
if layout.cp_size > 1:
if current_slot_spans is None:
if prefix_slot_span is not None:
if prefix_slot_span is not None or prefix_slot_spans is not None:
raise ValueError(
"CP shared KV batched current compose requires explicit "
"current_slot_spans to avoid reducing prefix slots twice."
@@ -3541,6 +3633,7 @@ def materialize_prefix_and_reuse_current_index_page_slots(
index_head_dim: int,
prefix_pages: int,
prefix_slot_span: tuple[int, int] | None = None,
prefix_slot_spans: list[tuple[int, int]] | None = None,
current_slot_spans: list[tuple[int, int]] | None = None,
layer_id: int | None = None,
nvtx_source: str = "index.partial_current_sync",
@@ -3548,14 +3641,30 @@ def materialize_prefix_and_reuse_current_index_page_slots(
"""Synchronously compose prefix index materialization with current index rows."""
total_slots = int(slot_remap.slot_logical_pages.numel())
if prefix_slot_span is None:
if prefix_slot_spans is not None and prefix_slot_span is not None:
raise ValueError(
"Specify either prefix_slot_span or prefix_slot_spans, not both."
)
if prefix_slot_spans is not None:
prefix_spans = _merge_slot_spans(prefix_slot_spans)
for prefix_start_slot, prefix_end_slot in prefix_spans:
if (
prefix_start_slot < 0
or prefix_end_slot < prefix_start_slot
or prefix_end_slot > total_slots
):
raise ValueError(
"Invalid CP shared KV index partial-current prefix slot span: "
f"prefix_slot_span={(prefix_start_slot, prefix_end_slot)} "
f"total_slots={total_slots}"
)
elif prefix_slot_span is None:
if prefix_pages < 0 or prefix_pages > total_slots:
raise ValueError(
"Invalid CP shared KV index partial-current prefix range: "
f"prefix_pages={prefix_pages} total_slots={total_slots}"
)
prefix_start_slot = 0
prefix_end_slot = int(prefix_pages)
prefix_spans = [(0, int(prefix_pages))] if int(prefix_pages) > 0 else []
else:
prefix_start_slot, prefix_end_slot = (
int(prefix_slot_span[0]),
@@ -3570,20 +3679,19 @@ def materialize_prefix_and_reuse_current_index_page_slots(
"Invalid CP shared KV index partial-current prefix slot span: "
f"prefix_slot_span={prefix_slot_span} total_slots={total_slots}"
)
prefix_spans = (
[(prefix_start_slot, prefix_end_slot)]
if prefix_end_slot > prefix_start_slot
else []
)
dense_page_buffer = page_buffer.new_zeros(
(slot_remap.dense_num_pages, *page_buffer.shape[1:])
)
materialized_by_ipc = _try_tai_ipc_materialize_paged_buffer_page_slots_into(
page_buffer=page_buffer,
dense_page_buffer=dense_page_buffer,
slot_logical_pages=slot_remap.slot_logical_pages,
layout=layout,
start_slot=prefix_start_slot,
end_slot=prefix_end_slot,
)
if not materialized_by_ipc:
materialize_local_paged_buffer_page_slots_into(
materialized_by_ipc = False
if len(prefix_spans) == 1:
prefix_start_slot, prefix_end_slot = prefix_spans[0]
materialized_by_ipc = _try_tai_ipc_materialize_paged_buffer_page_slots_into(
page_buffer=page_buffer,
dense_page_buffer=dense_page_buffer,
slot_logical_pages=slot_remap.slot_logical_pages,
@@ -3591,16 +3699,26 @@ def materialize_prefix_and_reuse_current_index_page_slots(
start_slot=prefix_start_slot,
end_slot=prefix_end_slot,
)
prefix_rows = slot_range_to_page_slice(prefix_start_slot, prefix_end_slot)
_all_reduce_materialized_buffer_range(
dense_page_buffer,
layout.cp_size,
prefix_rows.start,
prefix_rows.stop,
nvtx_source=nvtx_source,
nvtx_layer_id=layer_id,
nvtx_cp_rank=layout.cp_rank,
)
if not materialized_by_ipc:
for prefix_start_slot, prefix_end_slot in prefix_spans:
materialize_local_paged_buffer_page_slots_into(
page_buffer=page_buffer,
dense_page_buffer=dense_page_buffer,
slot_logical_pages=slot_remap.slot_logical_pages,
layout=layout,
start_slot=prefix_start_slot,
end_slot=prefix_end_slot,
)
prefix_rows = slot_range_to_page_slice(prefix_start_slot, prefix_end_slot)
_all_reduce_materialized_buffer_range(
dense_page_buffer,
layout.cp_size,
prefix_rows.start,
prefix_rows.stop,
nvtx_source=nvtx_source,
nvtx_layer_id=layer_id,
nvtx_cp_rank=layout.cp_rank,
)
dense_page_buffer = fill_current_index_page_slots(
dense_page_buffer=dense_page_buffer,
current_index_k=current_index_k,
@@ -3612,7 +3730,7 @@ def materialize_prefix_and_reuse_current_index_page_slots(
)
if layout.cp_size > 1:
if current_slot_spans is None:
if prefix_slot_span is not None:
if prefix_slot_span is not None or prefix_slot_spans is not None:
raise ValueError(
"CP shared KV batched index current compose requires explicit "
"current_slot_spans to avoid reducing prefix slots twice."

View File

@@ -17,7 +17,7 @@ from sglang.srt.environ import envs
from sglang.srt.layers.attention.nsa import index_buf_accessor
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
build_batch_current_slot_spans,
build_batch_prefix_slot_span,
build_batch_prefix_slot_spans,
cp_shared_kv_debug_enabled,
cp_shared_kv_debug_log,
cp_shared_kv_mla_prefetch_enabled,
@@ -593,7 +593,7 @@ class Indexer(MultiPlatformOp):
f"current_scale_shape={tuple(current_index_kv[1].shape)} "
f"out_cache_loc_shape={tuple(current_locs.shape)}"
)
prefix_slot_span = None
prefix_slot_spans = None
current_slot_spans = build_batch_current_slot_spans(
logical_pages=logical_page_table,
prefix_lens_cpu=prefix_lens_cpu,
@@ -604,7 +604,7 @@ class Indexer(MultiPlatformOp):
prefix_pages = int(prefix_lens_cpu[0]) // page_size
else:
prefix_pages = 0
prefix_slot_span = build_batch_prefix_slot_span(
prefix_slot_spans = build_batch_prefix_slot_spans(
logical_pages=logical_page_table,
prefix_lens_cpu=prefix_lens_cpu,
page_size=page_size,
@@ -663,7 +663,7 @@ class Indexer(MultiPlatformOp):
page_size=page_size,
index_head_dim=forward_batch.token_to_kv_pool.index_head_dim,
prefix_pages=prefix_pages,
prefix_slot_span=prefix_slot_span,
prefix_slot_spans=prefix_slot_spans,
current_slot_spans=current_slot_spans,
layer_id=layer_id,
)
@@ -675,13 +675,13 @@ class Indexer(MultiPlatformOp):
cp_shared_kv_mla_prefetch_log(
"index_partial_current_sync_compose cp_rank=%s layer=%s "
"prefix_lens=%s extend_lens=%s prefix_pages=%s "
"prefix_slot_span=%s current_rows=%s dense_pages=%s",
"prefix_slot_spans=%s current_rows=%s dense_pages=%s",
layout.cp_rank,
layer_id,
prefix_lens,
extend_lens,
prefix_pages,
prefix_slot_span,
prefix_slot_spans,
int(current_index_kv[0].shape[0]),
int(materialized.shape[0]),
)

View File

@@ -16,7 +16,7 @@ from sglang.srt.layers.attention.nsa.cp_shared_kv_prefetch import (
)
from sglang.srt.layers.attention.nsa.cp_shared_kv_runtime import (
build_batch_current_slot_spans,
build_batch_prefix_slot_span,
build_batch_prefix_slot_spans,
build_current_loc_remap,
cp_shared_kv_debug_enabled,
cp_shared_kv_debug_log,
@@ -2137,7 +2137,7 @@ class NativeSparseAttnBackend(
f"current_locs_shape={tuple(current_locs_for_reuse.shape)} "
f"page_size={page_size}"
)
prefix_slot_span = None
prefix_slot_spans = None
current_slot_spans = build_batch_current_slot_spans(
logical_pages=metadata.real_page_table,
prefix_lens_cpu=prefix_lens_cpu,
@@ -2148,7 +2148,7 @@ class NativeSparseAttnBackend(
prefix_pages = int(prefix_lens_cpu[0]) // page_size
else:
prefix_pages = 0
prefix_slot_span = build_batch_prefix_slot_span(
prefix_slot_spans = build_batch_prefix_slot_spans(
logical_pages=metadata.real_page_table,
prefix_lens_cpu=prefix_lens_cpu,
page_size=page_size,
@@ -2170,7 +2170,7 @@ class NativeSparseAttnBackend(
layout=forward_batch.cp_shared_kv_layout,
page_size=page_size,
prefix_pages=prefix_pages,
prefix_slot_span=prefix_slot_span,
prefix_slot_spans=prefix_slot_spans,
current_slot_spans=current_slot_spans,
layer_id=layer.layer_id,
)
@@ -2184,7 +2184,7 @@ class NativeSparseAttnBackend(
cp_shared_kv_mla_prefetch_log(
"forward_partial_current_sync_compose cp_rank=%s "
"layer=%s reason=%s prefix_lens=%s extend_lens=%s "
"prefix_pages=%s prefix_slot_span=%s "
"prefix_pages=%s prefix_slot_spans=%s "
"current_rows=%s kv_rows=%s page_table_shape=%s",
forward_batch.cp_shared_kv_layout.cp_rank,
layer.layer_id,
@@ -2192,7 +2192,7 @@ class NativeSparseAttnBackend(
prefix_lens,
extend_lens,
prefix_pages,
prefix_slot_span,
prefix_slot_spans,
int(current_kv_cache.shape[0]),
int(kv_cache.shape[0]),
tuple(page_table_1.shape),

View File

@@ -1436,6 +1436,16 @@ class HiCacheController:
return device_indices
def load_cp(self, nodes_to_load, node_id: int = -1) -> Optional[torch.Tensor]:
start_time = time.perf_counter()
stage_start_time = start_time
stage_durations_ms: List[tuple[str, float]] = []
def record_stage(stage: str) -> None:
nonlocal stage_start_time
now = time.perf_counter()
stage_durations_ms.append((stage, (now - stage_start_time) * 1000.0))
stage_start_time = now
# Reproduce the original (write-time) CP owner pattern. Each node
# carries `page_owners` (one int8 per logical page, identical on all
# CP ranks) so we can ask the allocator for a fresh device range
@@ -1455,10 +1465,12 @@ class HiCacheController:
)
# page_owners is CPU int8; .tolist() returns list[int] directly.
page_owners.extend(meta.page_owners.tolist())
record_stage("collect_page_owners")
device_indices = self.mem_pool_device_allocator.alloc_pages_with_owners(
page_owners
)
record_stage("alloc_pages_with_owners")
# Fail closed: returning None lets the caller drop to cache miss
# (cold prefill). Never proceed with a non-matching owner pattern.
if device_indices is None:
@@ -1537,6 +1549,7 @@ class HiCacheController:
host_chunks.append(meta.host_indices)
if draft_host_indices is not None:
draft_host_chunks.append(draft_host_indices)
record_stage("build_chunks")
visible_device_indices = (
torch.cat(visible_chunks)
@@ -1566,6 +1579,7 @@ class HiCacheController:
draft_host_indices = None
if self.has_draft_hicache:
draft_host_indices = torch.cat(draft_host_chunks)
record_stage("concat_indices")
try:
self._validate_cp_hicache_page_indices(
@@ -1575,6 +1589,7 @@ class HiCacheController:
self._validate_cp_hicache_page_indices(
draft_host_indices, physical_device_indices
)
record_stage("validate_indices")
except Exception:
self.mem_pool_device_allocator.free(device_indices)
raise
@@ -1595,6 +1610,19 @@ class HiCacheController:
)
)
total_duration = time.perf_counter() - start_time
if total_duration >= 1.0:
logger.warning(
"[HiCache-load] slow load_cp planning: node_id=%d duration_ms=%.3f "
"pages=%d host_indices=%d physical_indices=%d draft=%s stages_ms=%s",
node_id,
total_duration * 1000.0,
len(page_owners),
int(host_indices.numel()),
int(physical_device_indices.numel()),
draft_host_indices is not None,
[(stage, round(ms, 3)) for stage, ms in stage_durations_ms],
)
return visible_device_indices
def move_indices(self, op: CacheOperation, mem_pool_host=None):

View File

@@ -356,7 +356,12 @@ class CpLoadBackPlan:
page_owners: List[int]
required_by_owner: List[int]
available_by_owner: List[int]
# Exact capacity deficit. This is the only deficit allowed to block the
# synchronous load-back admission path.
deficit_by_owner: List[int]
# Advisory free-room deficit. This is logged for observability but must not
# force synchronous load-back eviction when exact capacity already fits.
free_room_deficit_by_owner: List[int]
host_hit_len: int
@@ -1017,39 +1022,36 @@ class HiRadixCache(RadixCache):
f"page_size={self.page_size} page_owners={len(page_owners)}"
)
required, available, deficits = self._cp_load_back_owner_lane_stats(page_owners)
required, available, deficits, free_room_deficits = (
self._cp_load_back_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],
free_room_deficit_by_owner=[int(v) for v in free_room_deficits],
host_hit_len=host_hit_len,
)
def _cp_load_back_owner_lane_stats(
self, page_owners: List[int]
) -> Tuple[List[int], List[int], List[int]]:
) -> Tuple[List[int], List[int], List[int], List[int]]:
allocator = self.token_to_kv_pool_allocator
target_ratio = float(getattr(self, "hicache_l1_free_room_ratio", 0.0) or 0.0)
trigger_ratio = float(
getattr(self, "hicache_l1_free_room_trigger_ratio", 0.0) or 0.0
)
free_room_stats = getattr(allocator, "compute_owner_lane_free_room_stats", None)
if free_room_stats is not None:
return free_room_stats(
page_owners,
target_ratio=target_ratio,
trigger_ratio=trigger_ratio,
)
required, available, exact_deficits = allocator.compute_owner_lane_stats(
page_owners
)
lane_capacity_pages = getattr(allocator, "compute_owner_lane_capacity_pages", None)
if lane_capacity_pages is None:
return required, available, exact_deficits
return required, available, exact_deficits, exact_deficits
return (
required,
available,
exact_deficits,
compute_owner_lane_free_room_deficits(
required=required,
available=available,
@@ -1060,14 +1062,15 @@ class HiRadixCache(RadixCache):
)
def _refresh_cp_load_back_plan(self, plan: CpLoadBackPlan) -> CpLoadBackPlan:
required, available, deficits = self._cp_load_back_owner_lane_stats(
plan.page_owners
required, available, deficits, free_room_deficits = (
self._cp_load_back_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],
free_room_deficit_by_owner=[int(v) for v in free_room_deficits],
host_hit_len=plan.host_hit_len,
)
@@ -1327,6 +1330,7 @@ class HiRadixCache(RadixCache):
required_by_owner=[],
available_by_owner=[],
deficit_by_owner=deficits,
free_room_deficit_by_owner=deficits,
host_hit_len=0,
)
eviction_plan = self._plan_cp_load_back_owner_lane_evictions(plan)
@@ -3362,6 +3366,15 @@ class HiRadixCache(RadixCache):
if self._uses_cp_hicache:
start_time = time.perf_counter()
stage_start_time = start_time
stage_durations_ms: List[Tuple[str, float]] = []
def record_stage(stage: str) -> None:
nonlocal stage_start_time
now = time.perf_counter()
stage_durations_ms.append((stage, (now - stage_start_time) * 1000.0))
stage_start_time = now
last_hit_node = node
nodes_to_load = []
while node.evicted:
@@ -3388,6 +3401,7 @@ class HiRadixCache(RadixCache):
load_back_plan = self._build_cp_load_back_plan(
nodes_to_load, node_id=last_hit_node.id
)
record_stage("build_plan")
except Exception:
self.dec_lock_ref(ancester_node)
raise
@@ -3395,7 +3409,8 @@ class HiRadixCache(RadixCache):
logger.info(
"[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",
"available_by_owner=%s deficit_by_owner=%s "
"free_room_deficit_by_owner=%s",
last_hit_node.id,
len(nodes_to_load),
host_hit_len,
@@ -3403,6 +3418,7 @@ class HiRadixCache(RadixCache):
load_back_plan.required_by_owner,
load_back_plan.available_by_owner,
load_back_plan.deficit_by_owner,
load_back_plan.free_room_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
@@ -3420,6 +3436,7 @@ class HiRadixCache(RadixCache):
load_back_plan = self._evict_cp_load_back_owner_lanes(
load_back_plan, node_id=last_hit_node.id
)
record_stage("owner_lane_evict")
if any(v > 0 for v in load_back_plan.deficit_by_owner):
self.dec_lock_ref(ancester_node)
@@ -3427,7 +3444,8 @@ class HiRadixCache(RadixCache):
"[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 "
"deficit_by_owner=%s free_room_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,
@@ -3435,6 +3453,7 @@ class HiRadixCache(RadixCache):
load_back_plan.required_by_owner,
load_back_plan.available_by_owner,
load_back_plan.deficit_by_owner,
load_back_plan.free_room_deficit_by_owner,
int(getattr(self, "evictable_size_", 0)),
int(getattr(self, "protected_size_", 0)),
len(getattr(self, "ongoing_load_back", {})),
@@ -3445,18 +3464,21 @@ class HiRadixCache(RadixCache):
device_indices = self.cache_controller.load_cp(
nodes_to_load, node_id=last_hit_node.id
)
record_stage("load_cp_plan")
if device_indices is None:
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 "
"free_room_deficit_by_owner=%s "
"allocator_state=%s",
last_hit_node.id,
host_hit_len,
failed_plan.required_by_owner,
failed_plan.available_by_owner,
failed_plan.deficit_by_owner,
failed_plan.free_room_deficit_by_owner,
self.token_to_kv_pool_allocator.allocator_state_str(),
)
self.dec_lock_ref(ancester_node)
@@ -3482,6 +3504,7 @@ class HiRadixCache(RadixCache):
physical_loaded_len += int(
getattr(metadata, "padded_len", host_len)
)
record_stage("assign_tree")
self.evictable_size_ += physical_loaded_len
self.inc_lock_ref(last_hit_node)
@@ -3499,6 +3522,22 @@ class HiRadixCache(RadixCache):
len(device_indices),
physical_loaded_len,
)
total_duration = time.perf_counter() - start_time
if total_duration >= 1.0:
logger.warning(
"[HiCache-load] slow CP load_back planning: node_id=%d "
"duration_ms=%.3f host_hit_len=%d required_by_owner=%s "
"available_by_owner=%s exact_deficit_by_owner=%s "
"free_room_deficit_by_owner=%s stages_ms=%s",
last_hit_node.id,
total_duration * 1000.0,
host_hit_len,
load_back_plan.required_by_owner,
load_back_plan.available_by_owner,
load_back_plan.deficit_by_owner,
load_back_plan.free_room_deficit_by_owner,
[(stage, round(ms, 3)) for stage, ms in stage_durations_ms],
)
return device_indices
start_time = time.perf_counter()