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