CP shared KV and HiCache need a stable contract where physical cache coverage is page-aligned, while scheduler/radix-visible hit length remains the valid token length. This records the contract, adds page-aligned extent metadata, keeps owner assignment on actual tail pages instead of short-prefix fallback, and updates partial current reuse tests around tail-page masking. Constraint: CP owner lanes operate on page units while scheduler and radix hit accounting must remain token-valid. Rejected: Pad short suffixes to cp_size or 2*cp_size pages | wastes KV capacity and can turn a small tail into a much larger physical span. Rejected: Silent direct-write or prefetch fallback | production fallback must be warning-visible for diagnosis. Confidence: medium Scope-risk: moderate Directive: Do not reintroduce replicated short-radix fallback without checking docs/advanced_features/nsa_prefill_cp_page_aligned_cache_contract.md. Tested: local py_compile for touched runtime, utility, owner, and unit-test files. Tested: remote g0034 container three-file suite: 122 passed, 5 warnings. Not-tested: full local pytest, blocked by missing runtime dependencies such as orjson. Not-tested: CUDA E2E runtime for this commit. Co-authored-by: OmX <omx@oh-my-codex.dev>
829 lines
29 KiB
Python
829 lines
29 KiB
Python
import unittest
|
|
import sys
|
|
from types import SimpleNamespace
|
|
from unittest.mock import patch
|
|
|
|
from sglang.srt.layers.attention.nsa.utils import (
|
|
NSAContextParallelMetadata,
|
|
PageAlignedCacheExtent,
|
|
build_page_aligned_cache_extent,
|
|
_get_in_seq_last_token_owner_and_offset,
|
|
build_page_aligned_in_seq_split_list,
|
|
build_token_balanced_in_seq_split_list,
|
|
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,
|
|
get_cp_local_embedding_padded_token_count,
|
|
pad_cp_local_input_ids_for_embedding,
|
|
split_in_seq_cp_local_pair,
|
|
)
|
|
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
|
from sglang.test.ci.ci_register import register_cpu_ci
|
|
|
|
register_cpu_ci(est_time=1, suite="stage-a-test-cpu")
|
|
|
|
|
|
class TestPageAlignedCacheExtent(unittest.TestCase):
|
|
def test_extent_uses_page_boundary_not_cp_size(self):
|
|
extent = build_page_aligned_cache_extent(valid_tokens=100, page_size=64)
|
|
|
|
self.assertEqual(extent.valid_tokens, 100)
|
|
self.assertEqual(extent.padded_pages, 2)
|
|
self.assertEqual(extent.padded_tokens, 128)
|
|
self.assertEqual(extent.padding_tokens, 28)
|
|
|
|
def test_extent_handles_empty_and_aligned_lengths(self):
|
|
self.assertEqual(
|
|
build_page_aligned_cache_extent(valid_tokens=0, page_size=64),
|
|
PageAlignedCacheExtent(
|
|
valid_tokens=0,
|
|
padded_pages=0,
|
|
padded_tokens=0,
|
|
padding_tokens=0,
|
|
),
|
|
)
|
|
self.assertEqual(
|
|
build_page_aligned_cache_extent(valid_tokens=128, page_size=64),
|
|
PageAlignedCacheExtent(
|
|
valid_tokens=128,
|
|
padded_pages=2,
|
|
padded_tokens=128,
|
|
padding_tokens=0,
|
|
),
|
|
)
|
|
|
|
|
|
class TestNSAInSeqCPUtils(unittest.TestCase):
|
|
def test_contiguous_valid_cp_query_count(self):
|
|
from sglang.srt.layers.attention.nsa.nsa_indexer import (
|
|
_compute_contiguous_valid_cp_query_count,
|
|
)
|
|
|
|
self.assertEqual(
|
|
_compute_contiguous_valid_cp_query_count(
|
|
cp_kv_end=1024,
|
|
actual_seq_q=128,
|
|
logical_kv_limit=1024,
|
|
),
|
|
128,
|
|
)
|
|
self.assertEqual(
|
|
_compute_contiguous_valid_cp_query_count(
|
|
cp_kv_end=1100,
|
|
actual_seq_q=128,
|
|
logical_kv_limit=1024,
|
|
),
|
|
52,
|
|
)
|
|
self.assertEqual(
|
|
_compute_contiguous_valid_cp_query_count(
|
|
cp_kv_end=1100,
|
|
actual_seq_q=64,
|
|
logical_kv_limit=1000,
|
|
),
|
|
0,
|
|
)
|
|
self.assertEqual(
|
|
_compute_contiguous_valid_cp_query_count(
|
|
cp_kv_end=100,
|
|
actual_seq_q=0,
|
|
logical_kv_limit=100,
|
|
),
|
|
0,
|
|
)
|
|
|
|
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_exposes_padded_extent_without_padding_split_list(self):
|
|
split_list, info = build_page_aligned_in_seq_split_list(
|
|
total_len=100,
|
|
extend_len=100,
|
|
extend_prefix_len=0,
|
|
page_size=64,
|
|
cp_size=8,
|
|
)
|
|
|
|
self.assertTrue(info.page_aligned)
|
|
self.assertEqual(sum(split_list), 100)
|
|
self.assertEqual(split_list[:2], [64, 36])
|
|
self.assertEqual(split_list[2:], [0] * 14)
|
|
self.assertEqual(info.extend_valid_tokens, 100)
|
|
self.assertEqual(info.extend_padded_pages, 2)
|
|
self.assertEqual(info.extend_padded_tokens, 128)
|
|
self.assertEqual(info.extend_padding_tokens, 28)
|
|
|
|
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_pads_zero_segments_when_page_units_are_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.assertTrue(info.page_aligned)
|
|
self.assertEqual(sum(split_list), 512)
|
|
self.assertEqual(split_list[:8], [64] * 8)
|
|
self.assertEqual(split_list[8:], [0] * 8)
|
|
self.assert_page_aligned_boundaries(
|
|
split_list, extend_prefix_len=0, extend_len=512, page_size=64
|
|
)
|
|
|
|
def test_page_aligned_split_allows_radix_hit_suffix_with_one_page_per_rank(self):
|
|
split_list, info = build_page_aligned_in_seq_split_list(
|
|
total_len=512,
|
|
extend_len=512,
|
|
extend_prefix_len=54464,
|
|
page_size=64,
|
|
cp_size=8,
|
|
)
|
|
|
|
self.assertTrue(info.page_aligned)
|
|
self.assertEqual(sum(split_list), 512)
|
|
self.assertEqual(split_list[:8], [64] * 8)
|
|
self.assertEqual(split_list[8:], [0] * 8)
|
|
self.assert_page_aligned_boundaries(
|
|
split_list, extend_prefix_len=54464, extend_len=512, page_size=64
|
|
)
|
|
|
|
def test_page_aligned_split_keeps_short_radix_hit_suffix_page_aligned(self):
|
|
split_list, info = build_page_aligned_in_seq_split_list(
|
|
total_len=256,
|
|
extend_len=256,
|
|
extend_prefix_len=54464,
|
|
page_size=64,
|
|
cp_size=8,
|
|
)
|
|
|
|
self.assertTrue(info.page_aligned)
|
|
self.assertEqual(sum(split_list), 256)
|
|
self.assertEqual(split_list[:4], [64] * 4)
|
|
self.assertEqual(split_list[4:], [0] * 12)
|
|
self.assert_page_aligned_boundaries(
|
|
split_list, extend_prefix_len=54464, extend_len=256, page_size=64
|
|
)
|
|
|
|
def test_can_cp_split_keeps_cp_for_short_radix_hit_suffix(self):
|
|
class Mode:
|
|
def is_context_parallel_extend(self):
|
|
return True
|
|
|
|
forward_batch = SimpleNamespace(
|
|
uses_cp_shared_kv=True,
|
|
extend_seq_lens_cpu=[256],
|
|
extend_prefix_lens_cpu=[54464],
|
|
token_to_kv_pool=SimpleNamespace(page_size=64),
|
|
forward_mode=Mode(),
|
|
)
|
|
|
|
with (
|
|
patch(
|
|
"sglang.srt.layers.attention.nsa.utils.is_nsa_prefill_cp_round_robin_split",
|
|
return_value=False,
|
|
),
|
|
patch(
|
|
"sglang.srt.layers.attention.nsa.utils.is_nsa_enable_prefill_cp",
|
|
return_value=True,
|
|
),
|
|
):
|
|
self.assertTrue(can_cp_split(256, 8, True, forward_batch))
|
|
|
|
def test_can_cp_split_keeps_cp_for_radix_hit_suffix_with_one_page_per_rank(self):
|
|
class Mode:
|
|
def is_context_parallel_extend(self):
|
|
return True
|
|
|
|
forward_batch = SimpleNamespace(
|
|
uses_cp_shared_kv=True,
|
|
extend_seq_lens_cpu=[512],
|
|
extend_prefix_lens_cpu=[54464],
|
|
token_to_kv_pool=SimpleNamespace(page_size=64),
|
|
forward_mode=Mode(),
|
|
)
|
|
|
|
with (
|
|
patch(
|
|
"sglang.srt.layers.attention.nsa.utils.is_nsa_prefill_cp_round_robin_split",
|
|
return_value=False,
|
|
),
|
|
patch(
|
|
"sglang.srt.layers.attention.nsa.utils.is_nsa_enable_prefill_cp",
|
|
return_value=True,
|
|
),
|
|
):
|
|
self.assertTrue(can_cp_split(512, 8, True, forward_batch))
|
|
|
|
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
|
|
# segment 2, not in rank 0's trailing padded segment.
|
|
split_list = [4] * 16
|
|
|
|
owner, local_offset = _get_in_seq_last_token_owner_and_offset(
|
|
split_list=split_list,
|
|
cp_size=8,
|
|
actual_token_count=11,
|
|
)
|
|
|
|
self.assertEqual(owner, 2)
|
|
self.assertEqual(local_offset, 2)
|
|
|
|
def test_last_token_owner_keeps_existing_unpadded_fast_path_location(self):
|
|
split_list = [4] * 16
|
|
|
|
owner, local_offset = _get_in_seq_last_token_owner_and_offset(
|
|
split_list=split_list,
|
|
cp_size=8,
|
|
actual_token_count=64,
|
|
)
|
|
|
|
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")
|
|
|
|
def test_cp_split_and_rebuild_1d_matches_in_seq_zigzag_order(self):
|
|
import torch
|
|
from types import SimpleNamespace
|
|
|
|
forward_batch = SimpleNamespace(
|
|
nsa_cp_metadata=NSAContextParallelMetadata(
|
|
split_list=[2, 2, 2, 2, 2, 2, 2, 2],
|
|
zigzag_index=[1, 6],
|
|
)
|
|
)
|
|
|
|
local_locs = cp_split_and_rebuild_1d(forward_batch, torch.arange(16))
|
|
|
|
self.assertEqual(local_locs.tolist(), [2, 3, 12, 13])
|
|
|
|
def test_cp_local_embedding_pad_len_uses_metadata_max_rank_len(self):
|
|
from types import SimpleNamespace
|
|
|
|
import torch
|
|
|
|
forward_batch = SimpleNamespace(
|
|
nsa_cp_metadata=NSAContextParallelMetadata(max_rank_len=[4096] * 8)
|
|
)
|
|
|
|
self.assertEqual(
|
|
get_cp_local_embedding_padded_token_count(forward_batch, 4040), 4096
|
|
)
|
|
self.assertEqual(
|
|
get_cp_local_embedding_padded_token_count(forward_batch, 4096), 4096
|
|
)
|
|
self.assertEqual(
|
|
pad_cp_local_input_ids_for_embedding(
|
|
SimpleNamespace(
|
|
nsa_cp_metadata=NSAContextParallelMetadata(max_rank_len=[6] * 8)
|
|
),
|
|
torch.tensor([11, 12, 13, 14]),
|
|
).tolist(),
|
|
[11, 12, 13, 14, 0, 0],
|
|
)
|
|
self.assertEqual(
|
|
pad_cp_local_input_ids_for_embedding(
|
|
SimpleNamespace(
|
|
nsa_cp_metadata=NSAContextParallelMetadata(max_rank_len=[4] * 8)
|
|
),
|
|
torch.tensor([11, 12, 13, 14]),
|
|
).tolist(),
|
|
[11, 12, 13, 14],
|
|
)
|
|
|
|
missing_metadata = SimpleNamespace(nsa_cp_metadata=None)
|
|
self.assertIsNone(
|
|
get_cp_local_embedding_padded_token_count(missing_metadata, 4040)
|
|
)
|
|
self.assertIsNone(
|
|
pad_cp_local_input_ids_for_embedding(
|
|
missing_metadata, torch.tensor([11, 12, 13, 14])
|
|
)
|
|
)
|
|
|
|
stale_metadata = SimpleNamespace(
|
|
nsa_cp_metadata=NSAContextParallelMetadata(max_rank_len=[4039] * 8)
|
|
)
|
|
self.assertIsNone(
|
|
get_cp_local_embedding_padded_token_count(stale_metadata, 4040)
|
|
)
|
|
|
|
def test_local_out_cache_loc_requires_compute_owner_pages(self):
|
|
import torch
|
|
from types import SimpleNamespace
|
|
|
|
page_size = 4
|
|
# Segment order for cp_size=4, cp_rank=1 is segment 1 then 6.
|
|
# The logical page ids below deliberately encode the same owners through
|
|
# (logical_page - 1) % cp_size:
|
|
# segment 1 -> page 2 owner 1
|
|
# segment 6 -> page 6 owner 1
|
|
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,
|
|
)
|
|
|
|
local_locs = get_cp_shared_kv_local_out_cache_loc(forward_batch)
|
|
|
|
self.assertIsNotNone(local_locs)
|
|
self.assertEqual(
|
|
local_locs.tolist(),
|
|
list(range(2 * page_size, 3 * page_size))
|
|
+ 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
|
|
|
|
page_size = 4
|
|
out_cache_loc = torch.arange(page_size * 8, page_size * 16)
|
|
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,
|
|
)
|
|
|
|
self.assertIsNone(get_cp_shared_kv_local_out_cache_loc(forward_batch))
|
|
|
|
def test_local_out_cache_loc_logs_every_fallback_event(self):
|
|
import torch
|
|
from types import SimpleNamespace
|
|
|
|
from sglang.srt.layers.attention.nsa import utils as nsa_utils
|
|
|
|
page_size = 4
|
|
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=False,
|
|
page_size=page_size,
|
|
extend_prefix_len=0,
|
|
),
|
|
out_cache_loc=torch.arange(page_size * 8, page_size * 16),
|
|
)
|
|
|
|
with self.assertLogs(
|
|
"sglang.srt.layers.attention.nsa.utils", level="WARNING"
|
|
) as cm:
|
|
self.assertIsNone(get_cp_shared_kv_local_out_cache_loc(forward_batch))
|
|
self.assertIsNone(get_cp_shared_kv_local_out_cache_loc(forward_batch))
|
|
|
|
self.assertEqual(len(cm.output), 2)
|
|
self.assertIn("[CP_SHARED_KV_FALLBACK][direct_write]", cm.output[0])
|
|
self.assertIn("metadata is not page-aligned", cm.output[0])
|
|
self.assertIn("[CP_SHARED_KV_FALLBACK][direct_write]", cm.output[1])
|
|
self.assertIn("metadata is not page-aligned", cm.output[1])
|
|
|
|
def test_indexer_direct_write_does_not_log_missing_metadata_for_non_cp_batch(self):
|
|
import torch
|
|
from types import SimpleNamespace
|
|
|
|
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
|
|
|
|
indexer = object.__new__(Indexer)
|
|
indexer.nsa_enable_prefill_cp = True
|
|
forward_batch = SimpleNamespace(
|
|
uses_cp_shared_kv=True,
|
|
nsa_cp_metadata=None,
|
|
)
|
|
|
|
with self.assertNoLogs(
|
|
"sglang.srt.layers.attention.nsa.utils", level="INFO"
|
|
):
|
|
stored = Indexer._store_cp_shared_local_index_k_cache(
|
|
indexer,
|
|
forward_batch,
|
|
layer_id=0,
|
|
local_key=torch.empty(0),
|
|
act_quant=None,
|
|
)
|
|
|
|
self.assertFalse(stored)
|
|
|
|
def test_indexer_in_seq_cp_pair_materializes_index_once_for_prev_next(self):
|
|
import torch
|
|
|
|
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
|
|
|
|
indexer = object.__new__(Indexer)
|
|
logical_pages = torch.tensor([[1, 2, 3, 4]], dtype=torch.int32)
|
|
materialized_index = torch.tensor([11], dtype=torch.int32)
|
|
dense_pages = torch.tensor([[1, 2, 3, 4]], dtype=torch.int32)
|
|
materialize_calls = []
|
|
topk_calls = []
|
|
|
|
class Metadata:
|
|
def get_page_table_64(self):
|
|
return logical_pages
|
|
|
|
def fake_materialize(forward_batch, layer_id, logical_page_table):
|
|
materialize_calls.append((layer_id, logical_page_table))
|
|
return materialized_index, dense_pages
|
|
|
|
def fake_get_topk(
|
|
forward_batch,
|
|
layer_id,
|
|
q_fp8,
|
|
weights,
|
|
metadata,
|
|
kv_len,
|
|
actual_seq_q,
|
|
cp_index=None,
|
|
current_index_kv=None,
|
|
shared_index_buffer=None,
|
|
shared_block_tables=None,
|
|
actual_seq_q_tensor=None,
|
|
actual_seq_q_cu_tensor=None,
|
|
):
|
|
topk_calls.append(
|
|
{
|
|
"kv_len": kv_len,
|
|
"actual_seq_q": actual_seq_q,
|
|
"actual_seq_q_tensor": actual_seq_q_tensor,
|
|
"actual_seq_q_cu_tensor": actual_seq_q_cu_tensor,
|
|
"shared_index_buffer": shared_index_buffer,
|
|
"shared_block_tables": shared_block_tables,
|
|
"current_index_kv": current_index_kv,
|
|
}
|
|
)
|
|
return torch.full((actual_seq_q, 2), len(topk_calls), dtype=torch.int32)
|
|
|
|
indexer._maybe_materialize_shared_index_buffer = fake_materialize
|
|
indexer._get_topk_ragged_with_cp = fake_get_topk
|
|
|
|
forward_batch = type(
|
|
"ForwardBatchStub",
|
|
(),
|
|
{
|
|
"nsa_cp_metadata": NSAContextParallelMetadata(
|
|
kv_len_prev=5,
|
|
kv_len_next=9,
|
|
actual_seq_q_prev=3,
|
|
actual_seq_q_next=2,
|
|
actual_seq_q_prev_cu_tensor=torch.tensor([0, 3], dtype=torch.int32),
|
|
actual_seq_q_next_cu_tensor=torch.tensor([0, 2], dtype=torch.int32),
|
|
)
|
|
},
|
|
)()
|
|
q_fp8 = torch.arange(5 * 4, dtype=torch.float32).view(5, 4)
|
|
weights = torch.arange(5 * 2, dtype=torch.float32).view(5, 2)
|
|
|
|
result = Indexer._get_topk_in_seq_cp_pair(
|
|
indexer,
|
|
forward_batch,
|
|
layer_id=7,
|
|
q_fp8=q_fp8,
|
|
weights=weights,
|
|
metadata=Metadata(),
|
|
current_index_kv=None,
|
|
)
|
|
|
|
self.assertEqual(len(materialize_calls), 1)
|
|
self.assertIs(materialize_calls[0][1], logical_pages)
|
|
self.assertEqual(len(topk_calls), 2)
|
|
self.assertIs(topk_calls[0]["shared_index_buffer"], materialized_index)
|
|
self.assertIs(topk_calls[1]["shared_index_buffer"], materialized_index)
|
|
self.assertIs(topk_calls[0]["shared_block_tables"], dense_pages)
|
|
self.assertIs(topk_calls[1]["shared_block_tables"], dense_pages)
|
|
self.assertIsNone(topk_calls[0]["current_index_kv"])
|
|
self.assertEqual(topk_calls[0]["kv_len"], 5)
|
|
self.assertEqual(topk_calls[1]["kv_len"], 9)
|
|
self.assertEqual(topk_calls[0]["actual_seq_q_cu_tensor"].tolist(), [0, 3])
|
|
self.assertEqual(topk_calls[1]["actual_seq_q_cu_tensor"].tolist(), [0, 2])
|
|
self.assertEqual(result.tolist(), [[1, 1], [1, 1], [1, 1], [2, 2], [2, 2]])
|
|
|
|
def test_indexer_in_seq_cp_pair_skips_materialize_when_current_index_reused(self):
|
|
import torch
|
|
|
|
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
|
|
|
|
indexer = object.__new__(Indexer)
|
|
current_index_kv = (torch.tensor([1]), torch.tensor([2]))
|
|
materialize_calls = []
|
|
topk_calls = []
|
|
|
|
class Metadata:
|
|
def get_page_table_64(self):
|
|
raise AssertionError("current index reuse should not read page table")
|
|
|
|
def fake_materialize(forward_batch, layer_id, logical_page_table):
|
|
materialize_calls.append((layer_id, logical_page_table))
|
|
raise AssertionError("current index reuse should not materialize")
|
|
|
|
def fake_get_topk(
|
|
forward_batch,
|
|
layer_id,
|
|
q_fp8,
|
|
weights,
|
|
metadata,
|
|
kv_len,
|
|
actual_seq_q,
|
|
cp_index=None,
|
|
current_index_kv=None,
|
|
shared_index_buffer=None,
|
|
shared_block_tables=None,
|
|
actual_seq_q_tensor=None,
|
|
actual_seq_q_cu_tensor=None,
|
|
):
|
|
topk_calls.append(
|
|
{
|
|
"current_index_kv": current_index_kv,
|
|
"actual_seq_q_tensor": actual_seq_q_tensor,
|
|
"actual_seq_q_cu_tensor": actual_seq_q_cu_tensor,
|
|
"shared_index_buffer": shared_index_buffer,
|
|
"shared_block_tables": shared_block_tables,
|
|
}
|
|
)
|
|
return torch.full((actual_seq_q, 2), len(topk_calls), dtype=torch.int32)
|
|
|
|
indexer._maybe_materialize_shared_index_buffer = fake_materialize
|
|
indexer._get_topk_ragged_with_cp = fake_get_topk
|
|
|
|
forward_batch = type(
|
|
"ForwardBatchStub",
|
|
(),
|
|
{
|
|
"nsa_cp_metadata": NSAContextParallelMetadata(
|
|
kv_len_prev=5,
|
|
kv_len_next=9,
|
|
actual_seq_q_prev=3,
|
|
actual_seq_q_next=2,
|
|
actual_seq_q_prev_cu_tensor=torch.tensor([0, 3], dtype=torch.int32),
|
|
actual_seq_q_next_cu_tensor=torch.tensor([0, 2], dtype=torch.int32),
|
|
)
|
|
},
|
|
)()
|
|
|
|
result = Indexer._get_topk_in_seq_cp_pair(
|
|
indexer,
|
|
forward_batch,
|
|
layer_id=7,
|
|
q_fp8=torch.empty(5, 4),
|
|
weights=torch.empty(5, 2),
|
|
metadata=Metadata(),
|
|
current_index_kv=current_index_kv,
|
|
)
|
|
|
|
self.assertEqual(materialize_calls, [])
|
|
self.assertEqual(len(topk_calls), 2)
|
|
self.assertIs(topk_calls[0]["current_index_kv"], current_index_kv)
|
|
self.assertIs(topk_calls[1]["current_index_kv"], current_index_kv)
|
|
self.assertIsNone(topk_calls[0]["shared_index_buffer"])
|
|
self.assertIsNone(topk_calls[1]["shared_block_tables"])
|
|
self.assertEqual(topk_calls[0]["actual_seq_q_cu_tensor"].tolist(), [0, 3])
|
|
self.assertEqual(topk_calls[1]["actual_seq_q_cu_tensor"].tolist(), [0, 2])
|
|
self.assertEqual(result.tolist(), [[1, 1], [1, 1], [1, 1], [2, 2], [2, 2]])
|
|
|
|
def test_paged_topk_transform_uses_cu_override_without_scan_metadata_ops(self):
|
|
import torch
|
|
|
|
from sglang.srt.layers.attention.nsa_backend import (
|
|
NSAMetadata,
|
|
NSAIndexerMetadata,
|
|
TopkTransformMethod,
|
|
)
|
|
|
|
cu_override = torch.tensor([0, 4], dtype=torch.int32)
|
|
attn_metadata = NSAMetadata(
|
|
page_size=64,
|
|
cache_seqlens_int32=torch.tensor([4], dtype=torch.int32),
|
|
max_seq_len_q=4,
|
|
max_seq_len_k=8,
|
|
cu_seqlens_q=torch.tensor([0, 4], dtype=torch.int32),
|
|
cu_seqlens_k=torch.tensor([0, 8], dtype=torch.int32),
|
|
page_table_1=torch.arange(8, dtype=torch.int32).view(1, 8),
|
|
real_page_table=torch.arange(8, dtype=torch.int32).view(1, 8),
|
|
nsa_cache_seqlens_int32=torch.tensor([4], dtype=torch.int32),
|
|
nsa_cu_seqlens_q=torch.arange(2, dtype=torch.int32),
|
|
nsa_cu_seqlens_k=torch.tensor([0, 4], dtype=torch.int32),
|
|
nsa_extend_seq_lens_list=[4],
|
|
nsa_seqlens_expanded=torch.arange(1, 5, dtype=torch.int32),
|
|
topk_indices_offset=torch.zeros(4, dtype=torch.int32),
|
|
)
|
|
metadata = NSAIndexerMetadata(
|
|
attn_metadata=attn_metadata,
|
|
topk_transform_method=TopkTransformMethod.PAGED,
|
|
)
|
|
logits = torch.zeros(4, 8)
|
|
lengths = torch.arange(1, 5, dtype=torch.int32)
|
|
expected = torch.full((4, 2), 7, dtype=torch.int32)
|
|
|
|
def fake_fused(**kwargs):
|
|
self.assertIs(kwargs["cu_seqlens_q"], cu_override)
|
|
self.assertIs(kwargs["lengths"], lengths)
|
|
return expected
|
|
|
|
fake_sgl_kernel = SimpleNamespace(
|
|
fast_topk_transform_fused=fake_fused,
|
|
fast_topk_transform_ragged_fused=lambda **_: (_ for _ in ()).throw(
|
|
AssertionError("ragged path should not run")
|
|
),
|
|
fast_topk_v2=lambda *_, **__: (_ for _ in ()).throw(
|
|
AssertionError("unfused path should not run")
|
|
),
|
|
)
|
|
|
|
with (
|
|
patch.dict(sys.modules, {"sgl_kernel": fake_sgl_kernel}),
|
|
patch(
|
|
"sglang.srt.layers.attention.nsa_backend.envs.SGLANG_NSA_FUSE_TOPK.get",
|
|
return_value=True,
|
|
),
|
|
patch(
|
|
"sglang.srt.layers.attention.nsa_backend.compute_cu_seqlens",
|
|
side_effect=AssertionError("paged override should skip cumsum"),
|
|
),
|
|
patch(
|
|
"torch.repeat_interleave",
|
|
side_effect=AssertionError("paged topk should not build ragged offsets"),
|
|
),
|
|
):
|
|
actual = metadata.topk_transform(
|
|
logits,
|
|
topk=2,
|
|
cu_seqlens_q=torch.tensor([4], dtype=torch.int32),
|
|
ke_offset=lengths,
|
|
cu_seqlens_q_topk_override=cu_override,
|
|
)
|
|
|
|
self.assertIs(actual, expected)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|