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