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