Enable compute-owner KV layout by page-aligning NSA CP split

Phase 5 needs each current KV page to have exactly one CP compute owner before local KV/index direct writes can be safe. This change teaches in-seq NSA prefill CP to produce page-aligned split metadata under shared-KV mode, threads page size into the metadata builders, and fixes local pair splitting so unequal page-aligned zigzag segments do not corrupt topk inputs.

Constraint: Phase 5 direct-write layout requires page ownership to be expressible at page granularity
Constraint: Short page-unit batches remain on the token-balanced fallback to avoid zero-page segment risk
Rejected: Split local q/weights by half | page-aligned zigzag segments can have unequal token counts
Confidence: medium
Scope-risk: moderate
Directive: Do not enable compute-owner direct writes unless nsa_cp_metadata.page_aligned is true and local loc ownership is verified
Tested: python3 -m py_compile python/sglang/srt/layers/attention/nsa/utils.py python/sglang/srt/layers/attention/nsa/nsa_indexer.py python/sglang/srt/models/deepseek_v2.py python/sglang/srt/models/deepseek_nextn.py test/registered/unit/layers/test_nsa_cp_utils.py
Not-tested: Local pytest collection is blocked in this environment by missing pybase64; container/runtime tests were not rerun during this commit step
This commit is contained in:
laoyao0822
2026-05-01 00:54:16 +08:00
parent 47bd2fdf1f
commit 91fa31bcac
6 changed files with 344 additions and 16 deletions

View File

@@ -2,6 +2,9 @@ import unittest
from sglang.srt.layers.attention.nsa.utils import (
_get_in_seq_last_token_owner_and_offset,
build_page_aligned_in_seq_split_list,
build_token_balanced_in_seq_split_list,
split_in_seq_cp_local_pair,
)
from sglang.test.ci.ci_register import register_cpu_ci
@@ -9,6 +12,103 @@ register_cpu_ci(est_time=1, suite="stage-a-test-cpu")
class TestNSAInSeqCPUtils(unittest.TestCase):
def assert_page_aligned_boundaries(
self, split_list, *, extend_prefix_len, extend_len, page_size
):
cursor = 0
for segment_len in split_list[:-1]:
cursor += segment_len
if cursor < extend_len:
self.assertEqual((extend_prefix_len + cursor) % page_size, 0)
def test_page_aligned_split_keeps_boundaries_on_pages(self):
split_list, info = build_page_aligned_in_seq_split_list(
total_len=32768,
extend_len=32768,
extend_prefix_len=0,
page_size=64,
cp_size=8,
)
self.assertTrue(info.page_aligned)
self.assertEqual(sum(split_list), 32768)
self.assertEqual(len(split_list), 16)
self.assertTrue(all(segment_len > 0 for segment_len in split_list))
self.assert_page_aligned_boundaries(
split_list, extend_prefix_len=0, extend_len=32768, page_size=64
)
def test_page_aligned_split_uses_prefix_for_boundary_alignment(self):
split_list, info = build_page_aligned_in_seq_split_list(
total_len=1024,
extend_len=1024,
extend_prefix_len=128,
page_size=64,
cp_size=8,
)
self.assertTrue(info.page_aligned)
self.assertEqual(sum(split_list), 1024)
self.assert_page_aligned_boundaries(
split_list, extend_prefix_len=128, extend_len=1024, page_size=64
)
def test_page_aligned_split_keeps_tail_partial_page_unsplit(self):
split_list, info = build_page_aligned_in_seq_split_list(
total_len=1100,
extend_len=1100,
extend_prefix_len=0,
page_size=64,
cp_size=8,
)
self.assertTrue(info.page_aligned)
self.assertEqual(sum(split_list), 1100)
self.assertEqual(split_list[-1], 12)
self.assert_page_aligned_boundaries(
split_list, extend_prefix_len=0, extend_len=1100, page_size=64
)
def test_page_aligned_split_falls_back_when_prefix_is_not_page_aligned(self):
split_list, info = build_page_aligned_in_seq_split_list(
total_len=1024,
extend_len=1024,
extend_prefix_len=1,
page_size=64,
cp_size=8,
)
self.assertFalse(info.page_aligned)
self.assertEqual(split_list, build_token_balanced_in_seq_split_list(1024, 8))
def test_page_aligned_split_falls_back_when_page_units_are_too_short(self):
split_list, info = build_page_aligned_in_seq_split_list(
total_len=512,
extend_len=512,
extend_prefix_len=0,
page_size=64,
cp_size=8,
)
self.assertFalse(info.page_aligned)
self.assertEqual(split_list, build_token_balanced_in_seq_split_list(512, 8))
def test_page_aligned_split_adds_padding_tokens_to_last_segment(self):
split_list, info = build_page_aligned_in_seq_split_list(
total_len=1040,
extend_len=1024,
extend_prefix_len=0,
page_size=64,
cp_size=8,
)
self.assertTrue(info.page_aligned)
self.assertEqual(sum(split_list), 1040)
self.assertEqual(split_list[-1], 80)
self.assert_page_aligned_boundaries(
split_list, extend_prefix_len=0, extend_len=1024, page_size=64
)
def test_last_token_owner_uses_actual_token_count_when_batch_is_padded(self):
# Padded prefill can have 64 model tokens while the real prompt has only
# 11 tokens. In in-seq split with cp_size=8, the real last token is in
@@ -36,6 +136,22 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
self.assertEqual(owner, 0)
self.assertEqual(local_offset, 7)
def test_local_pair_split_uses_metadata_lengths_not_half_split(self):
import torch
tensor = torch.arange(9)
prev, next_ = split_in_seq_cp_local_pair(tensor, 6, 3)
self.assertEqual(prev.tolist(), [0, 1, 2, 3, 4, 5])
self.assertEqual(next_.tolist(), [6, 7, 8])
def test_local_pair_split_rejects_stale_metadata(self):
import torch
with self.assertRaisesRegex(RuntimeError, "local in-seq CP length mismatch"):
split_in_seq_cp_local_pair(torch.arange(9), 5, 5, name="q_fp8")
if __name__ == "__main__":
unittest.main()