diff --git a/benchmark/hicache/bench_cp_shared_kv_production_collective.py b/benchmark/hicache/bench_cp_shared_kv_production_collective.py new file mode 100644 index 000000000..52f6bf269 --- /dev/null +++ b/benchmark/hicache/bench_cp_shared_kv_production_collective.py @@ -0,0 +1,673 @@ +"""Production-style CP shared-KV collective benchmark. + +This benchmark keeps the useful part of +``bench_cp_shared_kv_prefetch_collective.py`` but replaces the synthetic torch +materialize/pack work with the TAI kernels used by the production CP shared-KV +path. + +Measured paths: + + tai_dense_all_reduce: + TAI materialize owner pages into the current dense slot layout, then NCCL + all-reduce the dense buffer. The legacy dense kernel still carries page-0 + dummy internally; the benchmark validates/uses ``dense[1:]`` as the + no-dummy logical final product. + + tai_rank_major_metadata: + TAI build of logical-page -> rank-major communication metadata. This is + intended to be per-request/cache-entry work, not per-layer work. + + tai_pack_all_gather_rank_major_intermediate: + TAI pack only this rank's owner pages, then NCCL all-gather rank-major owner + pages. This buffer is a communication intermediate, not the consumer-facing + result. + + tai_pack_all_gather_reorder_logical: + Full proposed path: pack owner pages, all-gather rank-major pages, then TAI + unpack/reorder back to a no-dummy logical-order dense buffer. This final + product must match ``tai_dense_all_reduce`` after dropping the legacy dummy + page. + + *_remap variants: + Include the TAI logical-loc remap needed by MLA consumers against the final + no-dummy logical dense layout. Invalid/dummy logical locs map to -1. + +Run on the target CUDA machine, for example: + + PYTHONPATH=/mnt/beegfs/cjy/tai-kernel/python:$PYTHONPATH \\ + torchrun --standalone --nproc_per_node=8 \\ + benchmark/hicache/bench_cp_shared_kv_production_collective.py \\ + --payload both --tokens 4096 8192 16384 32768 40320 65536 + +This is still a microbenchmark. It does not measure full attention, prefetch +wait placement, or ETE scheduler interaction. +""" + +from __future__ import annotations + +import argparse +import json +import math +import os +import statistics +import time +from dataclasses import asdict, dataclass +from typing import Callable + +import torch +import torch.distributed as dist + +from tai_kernel.nsa_prefill.cp_shared_kv_materialize import ( + build_logical_dense_page_inverse, + build_rank_major_page_inverse, + materialize_shared_pages, + materialize_shared_token_kv_pages, + pack_shared_pages_rank_major, + pack_shared_token_kv_pages_rank_major, + remap_logical_locs_to_slot_dense_locs, + unpack_rank_major_pages_to_logical_dense, +) + + +@dataclass(frozen=True) +class CaseConfig: + payload: str + prefix_tokens: int + page_size: int + kv_dim: int + index_page_bytes: int + dtype: str + logical_pattern: str + remap_q_tokens: int + remap_topk: int + warmup: int + repeat: int + check: bool + + +@dataclass(frozen=True) +class BenchResult: + payload: str + prefix_tokens: int + prefix_pages: int + path: str + full_mib: float + packed_mib_per_rank: float + gpu_ms_avg: float + gpu_ms_p50: float + gpu_ms_min: float + gpu_ms_max_rank_avg: float + cpu_ms_avg: float + cpu_ms_p50: float + cpu_ms_max_rank_avg: float + + +def _rank_env() -> tuple[int, int, int]: + rank = int(os.environ.get("RANK", "0")) + world_size = int(os.environ.get("WORLD_SIZE", "1")) + local_rank = int(os.environ.get("LOCAL_RANK", str(rank))) + return rank, world_size, local_rank + + +def _dtype(name: str) -> torch.dtype: + if name == "bf16": + return torch.bfloat16 + if name == "fp16": + return torch.float16 + if name == "fp32": + return torch.float32 + raise ValueError(f"unsupported dtype: {name}") + + +def _payload_shape_and_dtype(cfg: CaseConfig) -> tuple[tuple[int, ...], torch.dtype]: + if cfg.payload == "mla": + return (cfg.page_size, cfg.kv_dim), _dtype(cfg.dtype) + if cfg.payload == "index": + return (cfg.index_page_bytes,), torch.uint8 + raise ValueError(f"unsupported payload: {cfg.payload}") + + +def _zigzag_owners(num_pages: int, world_size: int) -> list[int]: + segment_num = world_size * 2 + base = num_pages // segment_num + rem = num_pages % segment_num + owners: list[int] = [] + for segment_idx in range(segment_num): + count = base + (1 if segment_idx < rem else 0) + if segment_idx < world_size: + owner = segment_idx + else: + owner = segment_num - segment_idx - 1 + owners.extend([owner] * count) + return owners + + +def _make_logical_pages( + *, + prefix_pages: int, + world_size: int, + pattern: str, + device: torch.device, +) -> tuple[torch.Tensor, list[int]]: + if pattern == "sequential": + pages = list(range(1, prefix_pages + 1)) + elif pattern == "zigzag-owners": + owner_ordinals = [0 for _ in range(world_size)] + pages = [] + for owner in _zigzag_owners(prefix_pages, world_size): + pages.append(owner + world_size * owner_ordinals[owner] + 1) + owner_ordinals[owner] += 1 + else: + raise ValueError(f"unsupported logical page pattern: {pattern}") + + owner_counts = [0 for _ in range(world_size)] + for page in pages: + if page > 0: + owner_counts[(page - 1) % world_size] += 1 + return torch.tensor(pages, dtype=torch.long, device=device).contiguous(), owner_counts + + +def _make_source_pages( + *, + num_physical_pages: int, + page_shape: tuple[int, ...], + dtype: torch.dtype, + device: torch.device, + rank: int, +) -> torch.Tensor: + shape = (num_physical_pages, *page_shape) + if dtype is torch.uint8: + src = torch.empty(shape, dtype=dtype, device=device) + src.random_(0, 127) + else: + src = torch.empty(shape, dtype=dtype, device=device) + src.normal_(mean=float(rank + 1), std=0.01) + src[0].zero_() + return src.contiguous() + + +def _make_logical_locs( + *, + prefix_pages: int, + page_size: int, + q_tokens: int, + topk: int, + device: torch.device, +) -> torch.Tensor: + if q_tokens <= 0 or topk <= 0 or prefix_pages <= 0: + return torch.empty( + (max(q_tokens, 0), max(topk, 0)), + dtype=torch.long, + device=device, + ) + + num_locs = q_tokens * topk + ordinals = torch.arange(num_locs, dtype=torch.long, device=device) + pages = torch.remainder(ordinals, prefix_pages) + 1 + offsets = torch.remainder(ordinals * 17, page_size) + locs = (pages * page_size + offsets).reshape(q_tokens, topk).contiguous() + if num_locs >= 1024: + locs.reshape(-1)[::1024] = -1 + return locs + + +def _time_path( + fn: Callable[[], None], + *, + warmup: int, + repeat: int, + group: dist.ProcessGroup, + device: torch.device, +) -> tuple[list[float], list[float]]: + for _ in range(warmup): + fn() + torch.cuda.synchronize(device) + dist.barrier(group=group) + + gpu_ms: list[float] = [] + cpu_ms: list[float] = [] + for _ in range(repeat): + dist.barrier(group=group) + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + torch.cuda.synchronize(device) + t0 = time.perf_counter() + start.record() + fn() + end.record() + end.synchronize() + t1 = time.perf_counter() + gpu_ms.append(float(start.elapsed_time(end))) + cpu_ms.append((t1 - t0) * 1000.0) + dist.barrier(group=group) + return gpu_ms, cpu_ms + + +def _summarize_across_ranks( + *, + local_gpu_ms: list[float], + local_cpu_ms: list[float], + group: dist.ProcessGroup, + device: torch.device, +) -> tuple[dict[str, float], dict[str, float]]: + local = torch.tensor( + [ + statistics.mean(local_gpu_ms), + statistics.median(local_gpu_ms), + min(local_gpu_ms), + statistics.mean(local_cpu_ms), + statistics.median(local_cpu_ms), + ], + dtype=torch.float64, + device=device, + ) + world_size = dist.get_world_size(group=group) + gathered = torch.empty((world_size, local.numel()), dtype=local.dtype, device=device) + dist.all_gather_into_tensor(gathered, local, group=group) + gathered_cpu = gathered.cpu() + gpu = { + "avg": float(gathered_cpu[:, 0].mean().item()), + "p50": float(gathered_cpu[:, 1].mean().item()), + "min": float(gathered_cpu[:, 2].mean().item()), + "max_rank_avg": float(gathered_cpu[:, 0].max().item()), + } + cpu = { + "avg": float(gathered_cpu[:, 3].mean().item()), + "p50": float(gathered_cpu[:, 4].mean().item()), + "max_rank_avg": float(gathered_cpu[:, 3].max().item()), + } + return gpu, cpu + + +def _check_rank_major_intermediate_matches_dense( + *, + dense_buffer: torch.Tensor, + rank_major_buffer: torch.Tensor, + slot_dense_pages: torch.Tensor, + rank_major_dense_pages: torch.Tensor, +) -> None: + flat_slot = slot_dense_pages.reshape(-1) + flat_rank_major = rank_major_dense_pages.reshape(-1) + valid = flat_slot > 0 + if not torch.any(valid): + return + dense_selected = dense_buffer[flat_slot[valid].to(torch.long)] + rank_major_selected = rank_major_buffer[flat_rank_major[valid].to(torch.long)] + if not torch.equal(dense_selected, rank_major_selected): + mismatch = int((dense_selected != rank_major_selected).sum().item()) + raise AssertionError(f"rank-major gathered buffer mismatch_count={mismatch}") + + +def _check_logical_dense_matches_dense( + *, + dense_buffer: torch.Tensor, + logical_dense_buffer: torch.Tensor, +) -> None: + expected = dense_buffer[1:] + if not torch.equal(expected, logical_dense_buffer): + mismatch = int((expected != logical_dense_buffer).sum().item()) + raise AssertionError(f"logical dense final buffer mismatch_count={mismatch}") + + +def _check_logical_dense_locs_are_self_consistent( + *, + logical_dense_tokens: torch.Tensor, + logical_locs: torch.Tensor, + logical_dense_locs: torch.Tensor, + logical_pages: torch.Tensor, + page_size: int, +) -> None: + dense_flat = logical_dense_locs.reshape(-1) + logical_flat = logical_locs.reshape(-1) + valid = dense_flat >= 0 + if not torch.any(valid): + return + dense_selected = logical_dense_tokens[dense_flat[valid].to(torch.long)] + + logical_pages_flat = logical_pages.reshape(-1) + loc_pages = torch.div(logical_flat[valid], page_size, rounding_mode="floor") + loc_offsets = torch.remainder(logical_flat[valid], page_size) + expected_slot = torch.empty_like(loc_pages) + for i, page in enumerate(loc_pages.detach().cpu().tolist()): + matches = torch.nonzero(logical_pages_flat == page, as_tuple=False).reshape(-1) + if matches.numel() != 1: + raise AssertionError(f"logical loc page {page} does not map to one slot") + expected_slot[i] = matches[0] + expected_locs = expected_slot * page_size + loc_offsets + expected_selected = logical_dense_tokens[expected_locs.to(torch.long)] + if not torch.equal(dense_selected, expected_selected): + mismatch = int((dense_selected != expected_selected).sum().item()) + raise AssertionError(f"logical dense remapped loc mismatch_count={mismatch}") + + +def _bench_case( + cfg: CaseConfig, + *, + group: dist.ProcessGroup, + device: torch.device, + rank: int, + world_size: int, +) -> list[BenchResult]: + prefix_pages = math.ceil(cfg.prefix_tokens / cfg.page_size) + logical_pages, owner_counts = _make_logical_pages( + prefix_pages=prefix_pages, + world_size=world_size, + pattern=cfg.logical_pattern, + device=device, + ) + max_owned_pages = max(owner_counts) if owner_counts else 0 + logical_page_capacity = max(logical_pages.detach().cpu().tolist(), default=0) + 1 + page_shape, dtype = _payload_shape_and_dtype(cfg) + num_physical_pages = max(owner_counts) + 1 if owner_counts else 1 + src_pages = _make_source_pages( + num_physical_pages=num_physical_pages, + page_shape=page_shape, + dtype=dtype, + device=device, + rank=rank, + ) + + slot_dense_pages = torch.where( + logical_pages > 0, + torch.arange(1, prefix_pages + 1, dtype=logical_pages.dtype, device=device), + logical_pages, + ) + rank_major_dense_pages, _ = build_rank_major_page_inverse( + logical_pages, + logical_page_capacity, + cp_size=world_size, + max_owned_pages=max_owned_pages, + ) + logical_dense_page_inverse = build_logical_dense_page_inverse( + logical_pages, + logical_page_capacity, + ) + logical_locs = _make_logical_locs( + prefix_pages=prefix_pages, + page_size=cfg.page_size, + q_tokens=cfg.remap_q_tokens, + topk=cfg.remap_topk, + device=device, + ) + + local_packed = torch.empty( + (max_owned_pages, *page_shape), + dtype=dtype, + device=device, + ) + gathered = torch.empty( + (world_size, max_owned_pages, *page_shape), + dtype=dtype, + device=device, + ) + full_mib = (prefix_pages + 1) * src_pages[0].numel() * src_pages.element_size() + full_mib /= 1024 * 1024 + packed_mib = local_packed.numel() * local_packed.element_size() / (1024 * 1024) + + def dense_all_reduce() -> torch.Tensor: + if cfg.payload == "mla": + dense = materialize_shared_token_kv_pages( + src_pages.reshape(src_pages.shape[0] * cfg.page_size, cfg.kv_dim), + logical_pages, + page_size=cfg.page_size, + cp_rank=rank, + cp_size=world_size, + ).reshape(prefix_pages + 1, *page_shape) + else: + dense, _ = materialize_shared_pages( + src_pages, + logical_pages, + cp_rank=rank, + cp_size=world_size, + ) + dist.all_reduce(dense, op=dist.ReduceOp.SUM, group=group) + return dense + + def build_metadata() -> None: + build_rank_major_page_inverse( + logical_pages, + logical_page_capacity, + cp_size=world_size, + max_owned_pages=max_owned_pages, + ) + + def logical_dense_loc_remap_cached_meta() -> torch.Tensor: + return remap_logical_locs_to_slot_dense_locs( + logical_locs, + logical_dense_page_inverse, + page_size=cfg.page_size, + ) + + def logical_dense_metadata_loc_remap() -> torch.Tensor: + page_inverse = build_logical_dense_page_inverse( + logical_pages, + logical_page_capacity, + ) + return remap_logical_locs_to_slot_dense_locs( + logical_locs, + page_inverse, + page_size=cfg.page_size, + ) + + def dense_all_reduce_with_remap() -> None: + dense_all_reduce() + if cfg.payload == "mla": + logical_dense_loc_remap_cached_meta() + + def pack_all_gather() -> None: + nonlocal local_packed + if cfg.payload == "mla": + local_packed = pack_shared_token_kv_pages_rank_major( + src_pages.reshape(src_pages.shape[0] * cfg.page_size, cfg.kv_dim), + logical_pages, + rank_major_dense_pages, + page_size=cfg.page_size, + max_owned_pages=max_owned_pages, + cp_rank=rank, + cp_size=world_size, + ) + else: + local_packed = pack_shared_pages_rank_major( + src_pages, + logical_pages, + rank_major_dense_pages, + max_owned_pages=max_owned_pages, + cp_rank=rank, + cp_size=world_size, + ) + dist.all_gather_into_tensor(gathered, local_packed, group=group) + + def unpack_rank_major_to_logical_dense() -> torch.Tensor: + return unpack_rank_major_pages_to_logical_dense( + gathered.reshape(world_size * max_owned_pages, *page_shape), + rank_major_dense_pages, + ) + + def pack_all_gather_reorder_logical() -> torch.Tensor: + pack_all_gather() + return unpack_rank_major_to_logical_dense() + + def pack_all_gather_reorder_logical_with_remap() -> None: + pack_all_gather_reorder_logical() + if cfg.payload == "mla": + logical_dense_loc_remap_cached_meta() + + if cfg.check: + dense = dense_all_reduce() + pack_all_gather() + rank_major_buffer = gathered.reshape(world_size * max_owned_pages, *page_shape) + _check_rank_major_intermediate_matches_dense( + dense_buffer=dense, + rank_major_buffer=rank_major_buffer, + slot_dense_pages=slot_dense_pages, + rank_major_dense_pages=rank_major_dense_pages, + ) + logical_dense = unpack_rank_major_to_logical_dense() + _check_logical_dense_matches_dense( + dense_buffer=dense, + logical_dense_buffer=logical_dense, + ) + if cfg.payload == "mla": + logical_dense_locs = logical_dense_loc_remap_cached_meta() + _check_logical_dense_locs_are_self_consistent( + logical_dense_tokens=logical_dense.reshape( + prefix_pages * cfg.page_size, + cfg.kv_dim, + ), + logical_locs=logical_locs, + logical_dense_locs=logical_dense_locs, + logical_pages=logical_pages, + page_size=cfg.page_size, + ) + dist.barrier(group=group) + + paths: list[tuple[str, Callable[[], None]]] = [ + ("tai_dense_all_reduce", lambda: dense_all_reduce()), + ("tai_rank_major_metadata", build_metadata), + ("tai_pack_all_gather_rank_major_intermediate", pack_all_gather), + ( + "tai_unpack_rank_major_to_logical_dense", + lambda: unpack_rank_major_to_logical_dense(), + ), + ( + "tai_pack_all_gather_reorder_logical", + lambda: pack_all_gather_reorder_logical(), + ), + ] + if cfg.payload == "mla": + paths.extend( + [ + ( + "tai_logical_dense_loc_remap_cached_meta", + logical_dense_loc_remap_cached_meta, + ), + ( + "tai_logical_dense_metadata_loc_remap", + logical_dense_metadata_loc_remap, + ), + ("tai_dense_all_reduce_remap", dense_all_reduce_with_remap), + ( + "tai_pack_all_gather_reorder_logical_remap", + pack_all_gather_reorder_logical_with_remap, + ), + ] + ) + results: list[BenchResult] = [] + for name, fn in paths: + gpu_ms, cpu_ms = _time_path( + fn, + warmup=cfg.warmup, + repeat=cfg.repeat, + group=group, + device=device, + ) + gpu, cpu = _summarize_across_ranks( + local_gpu_ms=gpu_ms, + local_cpu_ms=cpu_ms, + group=group, + device=device, + ) + results.append( + BenchResult( + payload=cfg.payload, + prefix_tokens=cfg.prefix_tokens, + prefix_pages=prefix_pages, + path=name, + full_mib=full_mib, + packed_mib_per_rank=packed_mib, + gpu_ms_avg=gpu["avg"], + gpu_ms_p50=gpu["p50"], + gpu_ms_min=gpu["min"], + gpu_ms_max_rank_avg=gpu["max_rank_avg"], + cpu_ms_avg=cpu["avg"], + cpu_ms_p50=cpu["p50"], + cpu_ms_max_rank_avg=cpu["max_rank_avg"], + ) + ) + return results + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--payload", choices=["mla", "index", "both"], default="both") + parser.add_argument( + "--tokens", + type=int, + nargs="+", + default=[4096, 8192, 16384, 32768, 40320, 65536], + ) + parser.add_argument("--page-size", type=int, default=64) + parser.add_argument("--kv-dim", type=int, default=576) + parser.add_argument("--index-page-bytes", type=int, default=8448) + parser.add_argument("--dtype", choices=["bf16", "fp16", "fp32"], default="bf16") + parser.add_argument( + "--logical-pattern", + choices=["sequential", "zigzag-owners"], + default="sequential", + ) + parser.add_argument("--remap-q-tokens", type=int, default=1024) + parser.add_argument("--remap-topk", type=int, default=16) + parser.add_argument("--warmup", type=int, default=5) + parser.add_argument("--repeat", type=int, default=20) + parser.add_argument("--no-check", action="store_true") + return parser.parse_args() + + +def main() -> None: + args = _parse_args() + rank, world_size, local_rank = _rank_env() + torch.cuda.set_device(local_rank) + device = torch.device("cuda", local_rank) + dist.init_process_group("nccl", init_method="env://") + group = dist.group.WORLD + + payloads = ["mla", "index"] if args.payload == "both" else [args.payload] + if rank == 0: + print( + "CP shared-KV production collective benchmark " + f"world_size={world_size} tokens={args.tokens} payload={args.payload} " + f"logical_pattern={args.logical_pattern} " + f"remap_q_tokens={args.remap_q_tokens} remap_topk={args.remap_topk}" + ) + + for payload in payloads: + for tokens in args.tokens: + cfg = CaseConfig( + payload=payload, + prefix_tokens=tokens, + page_size=args.page_size, + kv_dim=args.kv_dim, + index_page_bytes=args.index_page_bytes, + dtype=args.dtype, + logical_pattern=args.logical_pattern, + remap_q_tokens=args.remap_q_tokens, + remap_topk=args.remap_topk, + warmup=args.warmup, + repeat=args.repeat, + check=not args.no_check, + ) + for result in _bench_case( + cfg, + group=group, + device=device, + rank=rank, + world_size=world_size, + ): + if rank == 0: + print(json.dumps(asdict(result), sort_keys=True)) + print( + f"{result.payload:5s} tokens={result.prefix_tokens:6d} " + f"pages={result.prefix_pages:5d} {result.path:40s} " + f"gpu_avg={result.gpu_ms_avg:8.3f}ms " + f"gpu_max_rank={result.gpu_ms_max_rank_avg:8.3f}ms " + f"cpu_avg={result.cpu_ms_avg:8.3f}ms " + f"full={result.full_mib:8.1f}MiB " + f"packed/rank={result.packed_mib_per_rank:8.1f}MiB" + ) + + dist.barrier(group=group) + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/docs/advanced_features/nsa_prefill_cp_page_aligned_cache_contract.md b/docs/advanced_features/nsa_prefill_cp_page_aligned_cache_contract.md index d15d287c6..4c96191e6 100644 --- a/docs/advanced_features/nsa_prefill_cp_page_aligned_cache_contract.md +++ b/docs/advanced_features/nsa_prefill_cp_page_aligned_cache_contract.md @@ -3613,3 +3613,279 @@ Verification: `PYTHONPATH=python python -m pytest -q test/registered/unit/mem_cache/test_cp_hicache_metadata.py::TestHiCacheEvictLoggingLevels::test_evict_hot_path_success_logs_are_debug_only` → `1 passed`. - Local and remote `py_compile` passed for touched Python files. + +### C87 — 2026-05-31 all-gather owner-packed path must unpack to logical dense final product + +Decision: + +- Rank-major/owner-major packed pages are only a communication intermediate for + the proposed all-gather replacement of dense all-reduce. +- The consumer-facing final product remains a dense buffer in request + logical/page-table order. Dense slot `i` corresponds to + `logical_pages.reshape(-1)[i]`. +- The no-dummy final contract maps invalid/sentinel pages and invalid locs to + `-1`; page 0 is not used as a silent dummy destination in the new final + remap contract. + +Implication for benchmarking: + +- Measuring only `pack -> all_gather -> rank-major buffer` is incomplete and + overestimates the all-gather approach. +- The production-style benchmark must include the reorder/unpack cost: + `pack owner pages -> all_gather rank-major pages -> unpack to logical dense`. +- MLA remap must also be measured against the final no-dummy logical dense + layout, not against the rank-major intermediate. + +Implemented TAI support: + +- `build_logical_dense_page_inverse()` builds a no-dummy zero-based + logical-page-to-dense-slot inverse. +- `unpack_rank_major_pages_to_logical_dense()` copies the rank-major gathered + intermediate back to logical dense order and zero-fills invalid/sentinel slots. + +Verification: + +- RED before implementation: the two new TAI CUDA tests failed with missing + imports for `build_logical_dense_page_inverse` and + `unpack_rank_major_pages_to_logical_dense`. +- GREEN after implementation on g0034 container: + `test_build_logical_dense_page_inverse_remaps_without_dummy_page` and + `test_unpack_rank_major_pages_to_logical_dense_matches_logical_order_reference` + passed. + +### C88 — 2026-05-31 owner-packed all-gather benchmark after logical-dense unpack + +Remote benchmark environment: + +- Host/container: g0034 / `sglang-glm5-dev-2`. +- Pattern: `--logical-pattern zigzag-owners`, 8 CP ranks, correctness checks enabled. +- Logs: + - `/mnt/beegfs/cjy/log/cp_shared_kv_collective_bench_20260531_154608.log` + - `/mnt/beegfs/cjy/log/cp_shared_kv_collective_bench_mla_focus_20260531_154720.log` + - `/mnt/beegfs/cjy/log/cp_shared_kv_collective_bench_index_focus_20260531_154804.log` + +MLA focused run (`repeat=30`, p50 GPU ms): + +| tokens | dense all-reduce | pack+all-gather+unpack | dense+remap | gather+unpack+remap | +| ---: | ---: | ---: | ---: | ---: | +| 4096 | 0.148 | 0.153 | 0.150 | 0.184 | +| 8192 | 0.167 | 0.151 | 0.169 | 0.182 | +| 16384 | 0.211 | 0.171 | 0.216 | 0.182 | +| 40320 | 0.330 | 0.275 | 0.334 | 0.284 | + +MLA breakdown, avg GPU ms: + +- rank-major metadata: ~0.059–0.060 ms. +- rank-major unpack to logical dense: ~0.056 ms up to 16k, ~0.070 ms at 40k. +- cached logical-dense loc remap: ~0.041 ms. +- pack+all-gather intermediate grows with bytes: ~0.132 ms at 4k, ~0.249 ms at 40k. + +Conclusion for MLA: + +- After including the required logical-dense unpack, owner-packed all-gather is + not a clear win for very short prefixes. At 4k it is roughly equal/slower, + especially when remap is included. +- It becomes beneficial around 16k+ in this microbenchmark and remains beneficial + at 40k. +- The fixed overheads (metadata, unpack, remap) are large enough that this path + should be gated by payload/size and should remain per-request, not per-layer. + +Index focused run (`repeat=30`, p50 GPU ms): + +| tokens | dense all-reduce | pack+all-gather+unpack | +| ---: | ---: | ---: | +| 4096 | 0.111 | 0.144 | +| 8192 | 0.110 | 0.142 | +| 16384 | 0.119 | 0.142 | +| 40320 | 0.125 | 0.140 | + +Conclusion for index: + +- The index payload is too small for owner-packed all-gather plus unpack to win + in this benchmark. Dense all-reduce p50 remains faster. +- Do not switch index prefetch/materialize collective to owner-packed all-gather + based on the current evidence. + +### C89 — 2026-05-31 extended prefix sweep 4k–120k + +Remote benchmark logs: + +- MLA: `/mnt/beegfs/cjy/log/cp_shared_kv_collective_bench_mla_4k_120k_20260531_160744.log` +- Index: `/mnt/beegfs/cjy/log/cp_shared_kv_collective_bench_index_4k_120k_20260531_160839.log` + +Configuration: + +- 8 CP ranks, `--logical-pattern zigzag-owners`, `--warmup 8`, `--repeat 30`. +- Prefix tokens: `4096 8192 16384 32768 65536 98304 122880`. +- Final product for owner-packed path includes the required + rank-major-to-logical-dense unpack. + +MLA p50 GPU ms: + +| tokens | pages | dense all-reduce | gather+unpack | ratio | dense+remap | gather+unpack+remap | ratio | +| ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | +| 4096 | 64 | 0.142 | 0.154 | 1.09 | 0.147 | 0.183 | 1.25 | +| 8192 | 128 | 0.166 | 0.151 | 0.91 | 0.170 | 0.182 | 1.07 | +| 16384 | 256 | 0.211 | 0.171 | 0.81 | 0.215 | 0.184 | 0.86 | +| 32768 | 512 | 0.288 | 0.243 | 0.84 | 0.290 | 0.250 | 0.86 | +| 65536 | 1024 | 0.446 | 0.378 | 0.85 | 0.450 | 0.385 | 0.85 | +| 98304 | 1536 | 0.607 | 0.519 | 0.86 | 0.611 | 0.525 | 0.86 | +| 122880 | 1920 | 0.726 | 0.627 | 0.86 | 0.729 | 0.633 | 0.87 | + +MLA p50 breakdown: + +- Metadata is flat at ~0.058 ms. +- Cached remap is flat at ~0.040 ms. +- Pack+all-gather grows with transferred bytes: 0.124 ms at 4k, 0.542 ms at 120k. +- Unpack grows with final dense bytes: 0.055 ms at 4k, 0.125 ms at 120k. + +MLA conclusion: + +- With final logical-dense unpack included, owner-packed all-gather is still not + worth it at 4k. +- It starts becoming useful around 8k for transfer-only, and around 16k once + remap is included. +- From 16k to 120k it is consistently about 13–19% faster by p50 in this + microbenchmark. + +Index p50 GPU ms from this run: + +| tokens | dense all-reduce | gather+unpack | ratio | +| ---: | ---: | ---: | ---: | +| 4096 | 0.196 | 0.230 | 1.18 | +| 8192 | 0.197 | 0.233 | 1.18 | +| 16384 | 0.204 | 0.228 | 1.12 | +| 32768 | 0.210 | 0.136 | 0.65 | +| 65536 | 0.138 | 0.138 | 1.00 | +| 98304 | 0.154 | 0.140 | 0.91 | +| 122880 | 0.176 | 0.156 | 0.88 | + +Index caveat: + +- Index timings are small and showed more run-to-run inconsistency than MLA. +- Keep the earlier conclusion conservative: do not switch index to + owner-packed all-gather based only on this microbenchmark until the production + critical path and warmup/order effects are isolated. + +### C90 — 2026-05-31 multi-extend prefixes break single-zigzag ordered all-gather assumption + +Finding: + +- The two-pass ordered all-gather idea only directly applies when the requested + prefix is one contiguous in-seq CP split extent whose logical pages follow a + single zigzag owner run: + `0,1,...,cp_size-1,cp_size-1,...,1,0`. +- Production prefixes can be assembled from multiple historical extend/cache + extents through radix/HiCache reuse. The global logical prefix can therefore + look like a concatenation of several independently created zigzag runs, not one + single zigzag over the full prefix length. +- A global two-pass all-gather keyed only by CP rank would mis-order pages for + these multi-extent prefixes unless it also carries extent/run metadata and + performs the ordered gather per run, or falls back to the generic rank-major + gather plus logical-dense unpack. + +Implication: + +- Treat single-zigzag two-pass ordered all-gather as a fast path gated by an + explicit single-run/regular-run descriptor check. +- Keep the generic rank-major + unpack path as the correctness baseline for + multi-extend prefixes until segmented ordered-gather is implemented and + benchmarked. + +### C91 — 2026-05-31 custom P2P allgatherv candidate to replace dense all-reduce + +Finding: + +- Dense all-reduce is the wrong semantic primitive for CP shared KV materialize: + each rank owns a sparse subset of pages, but all-reduce communicates the full + dense buffer and performs a reduction even though non-owner entries are zeros. +- Padded owner-packed all-gather removes the reduction and most zero traffic, but + still requires equal per-rank payload size (`max_owned_pages`) and an unpack to + logical dense order. +- Production prefixes may be composed from multiple extend/cache runs, so a + single ordered zigzag all-gather is only a gated fast path, not the generic + replacement. + +Candidate generic replacement: + +- Implement a custom P2P ring allgatherv for owner page payloads: + 1. All ranks derive per-owner page counts and compact offsets from the same + `logical_pages`/layout metadata; no runtime size collective is needed in the + common deterministic case. + 2. Each rank packs only its owned pages into a compact contiguous segment. + 3. Ring P2P forwards variable-size segments for `cp_size - 1` steps using + `batch_isend_irecv`/NCCL P2P. + 4. Receives write directly into a compact owner-concat buffer at known offsets. + 5. Existing logical-dense unpack can consume a compact page inverse + (`logical_slot -> owner_offset + local_owner_ordinal`) instead of padded + `rank * max_owned_pages + ordinal`. + +Risks / benchmark requirements: + +- Naive direct send-to-all is not acceptable; it increases per-rank injection + versus ring all-gather. The P2P implementation must use a ring or another + bandwidth-efficient schedule. +- P2P has more launch/control overhead than one NCCL collective. It must be + benchmarked against both dense all-reduce and padded all-gather+unpack for + 4k--120k prefixes and both MLA/index payloads. +- NCCL P2P is not guaranteed to be zero-SM; do not claim compute-overlap benefits + without profiling. + +### C92 — 2026-05-31 CUDA IPC P2P ops are viable but should reuse existing SGLang IPC infrastructure + +Finding: + +- The repo already has CUDA IPC support in several places: + - `python/sglang/srt/distributed/device_communicators/cuda_wrapper.py` + exposes `cudaIpcGetMemHandle` / `cudaIpcOpenMemHandle`. + - `python/sglang/srt/distributed/device_communicators/custom_all_reduce.py` + and `custom_all_reduce_utils.py` already register/open IPC buffers and test + real P2P access. + - `python/sglang/jit_kernel/include/sgl_kernel/distributed/custom_all_reduce.cuh` + already caches opened IPC handles and maintains peer pointer tables. +- Therefore a CP shared KV CUDA-IPC transport should not start from scratch. It + should reuse the existing handle exchange, peer-access validation, and handle + lifetime patterns, then add a narrower owner-page copy/gather primitive. + +Design direction: + +- Build an intra-node CUDA IPC owner-page transport for CP ranks on the same host. +- Register long-lived staging/output buffers once per scheduler/rank, exchange IPC + handles during CP group initialization, and keep opened peer pointers cached. +- Launch a custom copy kernel that pulls remote owner page segments from peer + pointers into the local compact owner-concat or logical-dense destination. +- Keep NCCL/PyTorch collective fallback only for unsupported topology or failed + IPC/P2P validation, and log it explicitly. + +Constraints: + +- CUDA IPC is intra-node only. It does not replace RDMA or cross-node transfer. +- Correctness requires explicit stream/event ordering: a rank must not read a + peer staging segment before the peer pack kernel has completed. +- Lifetime must be explicit: registered peer buffers cannot be freed/reallocated + while another rank holds an IPC pointer. +- This path may still use SMs for the copy kernel unless implemented with copy + engine primitives; do not assume it is 0-SM without profiling. + +### C93 — 2026-05-31 tai-kernel is the right boundary for CUDA IPC owner-page gather + +Finding: + +- `tai-kernel` already has a JIT C++/CUDA extension path for NSA prefill KV IO: + `python/tai_kernel/nsa_prefill/_kvcacheio_extension_loader.py` builds + `csrc/kvcacheio_ops.cc` and `csrc/kvcacheio_lf_pf.cu` with + `torch.utils.cpp_extension.load` and registers `tai_kernel_ops` via + `TORCH_LIBRARY_FRAGMENT`. +- CUDA IPC cannot be implemented in Triton alone because handle creation/opening, + peer pointer lifetime, and optional interprocess event handling require CUDA + runtime APIs. The right shape is a small C++/CUDA extension in `tai-kernel` + plus CUDA kernels for peer-pointer copy/gather. + +Design implication: + +- Add IPC primitives to the existing tai-kernel extension boundary rather than + adding another SGLang-local extension or relying on Python/ctypes in the hot + path. +- Expose handles as fixed-size uint8 tensors or keep opened peer pointers in a + C++ registry keyed by an integer transport id. Avoid Python object handles in + production hot paths.