Avoid redundant CP collectives on sync shared-KV materialize

Synchronous CP shared KV full-hit and partial-current paths can now use the tai-kernel CUDA IPC slot-dense materialize path instead of first copying local owner pages and then running a dense CP all-reduce. This keeps the async prefetch pipeline unchanged while routing the safer synchronous runtime path through descriptor-driven owner-page reads.

The runtime builds descriptors from the existing page-aligned slot contract, caches peer pointer tables for long-lived KV/index buffers, and falls back with explicit warnings if the tai-kernel IPC capability is missing. Unit coverage locks offset handle exchange, descriptor construction, and both MLA/index full and partial-current calls.

Constraint: Async prefetch currently has unresolved scheduling contention with GEMM/MoE and current-layer KV collectives, so this commit intentionally does not wire IPC into prefetch.

Constraint: tai-kernel must provide offset-aware CUDA IPC symbols from af9fb67.

Rejected: Replace prefetch all-reduce in this slice | Nsight shows prefetch timing and communicator contention need a separate scheduling design.

Rejected: Fail-fast on missing tai IPC immediately | remote ETE still needs to validate production deployment capability before removing the warning fallback.

Confidence: medium

Scope-risk: moderate

Directive: Keep prefetch and synchronous materialize decisions separate until prefetch has a low-SM/copy-engine schedule.

Related: tai-kernel af9fb67

Tested: git diff --check; remote runtime/unit evidence recorded in docs/advanced_features/nsa_prefill_cp_page_aligned_cache_contract.md.

Not-tested: Fresh GLM5 ETE after this commit; async prefetch IPC path.

Co-authored-by: OmX <omx@oh-my-codex.dev>
This commit is contained in:
laoyao0822
2026-05-31 23:59:23 +08:00
parent b4f5c3bc0e
commit 4342de0463
3 changed files with 844 additions and 50 deletions

View File

