Reduce CP shared KV materialize and direct-write overhead
Shared KV now relies on page-aligned CP metadata and compute-owner page allocation so persistent MLA KV and NSA index shards can be written by the rank that computed them. The compatibility read path keeps the dense full-view contract for existing topk and attention kernels, but removes duplicated prev/next index materialize, adds optional tai materialize integration, and tightens tests/docs around the fallback boundaries. Constraint: Decode remains non-CP while prefill CP owns the shared-KV changes Constraint: Existing attention/topk kernels still expect dense full-view KV/index inputs Rejected: Change attention kernels to read owner-sharded KV directly | larger semantic change reserved for later phases Rejected: Merge index K/scale storage with MLA KV storage | would couple topk and attention cache lifecycles before materialize overhead is isolated Confidence: medium Scope-risk: broad Directive: Do not remove fallback logging or debug-gated assertions without reproducing long-context chunked/radix-hit paths Tested: git diff --check --cached Not-tested: Local pytest/runtime server verification not run in this commit step per current workflow constraints
This commit is contained in:
@@ -1,11 +1,18 @@
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from sglang.srt.layers.attention.nsa.utils import (
|
||||
NSAContextParallelMetadata,
|
||||
_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,
|
||||
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")
|
||||
@@ -93,6 +100,85 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
self.assertFalse(info.page_aligned)
|
||||
self.assertEqual(split_list, build_token_balanced_in_seq_split_list(512, 8))
|
||||
|
||||
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_falls_back_when_radix_hit_suffix_has_zero_rank(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.assertFalse(info.page_aligned)
|
||||
self.assertEqual(split_list, build_token_balanced_in_seq_split_list(256, 8))
|
||||
|
||||
def test_can_cp_split_uses_replicated_compute_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.assertFalse(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,
|
||||
@@ -152,6 +238,309 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
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_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_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="INFO"
|
||||
) 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 direct-write fallback", cm.output[0])
|
||||
self.assertIn("metadata is not page-aligned", cm.output[0])
|
||||
self.assertIn("CP shared KV direct-write fallback", 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,
|
||||
):
|
||||
topk_calls.append(
|
||||
{
|
||||
"kv_len": kv_len,
|
||||
"actual_seq_q": actual_seq_q,
|
||||
"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,
|
||||
)
|
||||
},
|
||||
)()
|
||||
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(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,
|
||||
):
|
||||
topk_calls.append(
|
||||
{
|
||||
"current_index_kv": current_index_kv,
|
||||
"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,
|
||||
)
|
||||
},
|
||||
)()
|
||||
|
||||
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(result.tolist(), [[1, 1], [1, 1], [1, 1], [2, 2], [2, 2]])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -52,6 +52,27 @@ class TestCpSharedKVLayout(unittest.TestCase):
|
||||
|
||||
|
||||
class TestCPSharedPagedAllocator(unittest.TestCase):
|
||||
def test_compute_owner_alloc_fallback_logs_every_event(self):
|
||||
from sglang.srt.mem_cache import common
|
||||
|
||||
with self.assertLogs("sglang.srt.mem_cache.common", level="INFO") as cm:
|
||||
common._log_cp_shared_kv_alloc_fallback(
|
||||
"too_short_for_page_aligned",
|
||||
"falling back for test event %s",
|
||||
1,
|
||||
)
|
||||
common._log_cp_shared_kv_alloc_fallback(
|
||||
"too_short_for_page_aligned",
|
||||
"falling back for test event %s",
|
||||
2,
|
||||
)
|
||||
|
||||
self.assertEqual(len(cm.output), 2)
|
||||
self.assertIn("too_short_for_page_aligned", cm.output[0])
|
||||
self.assertIn("test event 1", cm.output[0])
|
||||
self.assertIn("too_short_for_page_aligned", cm.output[1])
|
||||
self.assertIn("test event 2", cm.output[1])
|
||||
|
||||
def test_shared_allocator_exposes_logical_capacity(self):
|
||||
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
|
||||
|
||||
@@ -73,6 +94,229 @@ class TestCPSharedPagedAllocator(unittest.TestCase):
|
||||
allocator.free(locs)
|
||||
self.assertEqual(allocator.available_size(), 64 * 8)
|
||||
|
||||
def test_compute_owner_page_assignment_matches_in_seq_zigzag(self):
|
||||
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
|
||||
build_in_seq_page_compute_owners,
|
||||
)
|
||||
|
||||
owners = build_in_seq_page_compute_owners(
|
||||
extend_len=64 * 16,
|
||||
extend_prefix_len=0,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
owners,
|
||||
[0, 0, 1, 1, 2, 2, 3, 3, 3, 3, 2, 2, 1, 1, 0, 0],
|
||||
)
|
||||
|
||||
def test_compute_owner_page_assignment_falls_back_for_short_extend(self):
|
||||
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
|
||||
build_in_seq_page_compute_owners,
|
||||
get_in_seq_page_compute_owner_unavailable_reason,
|
||||
)
|
||||
|
||||
owners = build_in_seq_page_compute_owners(
|
||||
extend_len=64 * 7,
|
||||
extend_prefix_len=0,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
)
|
||||
|
||||
self.assertIsNone(owners)
|
||||
self.assertEqual(
|
||||
get_in_seq_page_compute_owner_unavailable_reason(
|
||||
extend_len=64 * 7,
|
||||
extend_prefix_len=0,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
),
|
||||
"too_short_for_page_aligned",
|
||||
)
|
||||
|
||||
def test_compute_owner_page_assignment_allows_radix_hit_suffix_with_one_page_per_rank(
|
||||
self,
|
||||
):
|
||||
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
|
||||
build_in_seq_page_compute_owners,
|
||||
get_in_seq_page_compute_owner_unavailable_reason,
|
||||
)
|
||||
|
||||
owners = build_in_seq_page_compute_owners(
|
||||
extend_len=64 * 4,
|
||||
extend_prefix_len=54464,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
)
|
||||
|
||||
self.assertEqual(owners, [0, 1, 2, 3])
|
||||
self.assertIsNone(
|
||||
get_in_seq_page_compute_owner_unavailable_reason(
|
||||
extend_len=64 * 4,
|
||||
extend_prefix_len=54464,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
)
|
||||
)
|
||||
|
||||
def test_compute_owner_page_assignment_allows_short_radix_hit_suffix_with_replicated_compute(
|
||||
self,
|
||||
):
|
||||
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
|
||||
build_in_seq_page_compute_owners,
|
||||
get_in_seq_page_compute_owner_unavailable_reason,
|
||||
)
|
||||
|
||||
owners = build_in_seq_page_compute_owners(
|
||||
extend_len=64 * 3,
|
||||
extend_prefix_len=54464,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
)
|
||||
|
||||
self.assertEqual(owners, [0, 1, 2])
|
||||
self.assertIsNone(
|
||||
get_in_seq_page_compute_owner_unavailable_reason(
|
||||
extend_len=64 * 3,
|
||||
extend_prefix_len=54464,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
)
|
||||
)
|
||||
|
||||
def test_compute_owner_page_assignment_reports_prefix_misalignment(self):
|
||||
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
|
||||
get_in_seq_page_compute_owner_unavailable_reason,
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
get_in_seq_page_compute_owner_unavailable_reason(
|
||||
extend_len=64 * 16,
|
||||
extend_prefix_len=1,
|
||||
page_size=64,
|
||||
cp_size=4,
|
||||
),
|
||||
"prefix_not_page_aligned",
|
||||
)
|
||||
|
||||
def test_shared_allocator_can_allocate_pages_from_compute_owner_lanes(self):
|
||||
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.cp_shared_kv_compute_owner import (
|
||||
build_in_seq_page_compute_owners,
|
||||
)
|
||||
|
||||
page_size = 64
|
||||
cp_size = 4
|
||||
owners = build_in_seq_page_compute_owners(
|
||||
extend_len=page_size * 16,
|
||||
extend_prefix_len=0,
|
||||
page_size=page_size,
|
||||
cp_size=cp_size,
|
||||
)
|
||||
allocator = CPSharedPagedTokenToKVPoolAllocator(
|
||||
logical_size=page_size * 32,
|
||||
physical_size=page_size * 8,
|
||||
page_size=page_size,
|
||||
dtype=torch.bfloat16,
|
||||
device="cpu",
|
||||
kvcache=None,
|
||||
need_sort=False,
|
||||
cp_size=cp_size,
|
||||
cp_rank=0,
|
||||
)
|
||||
|
||||
locs = allocator.alloc_extend_compute_owner(
|
||||
prefix_lens=torch.tensor([0], dtype=torch.int64),
|
||||
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
|
||||
seq_lens=torch.tensor([page_size * 16], dtype=torch.int64),
|
||||
seq_lens_cpu=torch.tensor([page_size * 16], dtype=torch.int64),
|
||||
last_loc=torch.tensor([-1], dtype=torch.int64),
|
||||
extend_num_tokens=page_size * 16,
|
||||
page_compute_owners=owners,
|
||||
)
|
||||
|
||||
self.assertIsNotNone(locs)
|
||||
logical_pages = locs.view(-1, page_size)[:, 0] // page_size
|
||||
self.assertEqual(
|
||||
((logical_pages - 1) % cp_size).tolist(),
|
||||
owners,
|
||||
)
|
||||
self.assertEqual(allocator.available_size(), page_size * 16)
|
||||
|
||||
def test_compute_owner_lane_eviction_recovers_exhausted_owner_lane(self):
|
||||
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import EvictResult
|
||||
from sglang.srt.mem_cache.common import _evict_for_compute_owner_lanes
|
||||
|
||||
page_size = 64
|
||||
cp_size = 4
|
||||
allocator = CPSharedPagedTokenToKVPoolAllocator(
|
||||
logical_size=page_size * 16,
|
||||
physical_size=page_size * 4,
|
||||
page_size=page_size,
|
||||
dtype=torch.bfloat16,
|
||||
device="cpu",
|
||||
kvcache=None,
|
||||
need_sort=False,
|
||||
cp_size=cp_size,
|
||||
cp_rank=0,
|
||||
)
|
||||
lane0_locs = allocator.alloc_extend_compute_owner(
|
||||
prefix_lens=torch.tensor([0], dtype=torch.int64),
|
||||
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
|
||||
seq_lens=torch.tensor([page_size * 4], dtype=torch.int64),
|
||||
seq_lens_cpu=torch.tensor([page_size * 4], dtype=torch.int64),
|
||||
last_loc=torch.tensor([-1], dtype=torch.int64),
|
||||
extend_num_tokens=page_size * 4,
|
||||
page_compute_owners=[0, 0, 0, 0],
|
||||
)
|
||||
self.assertIsNone(
|
||||
allocator.alloc_extend_compute_owner(
|
||||
prefix_lens=torch.tensor([0], dtype=torch.int64),
|
||||
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
|
||||
seq_lens=torch.tensor([page_size * 4], dtype=torch.int64),
|
||||
seq_lens_cpu=torch.tensor([page_size * 4], dtype=torch.int64),
|
||||
last_loc=torch.tensor([-1], dtype=torch.int64),
|
||||
extend_num_tokens=page_size * 4,
|
||||
page_compute_owners=[0, 0, 0, 0],
|
||||
)
|
||||
)
|
||||
|
||||
class FakeTreeCache:
|
||||
def __init__(self):
|
||||
self.evict_calls = []
|
||||
|
||||
def is_chunk_cache(self):
|
||||
return False
|
||||
|
||||
def evictable_size(self):
|
||||
return lane0_locs.numel()
|
||||
|
||||
def evict(self, params):
|
||||
self.evict_calls.append(params.num_tokens)
|
||||
allocator.free(lane0_locs)
|
||||
return EvictResult(num_tokens_evicted=lane0_locs.numel())
|
||||
|
||||
tree_cache = FakeTreeCache()
|
||||
_evict_for_compute_owner_lanes(
|
||||
tree_cache=tree_cache,
|
||||
allocator=allocator,
|
||||
page_compute_owners=[0, 0, 0, 0],
|
||||
)
|
||||
|
||||
self.assertGreaterEqual(len(tree_cache.evict_calls), 1)
|
||||
locs = allocator.alloc_extend_compute_owner(
|
||||
prefix_lens=torch.tensor([0], dtype=torch.int64),
|
||||
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
|
||||
seq_lens=torch.tensor([page_size * 4], dtype=torch.int64),
|
||||
seq_lens_cpu=torch.tensor([page_size * 4], dtype=torch.int64),
|
||||
last_loc=torch.tensor([-1], dtype=torch.int64),
|
||||
extend_num_tokens=page_size * 4,
|
||||
page_compute_owners=[0, 0, 0, 0],
|
||||
)
|
||||
self.assertIsNotNone(locs)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -612,5 +612,224 @@ class TestCpSharedKVLazyDebugLogging(unittest.TestCase):
|
||||
self.assertEqual(key_to_write.shape[0], 2)
|
||||
|
||||
|
||||
class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
|
||||
def test_paged_materialize_uses_tai_kernel_when_enabled(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
class FakeTaiKernels:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
|
||||
def materialize_shared_pages(
|
||||
self,
|
||||
page_buffer,
|
||||
logical_pages,
|
||||
*,
|
||||
cp_rank,
|
||||
cp_size,
|
||||
):
|
||||
self.calls.append((page_buffer, logical_pages, cp_rank, cp_size))
|
||||
return (
|
||||
torch.full((logical_pages.numel() + 1, 3), 7, dtype=torch.uint8),
|
||||
torch.tensor([1, 0, 3], dtype=logical_pages.dtype),
|
||||
)
|
||||
|
||||
fake_tai = FakeTaiKernels()
|
||||
layout = CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=1)
|
||||
page_buffer = torch.arange(0, 5 * 3, dtype=torch.uint8).view(5, 3)
|
||||
logical_pages = torch.tensor([1, 0, 3], dtype=torch.int64)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.environ.envs.SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE.get",
|
||||
return_value=True,
|
||||
), patch.object(
|
||||
runtime, "cp_shared_kv_debug_enabled", return_value=False
|
||||
), patch.object(
|
||||
runtime, "_load_tai_materialize_kernels", return_value=fake_tai
|
||||
), patch.object(
|
||||
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
|
||||
):
|
||||
dense_page_buffer, dense_pages = runtime.materialize_shared_paged_buffer(
|
||||
page_buffer=page_buffer,
|
||||
logical_pages=logical_pages,
|
||||
layout=layout,
|
||||
)
|
||||
|
||||
self.assertEqual(len(fake_tai.calls), 1)
|
||||
self.assertIs(fake_tai.calls[0][0], page_buffer)
|
||||
self.assertIs(fake_tai.calls[0][1], logical_pages)
|
||||
self.assertEqual(fake_tai.calls[0][2:], (1, 2))
|
||||
self.assertEqual(dense_pages.tolist(), [1, 0, 3])
|
||||
self.assertEqual(int(dense_page_buffer.sum().item()), 7 * 4 * 3)
|
||||
|
||||
def test_token_materialize_uses_tai_kernel_for_slot_remap_when_enabled(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
class FakeTaiKernels:
|
||||
def __init__(self):
|
||||
self.page_inverse_calls = []
|
||||
self.remap_calls = []
|
||||
self.token_calls = []
|
||||
|
||||
def build_slot_page_inverse(self, slot_logical_pages, logical_page_capacity):
|
||||
self.page_inverse_calls.append((slot_logical_pages, logical_page_capacity))
|
||||
return torch.tensor([0, 1, 2, -1, 3], dtype=torch.long)
|
||||
|
||||
def remap_logical_locs_to_slot_dense_locs(
|
||||
self,
|
||||
logical_locs,
|
||||
page_inverse,
|
||||
*,
|
||||
page_size,
|
||||
):
|
||||
self.remap_calls.append((logical_locs, page_inverse, page_size))
|
||||
return torch.tensor([4, 8, -1], dtype=logical_locs.dtype)
|
||||
|
||||
def materialize_shared_token_kv_pages(
|
||||
self,
|
||||
kv_cache,
|
||||
slot_logical_pages,
|
||||
*,
|
||||
page_size,
|
||||
cp_rank,
|
||||
cp_size,
|
||||
):
|
||||
self.token_calls.append(
|
||||
(kv_cache, slot_logical_pages, page_size, cp_rank, cp_size)
|
||||
)
|
||||
return torch.full((16, 1, 1), 3.0, dtype=kv_cache.dtype)
|
||||
|
||||
fake_tai = FakeTaiKernels()
|
||||
layout = CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=1)
|
||||
kv_cache = torch.arange(0, 24, dtype=torch.float32).view(24, 1, 1)
|
||||
logical_locs = torch.tensor([4, 8, -1], dtype=torch.int64)
|
||||
remap_logical_pages = torch.tensor([[1, 2, 4]], dtype=torch.int64)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.environ.envs.SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE.get",
|
||||
return_value=True,
|
||||
), patch.object(
|
||||
runtime, "cp_shared_kv_debug_enabled", return_value=False
|
||||
), patch.object(
|
||||
runtime, "_load_tai_materialize_kernels", return_value=fake_tai
|
||||
), patch.object(
|
||||
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
|
||||
):
|
||||
dense_kv, dense_locs = runtime.materialize_shared_token_kv_buffer(
|
||||
kv_cache=kv_cache,
|
||||
logical_locs=logical_locs,
|
||||
remap_logical_pages=remap_logical_pages,
|
||||
layout=layout,
|
||||
page_size=4,
|
||||
)
|
||||
|
||||
self.assertEqual(len(fake_tai.page_inverse_calls), 1)
|
||||
self.assertTrue(
|
||||
torch.equal(
|
||||
fake_tai.page_inverse_calls[0][0],
|
||||
remap_logical_pages.reshape(-1),
|
||||
)
|
||||
)
|
||||
self.assertEqual(len(fake_tai.remap_calls), 1)
|
||||
self.assertEqual(len(fake_tai.token_calls), 1)
|
||||
self.assertIs(fake_tai.token_calls[0][0], kv_cache)
|
||||
self.assertTrue(
|
||||
torch.equal(fake_tai.token_calls[0][1], remap_logical_pages.reshape(-1))
|
||||
)
|
||||
self.assertEqual(fake_tai.token_calls[0][2:], (4, 1, 2))
|
||||
self.assertEqual(dense_locs.tolist(), [4, 8, -1])
|
||||
self.assertEqual(float(dense_kv.sum().item()), 48.0)
|
||||
|
||||
def test_token_tai_path_skips_torch_slot_page_remap(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
class FakeTaiKernels:
|
||||
def build_slot_page_inverse(self, slot_logical_pages, logical_page_capacity):
|
||||
return torch.tensor([0, 1, 2, -1, 3], dtype=torch.long)
|
||||
|
||||
def remap_logical_locs_to_slot_dense_locs(
|
||||
self,
|
||||
logical_locs,
|
||||
page_inverse,
|
||||
*,
|
||||
page_size,
|
||||
):
|
||||
return torch.tensor([4, 8, -1], dtype=logical_locs.dtype)
|
||||
|
||||
def materialize_shared_token_kv_pages(
|
||||
self,
|
||||
kv_cache,
|
||||
slot_logical_pages,
|
||||
*,
|
||||
page_size,
|
||||
cp_rank,
|
||||
cp_size,
|
||||
):
|
||||
return torch.full((16, 1, 1), 3.0, dtype=kv_cache.dtype)
|
||||
|
||||
layout = CpSharedKVLayout(page_size=4, cp_size=2, cp_rank=1)
|
||||
kv_cache = torch.arange(0, 24, dtype=torch.float32).view(24, 1, 1)
|
||||
logical_locs = torch.tensor([4, 8, -1], dtype=torch.int64)
|
||||
remap_logical_pages = torch.tensor([[1, 2, 4]], dtype=torch.int64)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.environ.envs.SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE.get",
|
||||
return_value=True,
|
||||
), patch.object(
|
||||
runtime, "cp_shared_kv_debug_enabled", return_value=False
|
||||
), patch.object(
|
||||
runtime, "_load_tai_materialize_kernels", return_value=FakeTaiKernels()
|
||||
), patch.object(
|
||||
runtime,
|
||||
"build_slot_page_remap",
|
||||
side_effect=AssertionError("tai token path must not run torch remap"),
|
||||
), patch.object(
|
||||
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
|
||||
):
|
||||
dense_kv, dense_locs = runtime.materialize_shared_token_kv_buffer(
|
||||
kv_cache=kv_cache,
|
||||
logical_locs=logical_locs,
|
||||
remap_logical_pages=remap_logical_pages,
|
||||
layout=layout,
|
||||
page_size=4,
|
||||
)
|
||||
|
||||
self.assertEqual(dense_locs.tolist(), [4, 8, -1])
|
||||
self.assertEqual(float(dense_kv.sum().item()), 48.0)
|
||||
|
||||
def test_tai_materialize_is_not_used_when_debug_enabled(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
layout = CpSharedKVLayout(page_size=4, cp_size=1, cp_rank=0)
|
||||
page_buffer = torch.arange(0, 4 * 3, dtype=torch.uint8).view(4, 3)
|
||||
logical_pages = torch.tensor([1, 2], dtype=torch.int64)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.environ.envs.SGLANG_CP_SHARED_KV_USE_TAI_MATERIALIZE.get",
|
||||
return_value=True,
|
||||
), patch.object(
|
||||
runtime, "cp_shared_kv_debug_enabled", return_value=True
|
||||
), patch.object(
|
||||
runtime,
|
||||
"_load_tai_materialize_kernels",
|
||||
side_effect=AssertionError("tai path must stay off in debug mode"),
|
||||
), patch.object(
|
||||
runtime, "_all_reduce_materialized_buffer", lambda x, _: x
|
||||
):
|
||||
dense_page_buffer, dense_pages = runtime.materialize_shared_paged_buffer(
|
||||
page_buffer=page_buffer,
|
||||
logical_pages=logical_pages,
|
||||
layout=layout,
|
||||
)
|
||||
|
||||
self.assertEqual(dense_pages.tolist(), [1, 2])
|
||||
self.assertTrue(torch.equal(dense_page_buffer[1], page_buffer[1]))
|
||||
self.assertTrue(torch.equal(dense_page_buffer[2], page_buffer[2]))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user