Batch-size support needs request-first CP metadata; treating a batch as one long sequence breaks page ownership, top-k ranges, and phase1 compact output collection. This adds a batch CP plan that records per-request page-aligned splits, rank-local offsets, kv/actual-seq metadata, last-token owners, and flattened descriptors for downstream allocator/runtime workstreams. The scalar full-rerange path now fail-fasts for batch metadata so bs>1 cannot silently discard the narrow-output optimization or restore hidden states with single-request assumptions. Constraint: CP shared-KV cache state is page-owned and must preserve request boundaries under bs>1. Rejected: Let bs>1 fall back to scalar full hidden rerange | it loses the phase1 communication reduction and uses wrong single-request metadata. Rejected: Add a collective to confirm batch plans | all ranks can derive the same plan from CPU metadata and config. Confidence: medium Scope-risk: moderate Directive: Do not remove batch fail-fast guards until W2/W3 consumers use CPSharedKVBatchPlan end-to-end. Tested: python -m py_compile python/sglang/srt/layers/attention/nsa/utils.py test/registered/unit/layers/test_nsa_cp_utils.py Tested: remote g0034 container PYTHONPATH=python python -m pytest -q test/registered/unit/layers/test_nsa_cp_utils.py -> 39 passed Not-tested: full ETE bs>1 CP shared-KV runtime; W2/W3 allocator/direct-write consumers are not implemented yet
1131 lines
40 KiB
Python
1131 lines
40 KiB
Python
import ast
|
|
from pathlib import Path
|
|
import unittest
|
|
import sys
|
|
from types import SimpleNamespace
|
|
from unittest.mock import patch
|
|
|
|
from sglang.srt.layers.attention.nsa.utils import (
|
|
NSAContextParallelMetadata,
|
|
PageAlignedCacheExtent,
|
|
build_flat_page_owner_plan,
|
|
build_batch_page_aligned_in_seq_split_plan,
|
|
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_all_gather_rerange_output,
|
|
cp_collect_last_token_hidden,
|
|
cp_split_and_rebuild_1d,
|
|
cp_split_and_rebuild_data,
|
|
get_cp_shared_kv_batch_plan,
|
|
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_tensor_by_cp_batch_plan,
|
|
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_skips_cp_when_radix_hit_suffix_has_too_few_pages(self):
|
|
class Mode:
|
|
def is_context_parallel_extend(self):
|
|
return True
|
|
|
|
forward_batch = SimpleNamespace(
|
|
uses_cp_shared_kv=True,
|
|
extend_seq_lens_cpu=[65],
|
|
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.assertFalse(can_cp_split(128, 8, True, forward_batch))
|
|
|
|
def test_can_cp_split_skips_cp_when_page_units_do_not_cover_all_lanes(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.assertFalse(can_cp_split(256, 8, True, forward_batch))
|
|
|
|
def test_can_cp_split_skips_current_only_when_page_units_do_not_cover_all_lanes(
|
|
self,
|
|
):
|
|
class Mode:
|
|
def is_context_parallel_extend(self):
|
|
return True
|
|
|
|
forward_batch = SimpleNamespace(
|
|
uses_cp_shared_kv=True,
|
|
extend_seq_lens_cpu=[64],
|
|
extend_prefix_lens_cpu=[0],
|
|
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.assertFalse(can_cp_split(64, 8, True, forward_batch))
|
|
|
|
def test_can_cp_split_fails_on_non_page_aligned_cp_shared_prefix(self):
|
|
class Mode:
|
|
def is_context_parallel_extend(self):
|
|
return True
|
|
|
|
forward_batch = SimpleNamespace(
|
|
uses_cp_shared_kv=True,
|
|
extend_seq_lens_cpu=[1024],
|
|
extend_prefix_lens_cpu=[65],
|
|
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.assertRaisesRegex(
|
|
RuntimeError,
|
|
r"\[CP_SHARED_KV_FAIL_FAST\]\[cp_split_non_page_aligned_prefix\]",
|
|
),
|
|
):
|
|
can_cp_split(1089, 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_batch_page_aligned_plan_keeps_request_boundaries_and_last_token_owners(
|
|
self,
|
|
):
|
|
plan = build_batch_page_aligned_in_seq_split_plan(
|
|
extend_lens=[4, 9],
|
|
prefix_lens=[0, 8],
|
|
page_size=4,
|
|
cp_size=2,
|
|
cp_rank=0,
|
|
)
|
|
|
|
self.assertEqual(plan.batch_size, 2)
|
|
self.assertEqual(plan.request_split_lists, [[4, 0, 0, 0], [4, 4, 1, 0]])
|
|
self.assertEqual(plan.request_padded_pages, [1, 3])
|
|
self.assertEqual(plan.request_padded_tokens, [4, 12])
|
|
self.assertEqual(plan.request_token_offsets, [0, 4])
|
|
self.assertEqual(plan.request_padded_token_offsets, [0, 4])
|
|
self.assertEqual(plan.request_page_offsets, [0, 1])
|
|
self.assertEqual(plan.request_last_token_owner, [0, 1])
|
|
self.assertEqual(plan.request_last_token_local_offset, [3, 4])
|
|
self.assertEqual(plan.request_rank_local_tokens, [4, 4])
|
|
self.assertEqual(plan.request_rank_local_offsets, [0, 4])
|
|
self.assertEqual(plan.request_kv_len_prev, [4, 4])
|
|
self.assertEqual(plan.request_kv_len_next, [4, 9])
|
|
self.assertEqual(plan.request_actual_seq_q_prev, [4, 4])
|
|
self.assertEqual(plan.request_actual_seq_q_next, [0, 0])
|
|
self.assertEqual(plan.flat_segment_request_ids, [0, 0, 0, 0, 1, 1, 1, 1])
|
|
self.assertEqual(plan.flat_segment_offsets, [0, 4, 4, 4, 0, 4, 8, 9])
|
|
|
|
def test_batch_plan_stable_helpers_split_and_build_page_owner_plan(self):
|
|
import torch
|
|
|
|
plan = build_batch_page_aligned_in_seq_split_plan(
|
|
extend_lens=[4, 9],
|
|
prefix_lens=[0, 8],
|
|
page_size=4,
|
|
cp_size=2,
|
|
cp_rank=1,
|
|
)
|
|
forward_batch = SimpleNamespace(nsa_cp_metadata=SimpleNamespace(batch_plan=plan))
|
|
|
|
self.assertIs(get_cp_shared_kv_batch_plan(forward_batch), plan)
|
|
self.assertEqual(build_flat_page_owner_plan(plan), [0, 0, 1, 1])
|
|
|
|
local_1d = split_tensor_by_cp_batch_plan(torch.arange(13), plan, mode="1d")
|
|
self.assertEqual(local_1d.tolist(), [8, 9, 10, 11, 12])
|
|
|
|
local_data = split_tensor_by_cp_batch_plan(
|
|
torch.arange(13 * 2).view(13, 2), plan, mode="data"
|
|
)
|
|
self.assertEqual(local_data[:, 0].tolist(), list(range(16, 26, 2)))
|
|
|
|
def test_collect_last_token_hidden_uses_batch_owner_metadata(self):
|
|
import torch
|
|
|
|
hidden_states = torch.tensor(
|
|
[[10.0], [11.0], [12.0], [13.0], [20.0], [21.0], [22.0], [23.0]]
|
|
)
|
|
forward_batch = SimpleNamespace(
|
|
extend_seq_lens_cpu=[4, 9],
|
|
nsa_cp_metadata=NSAContextParallelMetadata(
|
|
batch_size=2,
|
|
request_last_token_owner=[0, 1],
|
|
request_last_token_local_offset=[3, 4],
|
|
request_rank_local_offsets=[0, 4],
|
|
),
|
|
)
|
|
|
|
def fake_all_gather(output, local_last):
|
|
self.assertEqual(local_last.tolist(), [[13.0], [0.0]])
|
|
output.copy_(torch.tensor([[13.0], [0.0], [0.0], [99.0]]))
|
|
|
|
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.get_attention_cp_rank",
|
|
return_value=0,
|
|
),
|
|
patch(
|
|
"sglang.srt.layers.attention.nsa.utils.attn_cp_all_gather_into_tensor",
|
|
side_effect=fake_all_gather,
|
|
),
|
|
):
|
|
collected = cp_collect_last_token_hidden(hidden_states, forward_batch, 2)
|
|
|
|
self.assertEqual(collected.tolist(), [[13.0], [99.0]])
|
|
|
|
def test_collect_last_token_hidden_fails_fast_without_batch_owner_metadata(self):
|
|
import torch
|
|
|
|
forward_batch = SimpleNamespace(
|
|
extend_seq_lens_cpu=[4, 9],
|
|
nsa_cp_metadata=NSAContextParallelMetadata(batch_size=2),
|
|
)
|
|
|
|
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.get_attention_cp_rank",
|
|
return_value=0,
|
|
),
|
|
self.assertRaisesRegex(
|
|
RuntimeError,
|
|
r"\[CP_SHARED_KV_FAIL_FAST\]\[batch_gt1_missing_last_token_metadata\]",
|
|
),
|
|
):
|
|
cp_collect_last_token_hidden(torch.zeros((8, 1)), forward_batch, 2)
|
|
|
|
def test_full_rerange_fails_fast_for_batch_metadata(self):
|
|
import torch
|
|
|
|
forward_batch = SimpleNamespace(
|
|
nsa_cp_metadata=NSAContextParallelMetadata(batch_size=2)
|
|
)
|
|
|
|
with self.assertRaisesRegex(
|
|
RuntimeError,
|
|
r"\[CP_SHARED_KV_FAIL_FAST\]\[batch_gt1_full_rerange_unsupported\]",
|
|
):
|
|
cp_all_gather_rerange_output(
|
|
torch.zeros((8, 1)), 2, forward_batch, stream=None
|
|
)
|
|
|
|
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_split_and_rebuild_data_keeps_batch_request_boundaries(self):
|
|
import torch
|
|
|
|
forward_batch = SimpleNamespace(
|
|
nsa_cp_metadata=NSAContextParallelMetadata(
|
|
batch_size=2,
|
|
request_extend_lens=[4, 9],
|
|
request_split_lists=[[4, 0, 0, 0], [4, 4, 1, 0]],
|
|
request_zigzag_indices=[[0, 3], [0, 3]],
|
|
)
|
|
)
|
|
tensor = torch.arange(13 * 2).view(13, 2)
|
|
|
|
with patch(
|
|
"sglang.srt.layers.attention.nsa.utils.is_nsa_prefill_cp_round_robin_split",
|
|
return_value=False,
|
|
):
|
|
local = cp_split_and_rebuild_data(forward_batch, tensor)
|
|
|
|
self.assertEqual(local[:, 0].tolist(), list(range(0, 16, 2)))
|
|
|
|
def test_cp_split_and_rebuild_1d_keeps_batch_request_boundaries(self):
|
|
import torch
|
|
|
|
forward_batch = SimpleNamespace(
|
|
nsa_cp_metadata=NSAContextParallelMetadata(
|
|
batch_size=2,
|
|
request_extend_lens=[4, 9],
|
|
request_split_lists=[[4, 0, 0, 0], [4, 4, 1, 0]],
|
|
request_zigzag_indices=[[1, 2], [1, 2]],
|
|
)
|
|
)
|
|
|
|
with patch(
|
|
"sglang.srt.layers.attention.nsa.utils.is_nsa_prefill_cp_round_robin_split",
|
|
return_value=False,
|
|
):
|
|
local = cp_split_and_rebuild_1d(forward_batch, torch.arange(13))
|
|
|
|
self.assertEqual(local.tolist(), [8, 9, 10, 11, 12])
|
|
|
|
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
|
|
|
|
class Mode:
|
|
def is_extend_without_speculative(self):
|
|
return True
|
|
|
|
forward_batch = type(
|
|
"ForwardBatchStub",
|
|
(),
|
|
{
|
|
"forward_mode": Mode(),
|
|
"extend_prefix_lens_cpu": [0],
|
|
"extend_seq_lens_cpu": [5],
|
|
"seq_lens_cpu": torch.tensor([5], dtype=torch.int64),
|
|
"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
|
|
|
|
class Mode:
|
|
def is_extend_without_speculative(self):
|
|
return True
|
|
|
|
forward_batch = type(
|
|
"ForwardBatchStub",
|
|
(),
|
|
{
|
|
"forward_mode": Mode(),
|
|
"extend_prefix_lens_cpu": [0],
|
|
"extend_seq_lens_cpu": [5],
|
|
"seq_lens_cpu": torch.tensor([5], dtype=torch.int64),
|
|
"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)
|
|
|
|
def test_mla_partial_current_path_fails_fast_instead_of_compact_fallback(self):
|
|
backend_path = (
|
|
Path(__file__).resolve().parents[4]
|
|
/ "python/sglang/srt/layers/attention/nsa_backend.py"
|
|
)
|
|
source = backend_path.read_text()
|
|
tree = ast.parse(source)
|
|
compact_merge_calls = [
|
|
node
|
|
for node in ast.walk(tree)
|
|
if isinstance(node, ast.Call)
|
|
and isinstance(node.func, ast.Name)
|
|
and node.func.id == "merge_materialized_and_current_kv"
|
|
]
|
|
|
|
self.assertEqual(
|
|
compact_merge_calls,
|
|
[],
|
|
"MLA partial-current reuse must not fall back to compact "
|
|
"materialize/current merge when page-slot prefetch compose is unavailable.",
|
|
)
|
|
self.assertIn(
|
|
"[CP_SHARED_KV_FAIL_FAST][mla_partial_current_sync]",
|
|
source,
|
|
)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|