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