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:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user