Ground CP shared KV collective choices in production-shaped evidence

Add a production-style CP shared KV collective benchmark and keep the page-aligned cache contract ledger current with the collective, unpack, multi-extend prefix, P2P, and CUDA IPC findings. The benchmark makes the final consumer layout explicit so dense all-reduce and owner-packed all-gather are compared after producing the same logical-dense result.\n\nConstraint: Prefixes may be assembled from multiple radix/cache extents, so single-run zigzag ordering is only a gated fast path.\nRejected: Treat rank-major all-gather output as the final product | consumers require logical-dense ordering.\nRejected: Switch index materialization based on noisy microbenchmarks | current evidence is not strong enough for production.\nConfidence: medium\nScope-risk: narrow\nDirective: Do not replace generic logical-dense materialization with zigzag ordered gather unless a run descriptor proves the prefix shape.\nTested: python -m py_compile benchmark/hicache/bench_cp_shared_kv_production_collective.py\nTested: Remote production-style benchmark logs under /mnt/beegfs/cjy/log/cp_shared_kv_collective_bench_*_20260531_*.log\nNot-tested: Production SGLang ETE with a CUDA IPC/P2P materialization path; this commit only adds benchmark/documentation surfaces.
This commit is contained in:
laoyao0822
2026-05-31 17:04:00 +08:00
parent 71872bb851
commit a149289554
2 changed files with 949 additions and 0 deletions
@@ -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()
@@ -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.