@@ -19,9 +19,11 @@ logger = logging.getLogger(__name__)
_DEBUG_LOG_COUNTS: dict[str, int] = {}
_TAI_MATERIALIZE_FALLBACK_LOG_COUNTS: dict[str, int] = {}
_TAI_IPC_MATERIALIZE_FALLBACK_LOG_COUNTS: dict[str, int] = {}
_TAI_FUSED_MLA_STORE_FALLBACK_LOG_COUNTS: dict[str, int] = {}
_CURRENT_REUSE_FALLBACK_LOG_COUNTS: dict[str, int] = {}
_SLOT_REMAP_CACHE_LOG_COUNTS: dict[str, int] = {}
_TAI_IPC_PEER_PTR_CACHE: dict[tuple[object, ...], torch.Tensor] = {}
_MLA_PREFETCH_LOG_PROBE_LAYER = 2
_MLA_PREFETCH_DEFAULT_MIN_PREFIX_TOKENS = max(
int(envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_PREFIX_TOKENS.get()),
@@ -300,6 +302,38 @@ def _load_tai_materialize_kernels():
return None
@lru_cache(maxsize=1)
def _load_tai_ipc_kernels():
try:
from tai_kernel.nsa_prefill import ipc as tai_ipc
required = (
"get_cuda_ipc_mem_handle_with_offset",
"open_cuda_ipc_mem_handles_with_offsets",
"materialize_cuda_ipc_peer_pages_slot_dense",
)
missing = [name for name in required if not hasattr(tai_ipc, name)]
if missing:
_log_tai_ipc_materialize_fallback(
"missing_symbols",
"CP shared KV tai IPC materialize kernels are incomplete; "
"falling back to collective materialize. missing=%s",
missing,
limit=1,
)
return None
return tai_ipc
except Exception as exc:
_log_tai_ipc_materialize_fallback(
"import_failed",
"CP shared KV tai IPC materialize kernels are unavailable; "
"falling back to collective materialize. error=%s",
exc,
limit=1,
)
return None
@lru_cache(maxsize=16)
def _tai_current_slot_fill_sparse_pages_self_test(
device_type: str,
@@ -480,6 +514,23 @@ def _log_tai_materialize_fallback(
)
def _log_tai_ipc_materialize_fallback(
key: str,
message: str,
*args,
limit: int = 8,
) -> None:
count = _TAI_IPC_MATERIALIZE_FALLBACK_LOG_COUNTS.get(key, 0)
if count >= limit:
return
_TAI_IPC_MATERIALIZE_FALLBACK_LOG_COUNTS[key] = count + 1
logger.warning(
"[CP_SHARED_KV_FALLBACK][tai_ipc_materialize] reason=%s " + message,
key,
*args,
)
def _log_tai_fused_mla_store_fallback(
key: str,
message: str,
@@ -1120,6 +1171,266 @@ def _try_tai_materialize_token_kv_page_slots_into(
return False
def build_cp_shared_kv_ipc_page_descriptors(
slot_logical_pages: torch.Tensor,
layout: CpSharedKVLayout,
*,
physical_page_capacity: int | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Build owner/source descriptors for slot-dense CUDA IPC materialize.
The returned vectors are flattened in consumer slot order. Source page
indices are physical page ids in the peer's CP-local pool, including dummy
page 0. Invalid/sentinel slots are encoded as ``-1`` owner/source and are
zero-filled by the IPC kernel.
"""
logical_pages = slot_logical_pages.reshape(-1).to(torch.long)
owner_ranks = layout.owner_for_logical_pages(logical_pages).to(torch.long)
src_page_indices = layout.logical_pages_to_physical(logical_pages).to(torch.long)
invalid = logical_pages <= 0
if physical_page_capacity is not None:
invalid = invalid | (src_page_indices >= int(physical_page_capacity))
invalid_value = torch.full_like(owner_ranks, -1)
owner_ranks = torch.where(invalid, invalid_value, owner_ranks)
src_page_indices = torch.where(invalid, invalid_value, src_page_indices)
return owner_ranks.contiguous(), src_page_indices.contiguous()
def _tai_ipc_peer_ptr_cache_key(
tensor: torch.Tensor,
layout: CpSharedKVLayout,
cp_group: Any,
) -> tuple[object, ...]:
return (
int(tensor.untyped_storage().data_ptr()),
int(tensor.data_ptr()),
tuple(int(dim) for dim in tensor.shape),
tuple(int(stride) for stride in tensor.stride()),
str(tensor.dtype),
str(tensor.device),
int(layout.cp_size),
int(layout.cp_rank),
getattr(cp_group, "unique_name", None),
)
def _get_or_open_tai_ipc_peer_ptrs(
tensor: torch.Tensor,
layout: CpSharedKVLayout,
) -> tuple[Any, torch.Tensor] | None:
if layout.cp_size <= 1:
return None
if not _tai_materialize_runtime_enabled():
return None
if not tensor.is_cuda or not tensor.is_contiguous():
_log_tai_ipc_materialize_fallback(
"unsupported_tensor",
"CP shared KV tai IPC materialize requires a contiguous CUDA tensor; "
"falling back to collective materialize. device=%s contiguous=%s",
tensor.device,
tensor.is_contiguous(),
)
return None
kernels = _load_tai_ipc_kernels()
if kernels is None:
return None
cp_group = get_attention_cp_group()
if int(getattr(cp_group, "world_size", layout.cp_size)) != int(layout.cp_size):
_log_tai_ipc_materialize_fallback(
"group_size_mismatch",
"CP shared KV tai IPC materialize cp group size mismatch; "
"falling back to collective materialize. layout_cp_size=%s group_world_size=%s",
layout.cp_size,
getattr(cp_group, "world_size", None),
limit=1,
)
return None
cache_key = _tai_ipc_peer_ptr_cache_key(tensor, layout, cp_group)
cached = _TAI_IPC_PEER_PTR_CACHE.get(cache_key)
if cached is not None:
return kernels, cached
try:
handle, offset = kernels.get_cuda_ipc_mem_handle_with_offset(tensor)
handle_cuda = handle.to(device=tensor.device, non_blocking=False)
offset_cuda = torch.tensor(
[int(offset)], dtype=torch.int64, device=tensor.device
)
gathered_handles = torch.empty(
(layout.cp_size, int(handle.numel())),
dtype=torch.uint8,
device=tensor.device,
)
gathered_offsets = torch.empty(
(layout.cp_size,),
dtype=torch.int64,
device=tensor.device,
)
cp_group.all_gather_into_tensor(gathered_handles, handle_cuda)
cp_group.all_gather_into_tensor(gathered_offsets, offset_cuda)
peer_ptrs = kernels.open_cuda_ipc_mem_handles_with_offsets(
gathered_handles.cpu(),
gathered_offsets.cpu(),
tensor,
self_rank=layout.cp_rank,
)
_TAI_IPC_PEER_PTR_CACHE[cache_key] = peer_ptrs
return kernels, peer_ptrs
except Exception as exc:
_log_tai_ipc_materialize_fallback(
"peer_open_failed",
"CP shared KV tai IPC peer pointer setup failed; falling back to "
"collective materialize. cp_rank=%s cp_size=%s tensor_shape=%s "
"dtype=%s error=%s",
layout.cp_rank,
layout.cp_size,
tuple(tensor.shape),
tensor.dtype,
exc,
limit=4,
)
return None
def _page_nbytes_from_page_tensor(tensor: torch.Tensor) -> int:
if tensor.ndim == 0:
raise ValueError("CP shared KV IPC page tensor must have at least one dim")
return int(tensor.stride(0)) * int(tensor.element_size())
def _token_kv_page_nbytes(kv_cache: torch.Tensor, page_size: int) -> int:
return int(page_size) * _page_nbytes_from_page_tensor(kv_cache)
def _try_tai_ipc_materialize_token_kv_page_slots_into(
*,
kv_cache: torch.Tensor,
dense_kv_cache: torch.Tensor,
slot_logical_pages: torch.Tensor,
layout: CpSharedKVLayout,
page_size: int,
start_slot: int,
end_slot: int,
) -> bool:
if start_slot == end_slot:
return True
if start_slot != 0:
return False
if not dense_kv_cache.is_cuda or not dense_kv_cache.is_contiguous():
return False
ipc_state = _get_or_open_tai_ipc_peer_ptrs(kv_cache, layout)
if ipc_state is None:
return False
kernels, peer_ptrs = ipc_state
flat_slot_logical_pages = slot_logical_pages.reshape(-1)
slot_logical_pages_range = _contiguous_for_tai(
flat_slot_logical_pages[start_slot:end_slot]
)
if slot_logical_pages_range.numel() == 0:
return True
try:
owner_ranks, src_page_indices = build_cp_shared_kv_ipc_page_descriptors(
slot_logical_pages_range,
layout,
physical_page_capacity=kv_cache.shape[0] // page_size,
)
kernels.materialize_cuda_ipc_peer_pages_slot_dense(
peer_ptrs,
dense_kv_cache,
owner_ranks,
src_page_indices,
page_nbytes=_token_kv_page_nbytes(kv_cache, page_size),
)
return True
except Exception as exc:
_log_tai_ipc_materialize_fallback(
"token_kernel_failed",
"CP shared KV tai IPC token materialize failed; falling back to "
"collective materialize. cp_rank=%s cp_size=%s start_slot=%s "
"end_slot=%s slot_count=%s page_size=%s kv_shape=%s dense_shape=%s "
"error=%s",
layout.cp_rank,
layout.cp_size,
start_slot,
end_slot,
int(slot_logical_pages_range.numel()),
page_size,
tuple(kv_cache.shape),
tuple(dense_kv_cache.shape),
exc,
)
return False
def _try_tai_ipc_materialize_paged_buffer_page_slots_into(
*,
page_buffer: torch.Tensor,
dense_page_buffer: torch.Tensor,
slot_logical_pages: torch.Tensor,
layout: CpSharedKVLayout,
start_slot: int,
end_slot: int,
) -> bool:
if start_slot == end_slot:
return True
if start_slot != 0:
return False
if not dense_page_buffer.is_cuda or not dense_page_buffer.is_contiguous():
return False
ipc_state = _get_or_open_tai_ipc_peer_ptrs(page_buffer, layout)
if ipc_state is None:
return False
kernels, peer_ptrs = ipc_state
flat_slot_logical_pages = slot_logical_pages.reshape(-1)
slot_logical_pages_range = _contiguous_for_tai(
flat_slot_logical_pages[start_slot:end_slot]
)
if slot_logical_pages_range.numel() == 0:
return True
try:
owner_ranks, src_page_indices = build_cp_shared_kv_ipc_page_descriptors(
slot_logical_pages_range,
layout,
physical_page_capacity=page_buffer.shape[0],
)
kernels.materialize_cuda_ipc_peer_pages_slot_dense(
peer_ptrs,
dense_page_buffer,
owner_ranks,
src_page_indices,
page_nbytes=_page_nbytes_from_page_tensor(page_buffer),
)
return True
except Exception as exc:
_log_tai_ipc_materialize_fallback(
"paged_kernel_failed",
"CP shared KV tai IPC paged materialize failed; falling back to "
"collective materialize. cp_rank=%s cp_size=%s start_slot=%s "
"end_slot=%s slot_count=%s page_shape=%s dense_shape=%s error=%s",
layout.cp_rank,
layout.cp_size,
start_slot,
end_slot,
int(slot_logical_pages_range.numel()),
tuple(page_buffer.shape),
tuple(dense_page_buffer.shape),
exc,
)
return False
def is_current_only_extend_batch(forward_batch) -> bool:
"""Return whether an extend batch has no cached/history tokens.
@@ -2256,7 +2567,7 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
dense_kv_cache = kv_cache.new_zeros(
(slot_remap.dense_num_pages * page_size, *kv_cache.shape[1:])
)
materialize_local_token_kv_page_slots_into(
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,
@@ -2265,17 +2576,27 @@ def materialize_prefix_and_reuse_current_kv_page_slots(
start_slot=0,
end_slot=prefix_pages,
)
if not materialized_by_ipc:
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=0,
end_slot=prefix_pages,
)
prefix_rows = slot_range_to_token_slice(page_size, 0, prefix_pages)
_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, 0, prefix_pages)
_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,
@@ -2326,7 +2647,7 @@ def materialize_prefix_and_reuse_current_index_page_slots(
dense_page_buffer = page_buffer.new_zeros(
(slot_remap.dense_num_pages, *page_buffer.shape[1:])
)
materialize_local_paged_buffer_page_slots_into(
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,
@@ -2334,16 +2655,25 @@ def materialize_prefix_and_reuse_current_index_page_slots(
start_slot=0,
end_slot=prefix_pages,
)
prefix_rows = slot_range_to_page_slice(0, prefix_pages)
_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:
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=0,
end_slot=prefix_pages,
)
prefix_rows = slot_range_to_page_slice(0, prefix_pages)
_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,
@@ -2777,13 +3107,33 @@ def materialize_shared_token_kv_buffer(
dense_kv_cache, dense_locs = tai_result
use_slot_materialize = False
materialized_by_ipc = False
if use_slot_materialize:
dense_kv_cache = materialize_local_token_kv_page_slots(
dense_kv_cache = kv_cache.new_zeros(
(
(int(materialized_logical_pages.numel()) + 1) * 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=materialized_logical_pages,
layout=layout,
page_size=page_size,
start_slot=0,
end_slot=int(materialized_logical_pages.numel()),
)
if not materialized_by_ipc:
materialize_local_token_kv_page_slots_into(
kv_cache=kv_cache,
dense_kv_cache=dense_kv_cache,
slot_logical_pages=materialized_logical_pages,
layout=layout,
page_size=page_size,
start_slot=0,
end_slot=int(materialized_logical_pages.numel()),
)
elif dense_kv_cache is None:
dense_kv_cache = materialize_local_token_kv_pages(
kv_cache=kv_cache,
@@ -2820,13 +3170,14 @@ def materialize_shared_token_kv_buffer(
),
)
dense_kv_cache = _all_reduce_materialized_buffer(
dense_kv_cache,
layout.cp_size,
nvtx_source=nvtx_source,
nvtx_layer_id=nvtx_layer_id,
nvtx_cp_rank=layout.cp_rank,
)
if not materialized_by_ipc:
dense_kv_cache = _all_reduce_materialized_buffer(
dense_kv_cache,
layout.cp_size,
nvtx_source=nvtx_source,
nvtx_layer_id=nvtx_layer_id,
nvtx_cp_rank=layout.cp_rank,
)
if cp_shared_kv_debug_enabled():
cp_shared_kv_debug_log(
@@ -2860,29 +3211,48 @@ def materialize_shared_paged_buffer(
layout=layout,
physical_page_capacity=page_buffer.shape[0],
)
tai_result = _try_tai_materialize_shared_pages(
page_buffer=page_buffer,
logical_pages=logical_pages,
layout=layout,
)
if tai_result is not None:
dense_page_buffer, dense_pages = tai_result
materialized_logical_pages = logical_pages.reshape(-1)
elif slot_remap is not None:
materialized_by_ipc = False
if slot_remap is not None:
materialized_logical_pages = slot_remap.slot_logical_pages
dense_pages = slot_remap.dense_pages
dense_page_buffer = materialize_local_paged_buffer_page_slots(
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=materialized_logical_pages,
layout=layout,
start_slot=0,
end_slot=int(materialized_logical_pages.numel()),
)
if not materialized_by_ipc:
materialize_local_paged_buffer_page_slots_into(
page_buffer=page_buffer,
dense_page_buffer=dense_page_buffer,
slot_logical_pages=materialized_logical_pages,
layout=layout,
start_slot=0,
end_slot=int(materialized_logical_pages.numel()),
)
else:
materialized_logical_pages, dense_pages = build_slot_page_remap(logical_pages)
dense_page_buffer = materialize_local_paged_buffer_page_slots(
tai_result = _try_tai_materialize_shared_pages(
page_buffer=page_buffer,
slot_logical_pages=materialized_logical_pages,
logical_pages=logical_pages,
layout=layout,
)
if tai_result is not None:
dense_page_buffer, dense_pages = tai_result
materialized_logical_pages = logical_pages.reshape(-1)
else:
materialized_logical_pages, dense_pages = build_slot_page_remap(
logical_pages
)
dense_page_buffer = materialize_local_paged_buffer_page_slots(
page_buffer=page_buffer,
slot_logical_pages=materialized_logical_pages,
layout=layout,
)
if cp_shared_kv_debug_enabled():
owned_pages = materialized_logical_pages[
@@ -2910,13 +3280,14 @@ def materialize_shared_paged_buffer(
),
)
dense_page_buffer = _all_reduce_materialized_buffer(
dense_page_buffer,
layout.cp_size,
nvtx_source=nvtx_source,
nvtx_layer_id=nvtx_layer_id,
nvtx_cp_rank=layout.cp_rank,
)
if not materialized_by_ipc:
dense_page_buffer = _all_reduce_materialized_buffer(
dense_page_buffer,
layout.cp_size,
nvtx_source=nvtx_source,
nvtx_layer_id=nvtx_layer_id,
nvtx_cp_rank=layout.cp_rank,
)
if cp_shared_kv_debug_enabled():
cp_shared_kv_debug_log(