Reduce CP shared KV overhead without changing ownership semantics

The shared-KV path now keeps more CP metadata on-device and reuses
physical out-cache locations across MLA and NSA index writes, so each
layer avoids repeating logical-to-physical remaps. The in-seq CP
all-gather rerange path now delegates to tai-kernel when available and
falls back to the existing torch split/cat path with an explicit log.

This also extends the Phase8 prefetch machinery to cover shared KV
materialization metadata and keeps debug/fallback behavior gated so the
fast path is not polluted by diagnostic checks.

Constraint: Custom CP kernels must live in tai-kernel and be imported lazily from SGLang
Constraint: Decode does not use CP; these changes target NSA prefill CP in-seq-split shared KV
Rejected: Recompute physical local cache locations separately for MLA and index writes | repeats the same remap work every layer
Rejected: Keep the in-seq rerange Triton code inline in SGLang | duplicates kernel ownership and blocks tai-kernel reuse
Confidence: medium
Scope-risk: moderate
Directive: Keep CP collective ordering identical across ranks; do not add rank-local fallback decisions inside shared KV materialize paths
Tested: Remote g0034 container py_compile for modified SGLang/tai-kernel files; remote pytest test/registered/unit/layers/test_nsa_cp_utils.py passed with 24 tests
Not-tested: Full multi-node GLM5 prefill/decode throughput after the final commit boundary
This commit is contained in:
laoyao0822
2026-05-06 05:27:43 +08:00
parent 5e5ac5e2e7
commit 43ad2fe52d
10 changed files with 1152 additions and 46 deletions
@@ -11,6 +11,7 @@ from sglang.srt.layers.attention.nsa.utils import (
can_cp_split,
cp_split_and_rebuild_1d,
get_cp_shared_kv_local_out_cache_loc,
get_cp_shared_kv_local_physical_out_cache_loc,
split_in_seq_cp_local_pair,
)
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
@@ -335,6 +336,45 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
+ list(range(6 * page_size, 7 * page_size)),
)
def test_local_physical_out_cache_loc_is_cached(self):
import torch
from types import SimpleNamespace
page_size = 4
segment_pages = [1, 2, 3, 4, 8, 7, 6, 5]
out_cache_loc = torch.cat(
[
torch.arange(page * page_size, (page + 1) * page_size)
for page in segment_pages
]
)
forward_batch = SimpleNamespace(
uses_cp_shared_kv=True,
cp_shared_kv_layout=CpSharedKVLayout(
page_size=page_size,
cp_size=4,
cp_rank=1,
),
nsa_cp_metadata=NSAContextParallelMetadata(
split_list=[page_size] * 8,
zigzag_index=[1, 6],
page_aligned=True,
page_size=page_size,
extend_prefix_len=0,
),
out_cache_loc=out_cache_loc,
)
physical_locs = get_cp_shared_kv_local_physical_out_cache_loc(forward_batch)
second_read = get_cp_shared_kv_local_physical_out_cache_loc(forward_batch)
self.assertIs(physical_locs, second_read)
self.assertEqual(
physical_locs.tolist(),
list(range(1 * page_size, 2 * page_size))
+ list(range(2 * page_size, 3 * page_size)),
)
def test_local_out_cache_loc_falls_back_when_owner_mismatch(self):
import torch
from types import SimpleNamespace